Compare commits

..
92 changed files with 4823 additions and 3865 deletions
-26
View File
@@ -1,26 +0,0 @@
---
category: Added
---
- **Wait framework capability** — adds the reviewed `Contract.Wait`
declaration (`contract.WaitSpec`) with three execution modes: `poll`
(cadence-poll the leaf's `WaitPoll` hook), `event` (consume the leaf's
`WaitEvents` push stream, correlate events to the accepted resource via
`match_field`/`resource_query`, apply the same terminal map), and `auto`
(event first, fall back to polling when the stream ends or the
subscription fails — one deadline spans both phases). Declared commands
must use the `ResultInvoke` dispatcher; mode and hooks are paired at
construction (poll↔WaitPoll, event↔WaitEvents, auto↔both; surplus hooks
are rejected too). Declared commands register `--wait` /
`--wait-timeout` (framework-owned flags that never enter MCP toolArgs);
undeclared commands reject the flags as unknown. The wait phase closes
the unified envelope exactly once: terminal success → `success`,
terminal failure → `failure` with new wire-stable `error.type: "wait"`
(exit code 8), timeout → `pending` with `meta.operation.timed_out: true`
and the last observed state (exit 0). Deadline exhaustion during a poll,
during event consumption, or between polls always closes as timed-out
pending, never as a poll/stream failure; a correlated event with an
unknown status fails closed exactly like a poll. The capability is
projected into the Schema catalog (`wait` key) alongside `dry_run`. No
business command declares it yet; approval/export/batch adoption lands
separately.
+13 -1
View File
@@ -10,7 +10,7 @@ SCHEMA_META_INDEX_OUTPUT ?= artifacts/schema_meta_index.gob
POLICY_ENV = DWS_POLICY_TMPDIR="$(DWS_POLICY_TMPDIR)" GOTMPDIR="$(POLICY_GOTMPDIR)"
GO_SOURCE_LIST = git ls-files -z --cached --others --exclude-standard -- '*.go'
.PHONY: all help build rebuild test test-plan test-auth-legacy-compat lint format-check fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity skill-context-budget multi-im-skill-chain-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema fetch-mcp-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
.PHONY: all help build check-safechat test-safechat rebuild test test-plan test-auth-legacy-compat lint format-check fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity skill-context-budget multi-im-skill-chain-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema fetch-mcp-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
all: setup-hooks fmt lint build test rebuild
@@ -18,6 +18,8 @@ help:
@printf "Available targets:\n"
@printf " make build - Build the dws CLI binary\n"
@printf " make test - Run the Go test suite\n"
@printf " make check-safechat - Compile and vet the SafeChat message-crypto backend (needs CGO)\n"
@printf " make test-safechat - Run the message-crypto tests against the SafeChat backend\n"
@printf " make test-plan - Verify CI test and full-suite coverage package plans cover their scopes exactly once\n"
@printf " make test-auth-legacy-compat - Run stable legacy authentication compatibility regressions\n"
@printf " make lint - Run formatting checks, go vet, and staticcheck\n"
@@ -52,6 +54,16 @@ build:
rebuild:
@./scripts/dev/build.sh
# No dws command imports internal/msgcrypto yet, so a tagged CLI build would
# link nothing extra and look identical to the default binary. Gate the package
# itself until a caller wires it in.
check-safechat:
@CGO_ENABLED=1 $(GO) build -tags safechat ./internal/msgcrypto/...
@CGO_ENABLED=1 $(GO) vet -tags safechat ./internal/msgcrypto/...
test-safechat:
@CGO_ENABLED=1 $(GO) test -count=1 -tags safechat ./internal/msgcrypto/...
test:
@DWS_PACKAGE_VERSION="$(DWS_PACKAGE_VERSION)" $(GO) test -count=1 -timeout=10m ./...
+2 -3
View File
@@ -482,7 +482,7 @@ Env vars: `DWS_SKILL_MODE=mono|multi` (also honored by `install.sh` / `install.p
<details>
<summary><strong>Personal Event Subscription</strong> — real-time DingTalk messages for event-driven agents</summary>
`dws event consume` subscribes as the currently logged-in user over a managed Stream WebSocket and emits each event as one NDJSON line on stdout. The public catalog covers scoped and all one-to-one/group messages, specified senders, read/recall/reaction events, group lifecycle events, and seven OA approval task/instance events.
`dws event consume` subscribes as the currently logged-in user over a managed Stream WebSocket and emits each event as one NDJSON line on stdout. The public catalog covers scoped and all one-to-one/group messages, specified senders, read/recall/reaction events, group lifecycle events, and six OA approval task/instance events.
The default `ndjson`, `json`, and `pretty` output preserves the transport envelope (`type`, `event_type`, string `data`, and `headers`) for existing scripts; `compact` retains its existing processor. Add `--flatten` to emit the stable top-level business fields used by Agent workflows. `--format` controls JSON serialization; `--flatten` controls the data structure and cannot be combined with `-f raw` or `--debug-raw-events`.
@@ -530,13 +530,12 @@ dws event consume user_im_group_disbanded --group <openConversationId> --flatten
dws event +listen-im --kind sender --user <userId> \
--events message,read,recall -f ndjson
# Listen for all seven public OA approval events in one process
# Listen for all six public OA approval events in one process
dws event consume \
user_oa_approval_task_created \
user_oa_approval_task_finished \
user_oa_approval_task_redirected \
user_oa_approval_instance_started \
user_oa_approval_instance_cc \
user_oa_approval_instance_terminated \
user_oa_approval_instance_finished \
--flatten -f ndjson
+2 -3
View File
@@ -476,7 +476,7 @@ multi setup 或 upgrade 后,DWS 会把官方 bundle 快照和统一所有权
<details>
<summary><strong>个人事件订阅</strong> — 实时接收钉钉消息,驱动事件触发的 Agent</summary>
`dws event consume` 使用当前 OAuth 登录用户建立托管的 Stream WebSocket 长连接,并把每条事件以 NDJSON 一行输出到 stdout。当前公开目录覆盖指定范围和全量单聊/群消息、指定发送人、已读/撤回/表情回应、群生命周期,以及七个 OA 审批任务/实例事件。
`dws event consume` 使用当前 OAuth 登录用户建立托管的 Stream WebSocket 长连接,并把每条事件以 NDJSON 一行输出到 stdout。当前公开目录覆盖指定范围和全量单聊/群消息、指定发送人、已读/撤回/表情回应、群生命周期,以及六个 OA 审批任务/实例事件。
默认 `ndjson`、`json`、`pretty` 输出保留兼容 transport envelope(`type`、`event_type`、字符串 `data`、`headers`),`compact` 继续沿用原 processor。Agent 或新脚本显式加 `--flatten` 后,输出稳定的顶层业务字段。`--format` 控制 JSON 序列化,`--flatten` 控制数据结构,且不能与 `-f raw` 或 `--debug-raw-events` 同时使用。
@@ -524,13 +524,12 @@ dws event consume user_im_group_disbanded --group <openConversationId> --flatten
dws event +listen-im --kind sender --user <userId> \
--events message,read,recall -f ndjson
# 一个进程监听全部七个公开 OA 审批事件
# 一个进程监听全部六个公开 OA 审批事件
dws event consume \
user_oa_approval_task_created \
user_oa_approval_task_finished \
user_oa_approval_task_redirected \
user_oa_approval_instance_started \
user_oa_approval_instance_cc \
user_oa_approval_instance_terminated \
user_oa_approval_instance_finished \
--flatten -f ndjson
@@ -360,7 +360,6 @@ Definition(仅声明;不可编译)
| | `idempotency` | 评审源(或未来 Contract) | reviewed metadata | 今日非框架声明;不得推断 |
| | `effect_source` / provenance | 组装派生物 | resolver 写入 `FieldProvenance` | 派生,不手写 |
| **DryRun** | `preview_kind`, `remote_reads` | 评审源 | `schema_dry_run_capabilities`(正能力声明) | 否;无条目 ≠ 推断「不支持」之外的假能力 |
| **Wait** | `mode`(`poll`/`event`/`auto`), `poll_command`, `status_query`, `terminal`(状态→success/failure), `pending_values`, `event_key`/`match_field`/`resource_query`(event/auto), `default_timeout_secs` | **声明**(`ContractDecl.Wait` 正能力声明,且必须搭配 ResultInvoke dispatcher + 按模式的 hook:poll↔`WaitPoll`、event↔`WaitEvents`、auto↔两者,构造期配对校验,多余 hook 同样拒绝) | 声明后注册 `--wait`/`--wait-timeout`(框架 flag,不进 toolArgs);Schema 投影 `wait` 键;auto = 事件优先、流终止/订阅失败回退轮询,一个 deadline 覆盖两阶段并传入 `WaitPoll`/`WaitEvents`(及 `Command().Context()`);仅 pending 初始结果进入等待,success/failure/partial 原样返回 | 否;未声明命令传 `--wait` = unknown flag。终态失败经统一信封 `error.type: "wait"`(rc=8),超时保持 pending + `meta.operation.timed_out`(rc=0);轮询间/轮询中/事件消费中超时一律按 pending 关闭 |
| **Interface** | `interface_mode`, `interface_ref`, `availability`, `reason` | 评审源 | MCP meta + agent metadata 解析 | 否;与 CLI Identity 分离 |
| **Selection** | `agent_summary`, `use_when`, `avoid_when`, `examples`, `prerequisites`, `tips`, `workflow_refs`, … | 声明(`ContractDecl.Selection` / `ProductDecl`) | `ContractDecl` / `ProductDecl`(`schema_hints/` 已退役) | 可声明;声明载荷**不得携带** `Reviewed`(旧路径专用),携带即组装报错 |
| **FieldProvenance** | 各字段 winner / candidates | 组装派生物 | Schema 组装器 | 派生;须与 delivered value 一致 |
+3
View File
@@ -4,6 +4,8 @@ go 1.25.9
replace gitlab.alibaba-inc.com/aes/aem-go-sdk => ./third_party/aem-go-sdk
replace safechat-go-sdk => ./third_party/safechat-go-sdk
require (
github.com/Microsoft/go-winio v0.6.2
github.com/RealAlexandreAI/json-repair v0.0.15
@@ -24,6 +26,7 @@ require (
golang.org/x/crypto v0.49.0
golang.org/x/sys v0.42.0
golang.org/x/text v0.35.0
safechat-go-sdk v0.0.0
)
require (
+1 -1
View File
@@ -403,7 +403,7 @@ SIGTERM、关 stdin,或先用 dws event stop <subscribe_id> --dry-run 预览
Selection: contract.SelectionSpec{
AgentSummary: "消费 OA、群生命周期或需要底层控制的个人事件流;Agent 通常使用 --flatten 输出 NDJSON",
UseWhen: []string{
"需要监听七个公开 OA 审批任务/实例 EventKey 中的一个或多个事件",
"需要监听六个公开 OA 审批任务/实例 EventKey 中的一个或多个事件",
"需要监听指定群的标题变更、成员进退群或群解散事件",
"用户显式给出原始 EventKey、Filter DSL、subscribe_id,要求原始 transport envelope,或需要普通 IM facade 不提供的高级多事件控制",
},
+2 -2
View File
@@ -152,8 +152,8 @@ func TestCrossPlatformCoveragePersonalSubscriptionProtectionCoversAllPublicEvent
}
}
if publicCount != 23 {
t.Fatalf("public personal events = %d, want 23 (16 IM + 7 OA)", publicCount)
if publicCount != 22 {
t.Fatalf("public personal events = %d, want 22 (16 IM + 6 OA)", publicCount)
}
for _, ruleType := range []string{"at", "all", "singleChat", "sender", "group"} {
if !ruleTypes[ruleType] {
-9
View File
@@ -61,13 +61,6 @@ func TestPersonalOAEventListAndSchemaCommands(t *testing.T) {
"process_code", "title", "status", "create_time", "event_time",
},
},
{
eventKey: personal.EventOAApprovalInstanceCC,
properties: []string{
"type", "event_id", "timestamp", "subscribe_id", "process_instance_id",
"process_code", "title", "status", "create_time", "event_time",
},
},
{
eventKey: personal.EventOAApprovalInstanceTerminated,
properties: []string{
@@ -168,7 +161,6 @@ func TestPersonalOAEventConsumeDryRunAndValidation(t *testing.T) {
personal.EventOAApprovalTaskFinished,
personal.EventOAApprovalTaskRedirected,
personal.EventOAApprovalInstanceStarted,
personal.EventOAApprovalInstanceCC,
personal.EventOAApprovalInstanceTerminated,
personal.EventOAApprovalInstanceFinished,
}
@@ -422,7 +414,6 @@ func TestPersonalOAMultiConsumeCreatesIndependentAllSubscriptionsOnSharedBus(t *
personal.EventOAApprovalTaskFinished,
personal.EventOAApprovalTaskRedirected,
personal.EventOAApprovalInstanceStarted,
personal.EventOAApprovalInstanceCC,
personal.EventOAApprovalInstanceTerminated,
personal.EventOAApprovalInstanceFinished,
}
+1 -67
View File
@@ -51,38 +51,7 @@ func (c *paramAliasCaptureCaller) CallTool(_ context.Context, server, tool strin
func (c *paramAliasCaptureCaller) paramAliasResponseForTool(tool string) string {
switch tool {
case "list_calendar_events":
return `{"success":true,"result":{"events":[],"hasMore":false,"nextCursor":""}}`
case "get_calendar_detail":
return c.paramAliasCalendarDetailResponse()
case "get_calendar_participants":
return `{"success":true,"result":{"participants":[{"userId":"fixture-user","displayName":"Fixture User"},{"userId":"user-2","displayName":"User Two"}]}}`
case "search_calendar":
return `{"success":true,"result":{"calendars":[]}}`
case "search_rooms":
return `{"success":true,"result":{"rooms":[]}}`
case "query_available_meeting_room":
return `{"success":true,"result":{"rooms":[],"hasMore":false}}`
case "list_meeting_room_groups":
return `{"success":true,"result":{"groups":[]}}`
case "query_busy_status":
return `{"success":true,"result":[]}`
case "list_suggested_event_times":
return `{"success":true,"result":{"recommendEventTimes":[]}}`
case "create_calendar_event":
return `{"success":true,"result":{"eventId":"event-1"}}`
case "update_calendar_event", "delete_calendar_event", "add_calendar_participant", "remove_calendar_participant":
return `{"success":true}`
case "respond":
status := "accepted"
if call := c.lastParamAliasCall(); call != nil {
if value, ok := call.args["responseStatus"].(string); ok && value != "" {
status = value
}
}
encoded, _ := json.Marshal(map[string]any{"success": true, "result": map[string]any{"responseStatus": status}})
return string(encoded)
case "get_current_user_profile":
return `{"success":true,"result":{"userId":"user-1","name":"Fixture Current User"}}`
return `{"result":{"events":[]}}`
case "query_records":
return `{"success":true,"status":"success","error":{},"data":{}}`
case "search_mail_users":
@@ -158,41 +127,6 @@ func (c *paramAliasCaptureCaller) paramAliasResponseForTool(tool string) string
}
}
func (c *paramAliasCaptureCaller) lastParamAliasCall() *paramAliasToolCall {
if len(c.calls) == 0 {
return nil
}
return &c.calls[len(c.calls)-1]
}
func (c *paramAliasCaptureCaller) paramAliasCalendarDetailResponse() string {
event := map[string]any{
"eventId": "event-1",
"summary": "Fixture Meeting",
"description": "fixture description",
"startDateTime": "2026-03-10T09:00:00+08:00",
"endDateTime": "2026-03-10T10:00:00+08:00",
}
for _, call := range c.calls {
switch call.tool {
case "create_calendar_event", "update_calendar_event":
for _, key := range []string{"eventId", "summary", "description", "startDateTime", "endDateTime", "timeZone", "location", "freeBusy"} {
if value, ok := call.args[key]; ok {
event[key] = value
}
}
case "respond":
if value, ok := call.args["responseStatus"]; ok {
event["responseStatus"] = value
}
case "delete_calendar_event":
event["status"] = "cancelled"
}
}
encoded, _ := json.Marshal(map[string]any{"success": true, "result": event})
return string(encoded)
}
func (*paramAliasCaptureCaller) Format() string { return "json" }
func (*paramAliasCaptureCaller) DryRun() bool { return false }
func (*paramAliasCaptureCaller) Fields() string { return "" }
@@ -5,8 +5,6 @@ package app
import (
"errors"
"os"
"os/exec"
"reflect"
"strings"
"testing"
@@ -19,8 +17,6 @@ import (
const (
appFixtureCurrentDOpenID = "DAAAAAAAAAAAiE"
appFixtureCurrentDOpenID2 = "DAQEBAQEBAQEiE"
paramAliasCalendarPayloadChildEnv = "DWS_TEST_CALENDAR_PARAM_ALIAS_PAYLOAD_CHILD"
)
// paramAliasCompleteCommands is deliberately keyed by the exact reviewed
@@ -50,37 +46,7 @@ var paramAliasCompleteCommands = map[string][]string{
"aitable workflow run": {"aitable", "workflow", "run", "--base-id", "base-1", "--workflow-id", "workflow-1", "--table-id", "table-1", "--record-ids", "record-1", "--yes"},
"attendance check result": {"attendance", "check", "result", "--users", "user-1,user-2", "--start", "2026-03-01", "--end", "2026-03-02"},
"attendance +check-result": {"attendance", "+check-result", "--users", "user-1,user-2", "--start", "2026-03-01", "--end", "2026-03-02"},
"calendar +agenda": {"calendar", "+agenda", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00", "--calendar-id", "primary", "--cursor", "cursor-1", "--limit", "7"},
"calendar +attendee-list": {"calendar", "+attendee-list", "--event", "event-1", "--calendar-id", "primary"},
"calendar +book": {"calendar", "+book", "--title", "Fixture Meeting", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T10:00:00+08:00", "--with", "Fixture User", "--yes"},
"calendar +book-search": {"calendar", "+book-search", "--query", "fixture"},
"calendar +cancel-event": {"calendar", "+cancel-event", "--event", "event-1", "--yes"},
"calendar +conflicts": {"calendar", "+conflicts", "--in-days", "1"},
"calendar +create": {"calendar", "+create", "--title", "Fixture Meeting", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T10:00:00+08:00", "--desc", "fixture description", "--attendees", "user-1,user-2", "--rooms", "room-1,room-2", "--calendar-id", "primary", "--yes"},
"calendar +free": {"calendar", "+free", "--who", "Fixture User", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00"},
"calendar +free-slots": {"calendar", "+free-slots", "--from", "9", "--to", "18", "--in-days", "1"},
"calendar +freebusy": {"calendar", "+freebusy", "--users", "user-1,user-2", "--rooms", "room-1,room-2", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00"},
"calendar +get": {"calendar", "+get", "--event", "event-1", "--calendar-id", "primary"},
"calendar +invite": {"calendar", "+invite", "--event", "event-1", "--with", "Fixture User", "--yes"},
"calendar +my-free": {"calendar", "+my-free", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00"},
"calendar +reschedule": {"calendar", "+reschedule", "--event", "event-1", "--start", "2026-03-10T10:00:00+08:00", "--end", "2026-03-10T11:00:00+08:00", "--yes"},
"calendar +room-find": {"calendar", "+room-find", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T10:00:00+08:00", "--room-name", "Fixture Room", "--group-id", "group-1", "--page", "1", "--limit", "7"},
"calendar +room-groups": {"calendar", "+room-groups", "--page", "1", "--limit", "7"},
"calendar +room-search": {"calendar", "+room-search", "--room-name", "Fixture Room"},
"calendar +rsvp": {"calendar", "+rsvp", "--event", "event-1", "--status", "accept", "--calendar-id", "primary", "--yes"},
"calendar +search-event": {"calendar", "+search-event", "--query", "fixture", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00", "--calendar-id", "primary", "--cursor", "cursor-1", "--limit", "7"},
"calendar +suggest-time": {"calendar", "+suggest-time", "--with", "Fixture User", "--duration", "30", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00"},
"calendar +suggestion": {"calendar", "+suggestion", "--users", "user-1,user-2", "--duration", "30", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00", "--timezone", "Asia/Shanghai"},
"calendar +update": {"calendar", "+update", "--event", "event-1", "--title", "Fixture Updated Meeting", "--desc", "fixture updated description", "--start", "2026-03-10T10:00:00+08:00", "--end", "2026-03-10T11:00:00+08:00", "--add-attendees", "user-2", "--remove-attendees", "user-1", "--yes"},
"calendar busy search": {"calendar", "busy", "search", "--users", "user-1,user-2", "--rooms", "room-1,room-2", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00"},
"calendar event create": {"calendar", "event", "create", "--title", "Fixture Meeting", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T10:00:00+08:00", "--remind-minutes", "15", "--timezone", "Asia/Shanghai", "--rooms", "room-1,room-2"},
"calendar event list": {"calendar", "event", "list", "--start", "2026-03-10T14:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00", "--calendar-id", "primary", "--cursor", "cursor-1", "--limit", "7"},
"calendar event respond": {"calendar", "event", "respond", "--id", "event-1", "--status", "accepted"},
"calendar event suggest": {"calendar", "event", "suggest", "--users", "user-1,user-2", "--duration", "30", "--start", "2026-03-10T09:00:00+08:00", "--end", "2026-03-10T18:00:00+08:00", "--timezone", "Asia/Shanghai"},
"calendar event update": {"calendar", "event", "update", "--id", "event-1", "--timezone", "Asia/Shanghai"},
"calendar room add": {"calendar", "room", "add", "--event", "event-1", "--rooms", "room-1,room-2"},
"calendar room delete": {"calendar", "room", "delete", "--event", "event-1", "--rooms", "room-1,room-2"},
"calendar room search": {"calendar", "room", "search", "--room-name", "Fixture Room", "--group-id", "group-1", "--start", "2027-03-10T09:00:00+08:00", "--end", "2027-03-10T10:00:00+08:00", "--page", "1", "--limit", "7"},
"chat +chat-messages": {"chat", "+chat-messages", "--group", "fixture-conversation"},
"chat +chat-add-bot": {"chat", "+chat-add-bot", "--id", "fixture-conversation", "--robot-code", "robot-1", "--yes"},
"chat +chat-audit-join": {"chat", "+chat-audit-join", "--group", "fixture-conversation", "--record-id", "7", "--applicant", "user-1", "--inviter", "user-2", "--status", "AuditApprove", "--yes"},
@@ -600,108 +566,6 @@ var paramAliasRepresentativePayloadCases = map[string]bool{
paramAliasPayloadCaseKey("report list", "from-date"): true, // date-range concept alias
}
// paramAliasCalendarPayloadCases keeps the full reviewed Calendar expansion
// separate from the long-lived app-c race process. Each case still executes
// both canonical and alias argv through the real PreParse/Cobra path and
// compares the final captured transport calls; the owning top-level test runs
// these allocations in a short-lived race-instrumented subprocess so all Root
// registrations are released together when that process exits.
var paramAliasCalendarPayloadCases = map[string]bool{
paramAliasPayloadCaseKey("calendar +agenda", "from"): true,
paramAliasPayloadCaseKey("calendar +agenda", "to"): true,
paramAliasPayloadCaseKey("calendar +agenda", "max-results"): true,
paramAliasPayloadCaseKey("calendar +agenda", "next-cursor"): true,
paramAliasPayloadCaseKey("calendar +agenda", "calendar-book-id"): true,
paramAliasPayloadCaseKey("calendar +attendee-list", "event-id"): true,
paramAliasPayloadCaseKey("calendar +attendee-list", "calendar-book-id"): true,
paramAliasPayloadCaseKey("calendar +book", "summary"): true,
paramAliasPayloadCaseKey("calendar +book", "attendee-names"): true,
paramAliasPayloadCaseKey("calendar +book-search", "keyword"): true,
paramAliasPayloadCaseKey("calendar +book-search", "search"): true,
paramAliasPayloadCaseKey("calendar +book-search", "name"): true,
paramAliasPayloadCaseKey("calendar +cancel-event", "event-id"): true,
paramAliasPayloadCaseKey("calendar +cancel-event", "id"): true,
paramAliasPayloadCaseKey("calendar +free", "name"): true,
paramAliasPayloadCaseKey("calendar +free-slots", "start-hour"): true,
paramAliasPayloadCaseKey("calendar +free-slots", "end-hour"): true,
paramAliasPayloadCaseKey("calendar +free-slots", "day-offset"): true,
paramAliasPayloadCaseKey("calendar +freebusy", "user-ids"): true,
paramAliasPayloadCaseKey("calendar +freebusy", "room-ids"): true,
paramAliasPayloadCaseKey("calendar +freebusy", "room-id"): true,
paramAliasPayloadCaseKey("calendar +my-free", "from"): true,
paramAliasPayloadCaseKey("calendar +my-free", "to"): true,
paramAliasPayloadCaseKey("calendar +invite", "id"): true,
paramAliasPayloadCaseKey("calendar +invite", "participant-names"): true,
paramAliasPayloadCaseKey("calendar +reschedule", "id"): true,
paramAliasPayloadCaseKey("calendar +reschedule", "from"): true,
paramAliasPayloadCaseKey("calendar +reschedule", "to"): true,
paramAliasPayloadCaseKey("calendar +room-groups", "page-size"): true,
paramAliasPayloadCaseKey("calendar +room-groups", "page-index"): true,
paramAliasPayloadCaseKey("calendar +room-search", "query"): true,
paramAliasPayloadCaseKey("calendar +suggest-time", "duration-minutes"): true,
paramAliasPayloadCaseKey("calendar +suggest-time", "attendee-names"): true,
paramAliasPayloadCaseKey("calendar +conflicts", "day-offset"): true,
paramAliasPayloadCaseKey("calendar busy search", "room-id"): true,
paramAliasPayloadCaseKey("calendar event create", "reminder-minutes"): true,
paramAliasPayloadCaseKey("calendar event create", "tz"): true,
paramAliasPayloadCaseKey("calendar event create", "room-id"): true,
paramAliasPayloadCaseKey("calendar event respond", "response-status"): true,
paramAliasPayloadCaseKey("calendar event suggest", "duration-minutes"): true,
paramAliasPayloadCaseKey("calendar event update", "tz"): true,
paramAliasPayloadCaseKey("calendar room add", "room-id"): true,
paramAliasPayloadCaseKey("calendar room delete", "room-id"): true,
paramAliasPayloadCaseKey("calendar room search", "room-group-id"): true,
paramAliasPayloadCaseKey("calendar +create", "summary"): true,
paramAliasPayloadCaseKey("calendar +create", "description"): true,
paramAliasPayloadCaseKey("calendar +create", "user-ids"): true,
paramAliasPayloadCaseKey("calendar +create", "room-ids"): true,
paramAliasPayloadCaseKey("calendar +create", "room-id"): true,
paramAliasPayloadCaseKey("calendar +create", "calendar-book-id"): true,
paramAliasPayloadCaseKey("calendar +create", "to"): true,
paramAliasPayloadCaseKey("calendar +create", "from"): true,
paramAliasPayloadCaseKey("calendar +get", "event-id"): true,
paramAliasPayloadCaseKey("calendar +get", "calendar-book-id"): true,
paramAliasPayloadCaseKey("calendar +room-find", "from"): true,
paramAliasPayloadCaseKey("calendar +room-find", "to"): true,
paramAliasPayloadCaseKey("calendar +room-find", "page-size"): true,
paramAliasPayloadCaseKey("calendar +room-find", "page-index"): true,
paramAliasPayloadCaseKey("calendar +room-find", "room-group-id"): true,
paramAliasPayloadCaseKey("calendar +room-find", "query"): true,
paramAliasPayloadCaseKey("calendar +rsvp", "event-id"): true,
paramAliasPayloadCaseKey("calendar +rsvp", "response-status"): true,
paramAliasPayloadCaseKey("calendar +search-event", "keyword"): true,
paramAliasPayloadCaseKey("calendar +search-event", "from"): true,
paramAliasPayloadCaseKey("calendar +search-event", "to"): true,
paramAliasPayloadCaseKey("calendar +search-event", "next-cursor"): true,
paramAliasPayloadCaseKey("calendar +search-event", "max-results"): true,
paramAliasPayloadCaseKey("calendar +suggestion", "user-ids"): true,
paramAliasPayloadCaseKey("calendar +suggestion", "duration-minutes"): true,
paramAliasPayloadCaseKey("calendar +suggestion", "from"): true,
paramAliasPayloadCaseKey("calendar +suggestion", "to"): true,
paramAliasPayloadCaseKey("calendar +suggestion", "tz"): true,
paramAliasPayloadCaseKey("calendar +update", "event-id"): true,
paramAliasPayloadCaseKey("calendar +update", "from"): true,
paramAliasPayloadCaseKey("calendar +update", "summary"): true,
paramAliasPayloadCaseKey("calendar +update", "description"): true,
paramAliasPayloadCaseKey("calendar +update", "add-user-ids"): true,
paramAliasPayloadCaseKey("calendar +update", "remove-user-ids"): true,
}
// paramAliasCalendarConfirmationCases selects one newly reviewed alias for
// every Calendar Shortcut whose runtime contract requires user confirmation.
// The complete Calendar matrix proves confirmed canonical/alias payload
// equality; these representatives additionally prove semantic normalization
// cannot cross the confirmation boundary before the first transport call.
var paramAliasCalendarConfirmationCases = map[string]bool{
paramAliasPayloadCaseKey("calendar +book", "summary"): true,
paramAliasPayloadCaseKey("calendar +cancel-event", "event-id"): true,
paramAliasPayloadCaseKey("calendar +create", "summary"): true,
paramAliasPayloadCaseKey("calendar +invite", "id"): true,
paramAliasPayloadCaseKey("calendar +reschedule", "from"): true,
paramAliasPayloadCaseKey("calendar +rsvp", "response-status"): true,
paramAliasPayloadCaseKey("calendar +update", "event-id"): true,
}
func TestCrossPlatformCoverageReviewedParamAliasesHaveCompleteTemplatesAndRepresentativeFinalPayloads(t *testing.T) {
concepts, err := cli.LoadParamConcepts()
if err != nil {
@@ -735,7 +599,28 @@ func TestCrossPlatformCoverageReviewedParamAliasesHaveCompleteTemplatesAndRepres
}
executedRepresentatives[caseKey] = true
t.Run(fixture.Command+"/"+fixture.Emitted, func(t *testing.T) {
assertParamAliasFinalPayloadEquivalent(t, fixture.Command, canonicalArgs, aliasArgs)
canonicalCaller := &paramAliasCaptureCaller{}
_, canonicalErr := executeParamAliasPayloadE2E(t, canonicalCaller, canonicalArgs...)
if canonicalErr != nil {
t.Fatalf("complete canonical command failed: %v\nargs=%v\ncalls=%#v", canonicalErr, canonicalArgs, canonicalCaller.calls)
}
if len(canonicalCaller.calls) == 0 {
t.Fatalf("complete canonical command reached no final transport payload: args=%v", canonicalArgs)
}
aliasCaller := &paramAliasCaptureCaller{}
ctx, aliasErr := executeParamAliasPayloadE2E(t, aliasCaller, aliasArgs...)
if aliasErr != nil {
t.Fatalf("complete alias command failed: %v\nargs=%v\ncalls=%#v", aliasErr, aliasArgs, aliasCaller.calls)
}
if ctx == nil {
t.Fatal("complete alias command skipped PreParse")
}
normalizeParamAliasVolatileDefaults(fixture.Command, canonicalCaller, aliasCaller)
if !reflect.DeepEqual(aliasCaller.calls, canonicalCaller.calls) {
t.Fatalf("final transport calls differ\ncanonical args: %v\nalias args: %v\ncanonical calls: %#v\nalias calls: %#v", canonicalArgs, aliasArgs, canonicalCaller.calls, aliasCaller.calls)
}
})
}
@@ -765,118 +650,6 @@ func TestCrossPlatformCoverageReviewedParamAliasesHaveCompleteTemplatesAndRepres
}
}
func TestCrossPlatformCoverageReviewedCalendarParamAliasesReachCanonicalEquivalentFinalPayloads(t *testing.T) {
if os.Getenv(paramAliasCalendarPayloadChildEnv) != "1" {
command := exec.Command(
os.Args[0],
"-test.run=^TestCrossPlatformCoverageReviewedCalendarParamAliasesReachCanonicalEquivalentFinalPayloads$",
"-test.count=1",
"-test.timeout=5m",
)
command.Env = append(os.Environ(), paramAliasCalendarPayloadChildEnv+"=1")
output, err := command.CombinedOutput()
if err != nil {
t.Fatalf("Calendar param-alias payload subprocess failed: %v\n%s", err, strings.TrimSpace(string(output)))
}
return
}
concepts, err := cli.LoadParamConcepts()
if err != nil {
t.Fatalf("LoadParamConcepts() error = %v", err)
}
executed := make(map[string]bool)
executedConfirmation := make(map[string]bool)
for _, fixture := range concepts.Fixture {
caseKey := paramAliasPayloadCaseKey(fixture.Command, fixture.Emitted)
if !paramAliasCalendarPayloadCases[caseKey] {
continue
}
executed[caseKey] = true
fixture := fixture
t.Run(fixture.Command+"/"+fixture.Emitted, func(t *testing.T) {
complete, ok := paramAliasCompleteCommand(fixture.Command, fixture.Expect)
if !ok {
t.Fatal("reviewed Calendar alias has no complete-command E2E template")
}
canonicalArgs := append([]string(nil), complete...)
aliasArgs, replacements := replaceLongFlag(canonicalArgs, fixture.Expect, fixture.Emitted)
if replacements != 1 {
t.Fatalf("complete Calendar command must contain canonical --%s exactly once; replacements=%d args=%v", fixture.Expect, replacements, canonicalArgs)
}
assertParamAliasFinalPayloadEquivalent(t, fixture.Command, canonicalArgs, aliasArgs)
if paramAliasCalendarConfirmationCases[caseKey] {
executedConfirmation[caseKey] = true
assertParamAliasCannotBypassConfirmation(t, aliasArgs)
}
})
}
for caseKey := range paramAliasCalendarPayloadCases {
if !executed[caseKey] {
t.Errorf("Calendar final-payload case %q has no active reviewed fixture", caseKey)
}
}
if len(executed) != len(paramAliasCalendarPayloadCases) {
t.Fatalf("Calendar final-payload coverage = %d, want %d", len(executed), len(paramAliasCalendarPayloadCases))
}
for caseKey := range paramAliasCalendarConfirmationCases {
if !executedConfirmation[caseKey] {
t.Errorf("Calendar confirmation case %q has no active reviewed fixture", caseKey)
}
}
if len(executedConfirmation) != len(paramAliasCalendarConfirmationCases) {
t.Fatalf("Calendar confirmation coverage = %d, want %d", len(executedConfirmation), len(paramAliasCalendarConfirmationCases))
}
}
func assertParamAliasCannotBypassConfirmation(t *testing.T, aliasArgs []string) {
t.Helper()
unconfirmedArgs, removals := removeExactArg(aliasArgs, "--yes")
if removals != 1 {
t.Fatalf("confirmation template must contain --yes exactly once; removals=%d args=%v", removals, aliasArgs)
}
caller := &paramAliasCaptureCaller{}
ctx, err := executeParamAliasPayloadE2E(t, caller, unconfirmedArgs...)
if ctx == nil {
t.Fatal("unconfirmed Calendar alias command skipped PreParse")
}
var appErr *apperrors.Error
if !errors.As(err, &appErr) || appErr.Reason != "confirmation_required" {
t.Fatalf("unconfirmed Calendar alias command error = %#v, want confirmation_required\nargs=%v", err, unconfirmedArgs)
}
if len(caller.calls) != 0 {
t.Fatalf("unconfirmed Calendar alias crossed the transport boundary: args=%v calls=%#v", unconfirmedArgs, caller.calls)
}
}
func assertParamAliasFinalPayloadEquivalent(t *testing.T, command string, canonicalArgs, aliasArgs []string) {
t.Helper()
canonicalCaller := &paramAliasCaptureCaller{}
_, canonicalErr := executeParamAliasPayloadE2E(t, canonicalCaller, canonicalArgs...)
if canonicalErr != nil {
t.Fatalf("complete canonical command failed: %v\nargs=%v\ncalls=%#v", canonicalErr, canonicalArgs, canonicalCaller.calls)
}
if len(canonicalCaller.calls) == 0 {
t.Fatalf("complete canonical command reached no final transport payload: args=%v", canonicalArgs)
}
aliasCaller := &paramAliasCaptureCaller{}
ctx, aliasErr := executeParamAliasPayloadE2E(t, aliasCaller, aliasArgs...)
if aliasErr != nil {
t.Fatalf("complete alias command failed: %v\nargs=%v\ncalls=%#v", aliasErr, aliasArgs, aliasCaller.calls)
}
if ctx == nil {
t.Fatal("complete alias command skipped PreParse")
}
normalizeParamAliasVolatileDefaults(command, canonicalCaller, aliasCaller)
if !reflect.DeepEqual(aliasCaller.calls, canonicalCaller.calls) {
t.Fatalf("final transport calls differ\ncanonical args: %v\nalias args: %v\ncanonical calls: %#v\nalias calls: %#v", canonicalArgs, aliasArgs, canonicalCaller.calls, aliasCaller.calls)
}
}
func TestCrossPlatformCoverageNewIMParamAliasesReachCanonicalEquivalentFinalPayloads(t *testing.T) {
activeAliases := 0
for _, test := range paramAliasNewIMCases {
+4 -3
View File
@@ -113,9 +113,10 @@ const (
ClientIDPath = "/cli/clientId"
// MCP OAuth endpoints (used when clientId is fetched from MCP).
MCPOAuthTokenPath = "/oauth2/getToken"
MCPRefreshTokenPath = "/oauth2/refreshToken"
MCPRevokeTokenPath = "/oauth2/revokeToken"
MCPOAuthTokenPath = "/oauth2/getToken"
MCPRefreshTokenPath = "/oauth2/refreshToken"
MCPRevokeTokenPath = "/oauth2/revokeToken"
MCPVendorAuthCodePath = "/oauth2/vendorAuthCode"
// App-level access token endpoints (for dws api raw calls).
+192
View File
@@ -0,0 +1,192 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
headerUserAccessToken = "x-user-access-token"
headerDWSClientID = "x-dws-client-id"
headerDWSCLIVersion = "x-dws-cli-version"
// VendorAuthCode default / documented portal expiresIn, used only when
// the success body omits a positive value. Callers must still prefer
// the response field.
DefaultVendorAuthCodeExpiresIn = 120
VendorAuthCodeParamError = "PARAM_ERROR"
VendorAuthCodeVendorUnsupported = "VENDOR_UNSUPPORTED"
VendorAuthCodeTokenInvalid = "TOKEN_INVALID"
VendorAuthCodeOrgMismatch = "ORG_MISMATCH"
VendorAuthCodeUserNotInOrg = "USER_NOT_IN_ORG"
VendorAuthCodeVendorNotEnabled = "VENDOR_NOT_ENABLED"
VendorAuthCodeRateLimited = "RATE_LIMITED"
VendorAuthCodeInternalError = "INTERNAL_ERROR"
)
// VendorAuthCodeInput is a POST /oauth2/vendorAuthCode call. The body is
// only vendor + corpId; redirectURI and domain must not be sent.
type VendorAuthCodeInput struct {
AccessToken string
ClientID string
CLIVersion string
LoginRegion LoginRegion
Vendor string
CorpID string
HTTPClient *http.Client
// BaseURL overrides MCPBaseURLForLoginRegion. Tests use it; production
// callers leave it empty.
BaseURL string
}
// VendorAuthCodeResult is the success VO from portal.
type VendorAuthCodeResult struct {
AuthCode string
ExpiresIn int
}
// VendorAuthCodeError is a portal business error carried in an HTTP 200
// ServiceResult body (same envelope as /oauth2/getToken).
type VendorAuthCodeError struct {
Code string
Message string
}
func (e *VendorAuthCodeError) Error() string {
if e == nil {
return "vendorAuthCode failed"
}
if strings.TrimSpace(e.Message) != "" {
return fmt.Sprintf("vendorAuthCode %s: %s", e.Code, e.Message)
}
return fmt.Sprintf("vendorAuthCode %s", e.Code)
}
// Retryable reports whether DWS should retry this portal error once.
func (e *VendorAuthCodeError) Retryable() bool {
if e == nil {
return false
}
switch e.Code {
case VendorAuthCodeTokenInvalid, VendorAuthCodeRateLimited, VendorAuthCodeInternalError:
return true
default:
return false
}
}
// FetchVendorAuthCode POSTs {vendor, corpId} to /oauth2/vendorAuthCode.
// HTTP is expected to be 200; errors are read from body.errorCode.
func FetchVendorAuthCode(ctx context.Context, in VendorAuthCodeInput) (*VendorAuthCodeResult, error) {
vendor := strings.ToLower(strings.TrimSpace(in.Vendor))
corpID := strings.TrimSpace(in.CorpID)
token := strings.TrimSpace(in.AccessToken)
clientID := strings.TrimSpace(in.ClientID)
if token == "" || clientID == "" || vendor == "" || corpID == "" {
return nil, &VendorAuthCodeError{
Code: VendorAuthCodeParamError,
Message: "token, clientId, vendor and corpId are required",
}
}
base := strings.TrimRight(strings.TrimSpace(in.BaseURL), "/")
if base == "" {
base = strings.TrimRight(MCPBaseURLForLoginRegion(in.LoginRegion), "/")
}
endpoint := base + MCPVendorAuthCodePath
payload, err := json.Marshal(struct {
Vendor string `json:"vendor"`
CorpID string `json:"corpId"`
}{Vendor: vendor, CorpID: corpID})
if err != nil {
return nil, fmt.Errorf("marshaling vendorAuthCode request: %w", err)
}
req, err := oauthNewRequest(ctx, http.MethodPost, endpoint, bytes.NewReader(payload))
if err != nil {
return nil, fmt.Errorf("creating vendorAuthCode request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set(headerUserAccessToken, token)
req.Header.Set(headerDWSClientID, clientID)
if ver := strings.TrimSpace(in.CLIVersion); ver != "" {
req.Header.Set(headerDWSCLIVersion, ver)
}
applyEditionEnterpriseCredentialHeaders(req)
client := in.HTTPClient
if client == nil {
client = oauthHTTPClient
}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("sending vendorAuthCode request: %w", err)
}
defer resp.Body.Close()
data, readErr := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if resp.StatusCode != http.StatusOK {
if readErr != nil {
data = nil
}
return nil, &HTTPStatusError{
StatusCode: resp.StatusCode,
responseBody: truncateBody(data, 200),
}
}
if readErr != nil {
return nil, fmt.Errorf("reading vendorAuthCode response: %w", readErr)
}
return parseVendorAuthCodeResponse(data)
}
func parseVendorAuthCodeResponse(body []byte) (*VendorAuthCodeResult, error) {
var resp struct {
AuthCode string `json:"authCode"`
ExpiresIn int `json:"expiresIn"`
Success *bool `json:"success"`
ErrorCode string `json:"errorCode"`
ErrorMsg string `json:"errorMsg"`
}
if err := json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("parsing vendorAuthCode response: %w", err)
}
if resp.ErrorCode != "" || resp.ErrorMsg != "" || (resp.Success != nil && !*resp.Success) {
code := strings.TrimSpace(resp.ErrorCode)
if code == "" {
code = VendorAuthCodeInternalError
}
return nil, &VendorAuthCodeError{Code: code, Message: resp.ErrorMsg}
}
authCode := strings.TrimSpace(resp.AuthCode)
if authCode == "" {
return nil, fmt.Errorf("vendorAuthCode response missing authCode")
}
expiresIn := resp.ExpiresIn
if expiresIn <= 0 {
expiresIn = DefaultVendorAuthCodeExpiresIn
}
return &VendorAuthCodeResult{AuthCode: authCode, ExpiresIn: expiresIn}, nil
}
+178
View File
@@ -0,0 +1,178 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestFetchVendorAuthCodeSuccessEnvelope(t *testing.T) {
var gotBody map[string]string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != MCPVendorAuthCodePath {
t.Fatalf("path = %q, want %s", r.URL.Path, MCPVendorAuthCodePath)
}
if r.Method != http.MethodPost {
t.Fatalf("method = %s, want POST", r.Method)
}
if got := r.Header.Get("x-user-access-token"); got != "user-token" {
t.Fatalf("x-user-access-token = %q", got)
}
if got := r.Header.Get("x-dws-client-id"); got != "dws-client" {
t.Fatalf("x-dws-client-id = %q", got)
}
if got := r.Header.Get("x-dws-cli-version"); got != "1.2.3" {
t.Fatalf("x-dws-cli-version = %q", got)
}
if err := json.NewDecoder(r.Body).Decode(&gotBody); err != nil {
t.Fatalf("decode body: %v", err)
}
w.Header().Set("Cache-Control", "no-store")
_, _ = io.WriteString(w, `{"authCode":"tmp-code","expiresIn":120}`)
}))
defer srv.Close()
got, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
CLIVersion: "1.2.3",
Vendor: "SafeChat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
if err != nil {
t.Fatalf("FetchVendorAuthCode() = %v", err)
}
if got.AuthCode != "tmp-code" || got.ExpiresIn != 120 {
t.Fatalf("result = %+v", got)
}
if gotBody["vendor"] != "safechat" || gotBody["corpId"] != "dingxxxxxxxxxxxx" {
t.Fatalf("posted body = %v", gotBody)
}
if _, ok := gotBody["redirectURI"]; ok {
t.Fatalf("posted redirectURI, body = %v", gotBody)
}
if _, ok := gotBody["domain"]; ok {
t.Fatalf("posted domain, body = %v", gotBody)
}
}
func TestFetchVendorAuthCodeParsesAlways200ServiceResult(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = io.WriteString(w, `{"success":false,"errorCode":"VENDOR_NOT_ENABLED","errorMsg":"not installed"}`)
}))
defer srv.Close()
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
var verr *VendorAuthCodeError
if !errors.As(err, &verr) || verr.Code != VendorAuthCodeVendorNotEnabled {
t.Fatalf("error = %v, want VENDOR_NOT_ENABLED", err)
}
if verr.Retryable() {
t.Fatal("VENDOR_NOT_ENABLED must not be retryable")
}
}
func TestFetchVendorAuthCodeRequiresLocalFields(t *testing.T) {
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{Vendor: "safechat", CorpID: "ding"})
var verr *VendorAuthCodeError
if !errors.As(err, &verr) || verr.Code != VendorAuthCodeParamError {
t.Fatalf("error = %v, want PARAM_ERROR", err)
}
}
func TestFetchVendorAuthCodeKeepsHTTPStatusOnNon200(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
http.Error(w, "oops", http.StatusBadGateway)
}))
defer srv.Close()
_, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
var statusErr *HTTPStatusError
if !errors.As(err, &statusErr) || statusErr.StatusCode != http.StatusBadGateway {
t.Fatalf("error = %v, want HTTP 502", err)
}
}
func TestParseVendorAuthCodeResponseDefaultsExpiresIn(t *testing.T) {
got, err := parseVendorAuthCodeResponse([]byte(`{"authCode":"x"}`))
if err != nil {
t.Fatalf("parse: %v", err)
}
if got.ExpiresIn != DefaultVendorAuthCodeExpiresIn {
t.Fatalf("expiresIn = %d, want %d", got.ExpiresIn, DefaultVendorAuthCodeExpiresIn)
}
}
func TestVendorAuthCodeErrorRetryable(t *testing.T) {
for _, code := range []string{VendorAuthCodeTokenInvalid, VendorAuthCodeRateLimited, VendorAuthCodeInternalError} {
if !(&VendorAuthCodeError{Code: code}).Retryable() {
t.Fatalf("%s should be retryable", code)
}
}
if (&VendorAuthCodeError{Code: VendorAuthCodeOrgMismatch}).Retryable() {
t.Fatal("ORG_MISMATCH must not be retryable")
}
}
func TestFetchVendorAuthCodeDoesNotSendBlankCLIVersionHeader(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if got := r.Header.Get("x-dws-cli-version"); got != "" {
t.Fatalf("x-dws-cli-version = %q, want empty", got)
}
_, _ = io.WriteString(w, `{"authCode":"tmp-code","expiresIn":90}`)
}))
defer srv.Close()
got, err := FetchVendorAuthCode(context.Background(), VendorAuthCodeInput{
AccessToken: "user-token",
ClientID: "dws-client",
Vendor: "safechat",
CorpID: "dingxxxxxxxxxxxx",
HTTPClient: srv.Client(),
BaseURL: srv.URL,
})
if err != nil {
t.Fatalf("FetchVendorAuthCode() = %v", err)
}
if got.ExpiresIn != 90 {
t.Fatalf("expiresIn = %d, want 90 from response", got.ExpiresIn)
}
if strings.Contains(got.AuthCode, "redirect") {
t.Fatalf("unexpected code %q", got.AuthCode)
}
}
-481
View File
@@ -1773,439 +1773,6 @@ var generatedParamAliases = []ParamAliasEntry{
},
Blocked: []string{"at-user-ids", "staff-id", "uid", "user", "user-id", "userid"},
},
{
CLIPath: "calendar +agenda",
Aliases: map[string]string{
"begin": "start",
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-result": "limit",
"max-results": "limit",
"max-time": "end",
"min-time": "start",
"next-cursor": "cursor",
"next-page-token": "cursor",
"next-token": "cursor",
"page-size": "limit",
"page-token": "cursor",
"per-page": "limit",
"since": "start",
"size": "limit",
"start-date": "start",
"start-time": "start",
"take": "limit",
"time-max": "end",
"time-min": "start",
"to": "end",
"top": "limit",
},
Blocked: []string{"acl-id", "count", "date", "event", "event-id", "id", "offset", "page", "room-id", "time"},
},
{
CLIPath: "calendar +attendee-list",
Aliases: map[string]string{
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"calendar-event-id": "event",
"event-id": "event",
},
Blocked: []string{"acl-id", "room-id"},
Ambiguous: []string{"id"},
},
{
CLIPath: "calendar +book",
Aliases: map[string]string{
"attendee-names": "with",
"begin": "start",
"end-time": "end",
"from": "start",
"names": "with",
"participant-names": "with",
"start-time": "start",
"subject": "title",
"summary": "title",
"to": "end",
},
Blocked: []string{"attendees", "calendar-id", "calendar-name", "date", "end-date", "name", "open-dingtalk-ids", "participants", "room-id", "room-ids", "room-name", "rooms", "start-date", "time", "time-max", "time-min", "user", "user-id", "user-ids", "users"},
},
{
CLIPath: "calendar +book-search",
Aliases: map[string]string{
"keyword": "query",
"keywords": "query",
"name": "query",
"q": "query",
"search": "query",
"search-word": "query",
},
Blocked: []string{"subject", "text", "title"},
},
{
CLIPath: "calendar +cancel-event",
Aliases: map[string]string{
"calendar-event-id": "event",
"event-id": "event",
"id": "event",
},
Blocked: []string{"acl-id", "calendar-book-id", "calendar-id", "room-id"},
},
{
CLIPath: "calendar +conflicts",
Aliases: map[string]string{
"day-offset": "in-days",
"days-from-today": "in-days",
},
Blocked: []string{"days", "duration", "end", "from", "start", "to"},
},
{
CLIPath: "calendar +create",
Aliases: map[string]string{
"attendee-ids": "attendees",
"begin": "start",
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"description": "desc",
"end-time": "end",
"freebusy": "free-busy",
"from": "start",
"room-id": "rooms",
"room-ids": "rooms",
"start-time": "start",
"subject": "title",
"summary": "title",
"time-zone": "timezone",
"to": "end",
"tz": "timezone",
"user-ids": "attendees",
"users": "attendees",
},
Blocked: []string{"acl-id", "attendee-name", "attendee-names", "calendar-name", "config", "date", "end-date", "event", "event-id", "field-description", "group-id", "id", "locale", "name", "offset", "open-dingtalk-ids", "participant-name", "participant-names", "rich-text-desc", "room", "room-name", "start-date", "time", "time-max", "time-min", "utc-offset", "who", "with"},
},
{
CLIPath: "calendar +free",
Aliases: map[string]string{
"begin": "start",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-time": "end",
"min-time": "start",
"name": "who",
"person": "who",
"person-name": "who",
"since": "start",
"start-date": "start",
"start-time": "start",
"time-max": "end",
"time-min": "start",
"to": "end",
},
Blocked: []string{"date", "time", "user", "user-id", "user-ids", "users", "with"},
},
{
CLIPath: "calendar +free-slots",
Aliases: map[string]string{
"day-offset": "in-days",
"days-from-today": "in-days",
"end-hour": "to",
"start-hour": "from",
},
Blocked: []string{"days", "duration", "end", "end-time", "start", "start-time", "time-max", "time-min"},
},
{
CLIPath: "calendar +freebusy",
Aliases: map[string]string{
"begin": "start",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-time": "end",
"min-time": "start",
"room-id": "rooms",
"room-ids": "rooms",
"since": "start",
"start-date": "start",
"start-time": "start",
"time-max": "end",
"time-min": "start",
"to": "end",
"user-ids": "users",
},
Blocked: []string{"at-user-ids", "attendee-names", "date", "group-id", "location", "name", "names", "participant-names", "room", "room-name", "staff-id", "time", "uid", "user", "user-id", "userid", "who", "with"},
},
{
CLIPath: "calendar +get",
Aliases: map[string]string{
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"calendar-event-id": "event",
"event-id": "event",
},
Blocked: []string{"acl-id", "room-id"},
Ambiguous: []string{"id"},
},
{
CLIPath: "calendar +invite",
Aliases: map[string]string{
"attendee-names": "with",
"calendar-event-id": "event",
"event-id": "event",
"id": "event",
"names": "with",
"participant-names": "with",
},
Blocked: []string{"acl-id", "attendees", "calendar-book-id", "calendar-id", "open-dingtalk-ids", "participants", "room-id", "user", "user-id", "user-ids", "users"},
},
{
CLIPath: "calendar +my-free",
Aliases: map[string]string{
"begin": "start",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-time": "end",
"min-time": "start",
"since": "start",
"start-date": "start",
"start-time": "start",
"time-max": "end",
"time-min": "start",
"to": "end",
},
Blocked: []string{"date", "time"},
},
{
CLIPath: "calendar +reschedule",
Aliases: map[string]string{
"begin": "start",
"calendar-event-id": "event",
"end-time": "end",
"event-id": "event",
"from": "start",
"id": "event",
"start-time": "start",
"to": "end",
},
Blocked: []string{"acl-id", "calendar-book-id", "calendar-id", "date", "end-date", "room-id", "start-date", "time", "time-max", "time-min"},
},
{
CLIPath: "calendar +room-find",
Aliases: map[string]string{
"begin": "start",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"group": "group-id",
"max-result": "limit",
"max-results": "limit",
"max-time": "end",
"min-time": "start",
"name": "room-name",
"page-index": "page",
"page-size": "limit",
"per-page": "limit",
"query": "room-name",
"room-group-id": "group-id",
"since": "start",
"size": "limit",
"start-date": "start",
"start-time": "start",
"take": "limit",
"time-max": "end",
"time-min": "start",
"to": "end",
"top": "limit",
},
Blocked: []string{"count", "cursor", "date", "location", "next-cursor", "page-token", "room", "room-id", "room-ids", "rooms", "time"},
},
{
CLIPath: "calendar +room-groups",
Aliases: map[string]string{
"max-result": "limit",
"max-results": "limit",
"page-index": "page",
"page-size": "limit",
"per-page": "limit",
"size": "limit",
"take": "limit",
"top": "limit",
},
Blocked: []string{"count", "cursor", "page-token"},
},
{
CLIPath: "calendar +room-search",
Aliases: map[string]string{
"name": "room-name",
"query": "room-name",
},
Blocked: []string{"group-id", "location", "room", "room-id", "room-ids", "rooms"},
},
{
CLIPath: "calendar +rsvp",
Aliases: map[string]string{
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"calendar-event-id": "event",
"event-id": "event",
"response": "status",
"response-status": "status",
},
Blocked: []string{"acl-id", "availability", "done", "free-busy", "room-id", "state"},
Ambiguous: []string{"id"},
},
{
CLIPath: "calendar +search-event",
Aliases: map[string]string{
"begin": "start",
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"keyword": "query",
"keywords": "query",
"max-result": "limit",
"max-results": "limit",
"max-time": "end",
"min-time": "start",
"next-cursor": "cursor",
"next-page-token": "cursor",
"next-token": "cursor",
"page-size": "limit",
"page-token": "cursor",
"per-page": "limit",
"q": "query",
"search": "query",
"search-word": "query",
"since": "start",
"size": "limit",
"start-date": "start",
"start-time": "start",
"take": "limit",
"time-max": "end",
"time-min": "start",
"to": "end",
"top": "limit",
},
Blocked: []string{"acl-id", "count", "date", "event", "event-id", "id", "name", "offset", "page", "page-index", "room-id", "subject", "text", "time", "title"},
},
{
CLIPath: "calendar +suggest-time",
Aliases: map[string]string{
"attendee-names": "with",
"begin": "start",
"duration-minutes": "duration",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-time": "end",
"meeting-duration-minutes": "duration",
"min-time": "start",
"names": "with",
"participant-names": "with",
"since": "start",
"start-date": "start",
"start-time": "start",
"time-max": "end",
"time-min": "start",
"to": "end",
},
Blocked: []string{"attendees", "date", "open-dingtalk-ids", "participants", "remind-minutes", "time", "user", "user-id", "user-ids", "users"},
},
{
CLIPath: "calendar +suggestion",
Aliases: map[string]string{
"begin": "start",
"duration-minutes": "duration",
"end-date": "end",
"end-time": "end",
"from": "start",
"from-date": "start",
"max-time": "end",
"meeting-duration-minutes": "duration",
"min-time": "start",
"since": "start",
"start-date": "start",
"start-time": "start",
"time-max": "end",
"time-min": "start",
"time-zone": "timezone",
"to": "end",
"tz": "timezone",
"user-ids": "users",
},
Blocked: []string{"at-user-ids", "attendee-name", "attendee-names", "date", "locale", "name", "names", "offset", "open-dingtalk-ids", "participant-name", "participant-names", "remind-minutes", "room-id", "room-ids", "room-name", "rooms", "staff-id", "time", "uid", "user", "user-id", "userid", "utc-offset", "who", "with"},
},
{
CLIPath: "calendar +update",
Aliases: map[string]string{
"add-user-ids": "add-attendees",
"add-users": "add-attendees",
"begin": "start",
"calendar": "calendar-id",
"calendar-book-id": "calendar-id",
"calendar-event-id": "event",
"description": "desc",
"end-time": "end",
"event-id": "event",
"freebusy": "free-busy",
"from": "start",
"remove-user-ids": "remove-attendees",
"remove-users": "remove-attendees",
"start-time": "start",
"subject": "title",
"summary": "title",
"time-zone": "timezone",
"to": "end",
"tz": "timezone",
},
Blocked: []string{"acl-id", "attendee-name", "attendee-names", "calendar-name", "config", "date", "end-date", "field-description", "group-id", "locale", "name", "offset", "open-dingtalk-ids", "participant-name", "participant-names", "remind-minutes", "reminder-minutes", "rich-text-desc", "room-id", "room-ids", "room-name", "rooms", "start-date", "time", "time-max", "time-min", "utc-offset", "who", "with"},
Ambiguous: []string{"attendees", "id", "user-ids", "users"},
},
{
CLIPath: "calendar acl delete",
Blocked: []string{"calendar-book-id", "calendar-id", "event", "event-id", "room-id", "user-id"},
},
{
CLIPath: "calendar attachment add",
Blocked: []string{"attachments", "file", "file-id", "file-ids"},
},
{
CLIPath: "calendar attendee add",
Blocked: []string{"attendee-name", "attendee-names", "open-dingtalk-ids", "participant-name", "participant-names", "who", "with"},
},
{
CLIPath: "calendar attendee delete",
Blocked: []string{"attendee-name", "attendee-names", "open-dingtalk-ids", "participant-name", "participant-names", "who", "with"},
},
{
CLIPath: "calendar busy search",
Aliases: map[string]string{
"room-id": "rooms",
},
Blocked: []string{"attendee-name", "attendee-names", "group-id", "location", "name", "names", "participant-names", "room", "room-name", "who", "with"},
},
{
CLIPath: "calendar event create",
Aliases: map[string]string{
"reminder-minutes": "remind-minutes",
"reminder-offset-minutes": "remind-minutes",
"room-id": "rooms",
"time-zone": "timezone",
"tz": "timezone",
},
Blocked: []string{"at", "attendee-name", "attendee-names", "due", "duration", "group-id", "locale", "offset", "participant-name", "participant-names", "remind-at", "reminder-time", "room", "room-name", "utc-offset", "who", "with"},
},
{
CLIPath: "calendar event list",
Aliases: map[string]string{
@@ -2222,54 +1789,6 @@ var generatedParamAliases = []ParamAliasEntry{
},
Blocked: []string{"offset", "page", "time"},
},
{
CLIPath: "calendar event respond",
Aliases: map[string]string{
"response": "status",
"response-status": "status",
},
Blocked: []string{"availability", "done", "free-busy", "state"},
},
{
CLIPath: "calendar event suggest",
Aliases: map[string]string{
"duration-minutes": "duration",
"meeting-duration-minutes": "duration",
"time-zone": "timezone",
"tz": "timezone",
},
Blocked: []string{"attendee-name", "attendee-names", "from", "locale", "name", "names", "offset", "open-dingtalk-ids", "participant-names", "remind-minutes", "to", "utc-offset", "who", "with"},
},
{
CLIPath: "calendar event update",
Aliases: map[string]string{
"time-zone": "timezone",
"tz": "timezone",
},
Blocked: []string{"attendees", "group-id", "locale", "offset", "participants", "remind-minutes", "reminder-minutes", "room-id", "room-ids", "room-name", "rooms", "utc-offset"},
},
{
CLIPath: "calendar room add",
Aliases: map[string]string{
"room-id": "rooms",
},
Blocked: []string{"group-id", "location", "room", "room-name"},
},
{
CLIPath: "calendar room delete",
Aliases: map[string]string{
"room-id": "rooms",
},
Blocked: []string{"group-id", "location", "room", "room-name"},
},
{
CLIPath: "calendar room search",
Aliases: map[string]string{
"group": "group-id",
"room-group-id": "group-id",
},
Blocked: []string{"location", "room", "room-id", "room-ids", "rooms"},
},
{
CLIPath: "chat +bot-find",
Aliases: map[string]string{
File diff suppressed because one or more lines are too long
+1 -1
View File
@@ -1058,7 +1058,7 @@ var schemaCompactPayloadKeys = map[string]bool{
"agent_summary": true, "description": true,
"effect": true, "risk": true, "confirmation": true, "idempotency": true,
"interface_mode": true, "availability": true, "interface_reason": true,
"parameters": true, "constraints": true, "positionals": true, "dry_run": true, "wait": true,
"parameters": true, "constraints": true, "positionals": true, "dry_run": true,
"result": true, "pagination": true,
"examples": true, "use_when": true, "avoid_when": true,
}
-1
View File
@@ -81,7 +81,6 @@ var schemaCatalogToolOptionalKeys = []string{
"pagination",
"positionals",
"result",
"wait",
}
var schemaCatalogToolEnums = map[string][]string{
-21
View File
@@ -60,7 +60,6 @@ type ToolSpec struct {
Constraints RuntimeSchemaConstraints
Positionals []contract.RuntimeSchemaPositional
DryRun *contract.DryRunSpec
Wait *contract.WaitSpec
Result *contract.ResultSpec
Pagination *contract.PaginationSpec
Safety contract.SafetySpec
@@ -136,7 +135,6 @@ type RuntimeToolSpecInput struct {
Constraints RuntimeSchemaConstraints
Positionals []contract.RuntimeSchemaPositional
DryRun *contract.DryRunSpec
Wait *contract.WaitSpec
Result *contract.ResultSpec
Pagination *contract.PaginationSpec
Safety contract.SafetySpec
@@ -544,11 +542,6 @@ func (t ToolSpec) Validate() error {
return err
}
}
if t.Wait != nil {
if err := t.Wait.Validate(id.CanonicalPath); err != nil {
return err
}
}
if t.Result != nil {
if _, err := contract.NormalizeResultSpec(t.Result, id.CanonicalPath); err != nil {
return err
@@ -751,16 +744,6 @@ func (t ToolSpec) normalized() ToolSpec {
dryRun.PreviewKind = strings.TrimSpace(dryRun.PreviewKind)
out.DryRun = &dryRun
}
if t.Wait != nil {
// NormalizeWaitSpec is the single canonical form shared with the
// declaration path: trimmed status values, duplicate/conflict
// rejection, defensive copy. Invalid declarations are rejected by
// ToolSpec.Validate below, which runs the same normalization
// through WaitSpec.Validate.
if wait, err := contract.NormalizeWaitSpec(t.Wait, id.CanonicalPath); err == nil {
out.Wait = wait
}
}
if t.Result != nil {
result, err := contract.NormalizeResultSpec(t.Result, id.CanonicalPath)
if err == nil {
@@ -999,10 +982,6 @@ func (t ToolSpec) ToPayload() (map[string]any, error) {
value, _ := typedJSONValue(t.DryRun)
payload["dry_run"] = value
}
if t.Wait != nil {
value, _ := typedJSONValue(t.Wait)
payload["wait"] = value
}
if t.Result != nil {
value, _ := typedJSONValue(t.Result)
payload["result"] = value
@@ -974,57 +974,3 @@ func TestFinalProvenanceCoverageDoesNotInventOptionalInterfaceReason(t *testing.
t.Fatalf("optional local interface_reason should not require invented provenance: %v", err)
}
}
func TestToolSpecWaitCapabilityIsPositiveOnly(t *testing.T) {
base := RuntimeToolSpecInput{Identity: contract.ToolIdentitySpec{
ProductID: "sample",
Name: "waitrun",
CLIName: "waitrun",
CLIPath: "sample waitrun",
}}
withoutCapability, err := ToolSpecFromRuntime(base)
if err != nil {
t.Fatalf("ToolSpecFromRuntime() error = %v", err)
}
payload, err := withoutCapability.ToPayload()
if err != nil {
t.Fatalf("ToPayload() error = %v", err)
}
if _, ok := payload["wait"]; ok {
t.Fatalf("nil capability unexpectedly emitted wait: %#v", payload["wait"])
}
base.Wait = &contract.WaitSpec{Mode: "webhook"}
if _, err := ToolSpecFromRuntime(base); err == nil || !strings.Contains(err.Error(), "unknown mode") {
t.Fatalf("invalid mode error = %v", err)
}
base.Wait = &contract.WaitSpec{Mode: contract.WaitModeEvent, StatusQuery: "status", Terminal: map[string]contract.ResultOutcome{"DONE": contract.ResultOutcomeSuccess}}
if _, err := ToolSpecFromRuntime(base); err == nil || !strings.Contains(err.Error(), "requires event_key") {
t.Fatalf("event mode body error = %v", err)
}
base.Wait = &contract.WaitSpec{
Mode: contract.WaitModePoll,
PollCommand: "sample status get",
StatusQuery: "result.status",
Terminal: map[string]contract.ResultOutcome{"COMPLETED": contract.ResultOutcomeSuccess},
PendingValues: []string{"NEW"},
DefaultTimeoutSecs: 120,
}
withCapability, err := ToolSpecFromRuntime(base)
if err != nil {
t.Fatalf("ToolSpecFromRuntime(valid wait) error = %v", err)
}
if withCapability.Wait == nil || withCapability.Wait.Mode != contract.WaitModePoll {
t.Fatalf("wait capability lost through normalization: %#v", withCapability.Wait)
}
payload, err = withCapability.ToPayload()
if err != nil {
t.Fatalf("ToPayload(valid wait) error = %v", err)
}
wait := payload["wait"].(map[string]any)
if wait["mode"] != contract.WaitModePoll || wait["poll_command"] != "sample status get" {
t.Fatalf("wait payload=%#v", wait)
}
}
-1
View File
@@ -354,7 +354,6 @@ func runtimeToolSpecFromContractFinal(entry runtimeSchemaEntry, final contract.C
Constraints: constraints,
Positionals: positionals,
DryRun: final.DryRun,
Wait: final.Wait,
Result: result,
Pagination: pagination,
Safety: safety,
-2
View File
@@ -65,7 +65,6 @@ type schemaToolWire struct {
Constraints RuntimeSchemaConstraints `json:"constraints"`
Positionals []contract.RuntimeSchemaPositional `json:"positionals"`
DryRun *contract.DryRunSpec `json:"dry_run"`
Wait *contract.WaitSpec `json:"wait"`
Result *contract.ResultSpec `json:"result"`
Pagination *contract.PaginationSpec `json:"pagination"`
Effect string `json:"effect"`
@@ -270,7 +269,6 @@ func schemaToolSpecFromWire(wire schemaToolWire) (ToolSpec, error) {
Constraints: wire.Constraints,
Positionals: wire.Positionals,
DryRun: wire.DryRun,
Wait: wire.Wait,
Result: wire.Result,
Pagination: wire.Pagination,
Safety: contract.SafetySpec{
-1
View File
@@ -30,7 +30,6 @@ type ContractFinalPayload struct {
Parameters []ParamDecl
Safety *SafetySpec
DryRun *DryRunSpec
Wait *WaitSpec
Result *ResultSpec
Pagination *PaginationSpec
Interface *InterfaceSpec
-149
View File
@@ -62,155 +62,6 @@ type DryRunSpec struct {
RemoteReads bool `json:"remote_reads,omitempty"`
}
// Wait modes. Poll executes the leaf's WaitPoll hook on a cadence. Event
// consumes the leaf's WaitEvents push stream and correlates events to the
// accepted resource. Auto prefers the event stream and falls back to polling
// when the stream ends before a terminal status.
const (
WaitModePoll = "poll"
WaitModeEvent = "event"
WaitModeAuto = "auto"
)
// WaitSpec is a positive capability declaration for terminal-state waiting
// (approval flows, async exports, batch jobs). A nil ToolSpec.Wait means the
// command has not declared reviewed --wait support; the flag is not
// registered and the Schema does not publish the capability.
//
// Like DryRunSpec, the object is one atomic contract field: Schema only
// projects the reviewed capability; runtime execution stays owned by the
// command runner through the leaf's WaitPoll / WaitEvents hooks. PollCommand
// names the read command that observes status — it is a declared,
// catalog-visible fact (the same command an agent would poll manually), not
// a framework-owned invocation: how one poll or event subscription executes
// is decided by the leaf.
type WaitSpec struct {
Mode string `json:"mode"`
PollCommand string `json:"poll_command,omitempty"`
StatusQuery string `json:"status_query"`
Terminal map[string]ResultOutcome `json:"terminal"`
PendingValues []string `json:"pending_values,omitempty"`
// EventKey is the push channel key the WaitEvents hook subscribes to
// (event/auto modes). Declared for the catalog; the transport stays
// leaf-owned.
EventKey string `json:"event_key,omitempty"`
// MatchField is the event-document path holding the resource identifier
// (event/auto modes); its value must equal the ResourceQuery resolution
// of the accepted result.
MatchField string `json:"match_field,omitempty"`
// ResourceQuery is the dotted path into the accepted result data that
// yields the resource identifier correlated against MatchField
// (event/auto modes).
ResourceQuery string `json:"resource_query,omitempty"`
// DefaultTimeoutSecs is the reviewed default for --wait-timeout. Zero
// means the framework default (300s); the user flag always wins.
DefaultTimeoutSecs int `json:"default_timeout_secs"`
}
// Validate checks mode requirements and the terminal/pending status maps.
// Unknown terminal outcomes, unknown modes, and mode/body mismatches fail at
// declaration so a malformed wait capability cannot reach the wire.
// Validation delegates to NormalizeWaitSpec so the acceptance rules can never
// drift from the normalization the wire and the runtime wait engine share.
func (w WaitSpec) Validate(canonical string) error {
_, err := NormalizeWaitSpec(&w, canonical)
return err
}
// NormalizeWaitSpec returns a validated, canonical, defensively copied wait
// contract. It is shared by declaration (corecmd.New / AttachContract),
// ToolSpec, and snapshot paths, mirroring NormalizeResultSpec. Status values
// are trimmed into their wire form: the wait engine compares backend
// statuses verbatim against these tables, so a padded declaration
// (" processing ") would publish a Schema that its own runtime treats as an
// unknown status. Values collapsing onto one value after trimming (duplicate
// pending values, duplicate terminal keys, terminal/pending conflicts) are
// rejected instead of silently merged.
func NormalizeWaitSpec(in *WaitSpec, canonical string) (*WaitSpec, error) {
if in == nil {
return nil, nil
}
canonical = defaultString(strings.TrimSpace(canonical), "<unknown>")
out := &WaitSpec{
Mode: strings.TrimSpace(in.Mode),
PollCommand: strings.TrimSpace(in.PollCommand),
StatusQuery: strings.TrimSpace(in.StatusQuery),
EventKey: strings.TrimSpace(in.EventKey),
MatchField: strings.TrimSpace(in.MatchField),
ResourceQuery: strings.TrimSpace(in.ResourceQuery),
DefaultTimeoutSecs: in.DefaultTimeoutSecs,
}
if out.Mode == "" {
return nil, fmt.Errorf("schema tool %s wait has no mode", canonical)
}
switch out.Mode {
case WaitModePoll, WaitModeEvent, WaitModeAuto:
default:
return nil, fmt.Errorf("schema tool %s wait has unknown mode %q", canonical, out.Mode)
}
needsPoll := out.Mode == WaitModePoll || out.Mode == WaitModeAuto
if needsPoll && out.PollCommand == "" {
return nil, fmt.Errorf("schema tool %s wait mode %s requires poll_command", canonical, out.Mode)
}
needsEvent := out.Mode == WaitModeEvent || out.Mode == WaitModeAuto
if needsEvent {
if out.EventKey == "" {
return nil, fmt.Errorf("schema tool %s wait mode %s requires event_key", canonical, out.Mode)
}
if out.MatchField == "" {
return nil, fmt.Errorf("schema tool %s wait mode %s requires match_field", canonical, out.Mode)
}
if out.ResourceQuery == "" {
return nil, fmt.Errorf("schema tool %s wait mode %s requires resource_query", canonical, out.Mode)
}
}
if out.StatusQuery == "" {
return nil, fmt.Errorf("schema tool %s wait mode %s requires status_query", canonical, out.Mode)
}
if len(in.Terminal) == 0 {
return nil, fmt.Errorf("schema tool %s wait has no terminal states", canonical)
}
out.Terminal = make(map[string]ResultOutcome, len(in.Terminal))
for status, outcome := range in.Terminal {
status = strings.TrimSpace(status)
if status == "" {
return nil, fmt.Errorf("schema tool %s wait has a blank terminal status", canonical)
}
if _, dup := out.Terminal[status]; dup {
return nil, fmt.Errorf("schema tool %s wait has duplicate terminal status %q", canonical, status)
}
// Terminal states must close into success or failure. Pending and
// partial are not wait outcomes: pending is expressed through
// timeout, and partial requires the typed multi-status payload only
// the leaf can construct.
if outcome != ResultOutcomeSuccess && outcome != ResultOutcomeFailure {
return nil, fmt.Errorf(
"schema tool %s wait terminal status %q must map to success or failure, got %q",
canonical, status, outcome)
}
out.Terminal[status] = outcome
}
seenPending := make(map[string]bool, len(in.PendingValues))
for _, value := range in.PendingValues {
value = strings.TrimSpace(value)
if value == "" {
return nil, fmt.Errorf("schema tool %s wait has a blank pending value", canonical)
}
if _, conflict := out.Terminal[value]; conflict {
return nil, fmt.Errorf("schema tool %s wait status %q is both terminal and pending", canonical, value)
}
if seenPending[value] {
return nil, fmt.Errorf("schema tool %s wait has duplicate pending value %q", canonical, value)
}
seenPending[value] = true
out.PendingValues = append(out.PendingValues, value)
}
if in.DefaultTimeoutSecs < 0 {
return nil, fmt.Errorf("schema tool %s wait default_timeout_secs must be >= 0", canonical)
}
return out, nil
}
// ResultOutcome is one closed unified-output envelope outcome.
type ResultOutcome string
-227
View File
@@ -1,227 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package contract
import "testing"
func TestWaitSpecValidateAcceptsReviewedShapes(t *testing.T) {
cases := []WaitSpec{
{
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "result.status",
Terminal: map[string]ResultOutcome{
"COMPLETED": ResultOutcomeSuccess,
"REJECTED": ResultOutcomeFailure,
},
PendingValues: []string{"NEW", "RUNNING"},
DefaultTimeoutSecs: 600,
},
}
for i, spec := range cases {
if err := spec.Validate("sample.tool"); err != nil {
t.Fatalf("case %d: unexpected error: %v", i, err)
}
}
}
func TestWaitSpecValidateRejectsMalformedShapes(t *testing.T) {
terminal := map[string]ResultOutcome{"COMPLETED": ResultOutcomeSuccess}
cases := map[string]WaitSpec{
"no mode": {
Terminal: terminal,
},
"unknown mode": {
Mode: "webhook",
Terminal: terminal,
},
"poll without poll_command": {
Mode: WaitModePoll,
StatusQuery: "status",
Terminal: terminal,
},
"poll without status_query": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
Terminal: terminal,
},
"event without event_key": {
Mode: WaitModeEvent,
MatchField: "process_instance_id",
ResourceQuery: "id",
StatusQuery: "result.status",
Terminal: terminal,
},
"event without match_field": {
Mode: WaitModeEvent,
EventKey: "bpms_instance_change",
ResourceQuery: "id",
StatusQuery: "result.status",
Terminal: terminal,
},
"event without resource_query": {
Mode: WaitModeEvent,
EventKey: "bpms_instance_change",
MatchField: "process_instance_id",
StatusQuery: "result.status",
Terminal: terminal,
},
"auto missing poll_command": {
Mode: WaitModeAuto,
EventKey: "export_finished",
MatchField: "job_id",
ResourceQuery: "job_id",
StatusQuery: "status",
Terminal: terminal,
},
"no terminal states": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
},
"blank terminal status": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{" ": ResultOutcomeSuccess},
},
"terminal outcome pending": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{"COMPLETED": ResultOutcomePending, "REJECTED": ResultOutcomeFailure},
},
"terminal outcome partial": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{"COMPLETED": ResultOutcomePartialFailure, "REJECTED": ResultOutcomeFailure},
},
"terminal outcome outside closed set": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{"COMPLETED": ResultOutcome("explosion")},
},
"only pending terminal outcome": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{"NEW": ResultOutcomePending},
},
"status both terminal and pending": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: terminal,
PendingValues: []string{"COMPLETED"},
},
"blank pending value": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: terminal,
PendingValues: []string{" "},
},
"negative timeout default": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: terminal,
DefaultTimeoutSecs: -1,
},
}
for name, spec := range cases {
if err := spec.Validate("sample.tool"); err == nil {
t.Fatalf("%s: expected error, got nil", name)
}
}
}
func TestNormalizeWaitSpecTrimsStatusValuesIntoWireForm(t *testing.T) {
in := &WaitSpec{
Mode: " poll ",
PollCommand: " oa approval-instance get ",
StatusQuery: " result.status ",
Terminal: map[string]ResultOutcome{" COMPLETED ": ResultOutcomeSuccess, "REJECTED": ResultOutcomeFailure},
PendingValues: []string{" NEW ", "RUNNING"},
DefaultTimeoutSecs: 60,
}
out, err := NormalizeWaitSpec(in, "sample.tool")
if err != nil {
t.Fatalf("NormalizeWaitSpec() error = %v", err)
}
if out.Mode != WaitModePoll || out.PollCommand != "oa approval-instance get" || out.StatusQuery != "result.status" {
t.Fatalf("normalized scalars: %#v", out)
}
if len(out.Terminal) != 2 {
t.Fatalf("terminal=%#v, want two trimmed keys", out.Terminal)
}
if got := out.Terminal["COMPLETED"]; got != ResultOutcomeSuccess {
t.Fatalf("terminal[COMPLETED]=%q, want success (key must be trimmed)", got)
}
if _, padded := out.Terminal[" COMPLETED "]; padded {
t.Fatal("padded terminal key survived normalization")
}
for i, want := range []string{"NEW", "RUNNING"} {
if out.PendingValues[i] != want {
t.Fatalf("pending[%d]=%q, want %q", i, out.PendingValues[i], want)
}
}
// The input declaration must stay untouched (defensive copy).
if _, padded := in.Terminal[" COMPLETED "]; !padded {
t.Fatal("NormalizeWaitSpec mutated its input terminal map")
}
if in.PendingValues[0] != " NEW " {
t.Fatal("NormalizeWaitSpec mutated its input pending values")
}
}
func TestNormalizeWaitSpecRejectsDuplicatesAndConflictsAfterTrim(t *testing.T) {
terminal := map[string]ResultOutcome{"COMPLETED": ResultOutcomeSuccess}
cases := map[string]*WaitSpec{
"terminal keys collapsing after trim": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: map[string]ResultOutcome{"COMPLETED": ResultOutcomeSuccess, " COMPLETED ": ResultOutcomeFailure},
},
"pending values collapsing after trim": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: terminal,
PendingValues: []string{"NEW", " NEW "},
},
"terminal/pending conflict hidden by padding": {
Mode: WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "status",
Terminal: terminal,
PendingValues: []string{" COMPLETED "},
},
}
for name, spec := range cases {
if _, err := NormalizeWaitSpec(spec, "sample.tool"); err == nil {
t.Fatalf("%s: expected error, got nil", name)
}
}
}
func TestNormalizeWaitSpecNilReturnsNil(t *testing.T) {
out, err := NormalizeWaitSpec(nil, "sample.tool")
if err != nil || out != nil {
t.Fatalf("NormalizeWaitSpec(nil) = %#v, %v", out, err)
}
}
-4
View File
@@ -42,7 +42,6 @@ type ContractDecl struct {
Positionals []contract.RuntimeSchemaPositional
Parameters []contract.ParamDecl
DryRun *contract.DryRunSpec
Wait *contract.WaitSpec
Result *contract.ResultSpec
Pagination *contract.PaginationSpec
Interface *contract.InterfaceSpec
@@ -147,9 +146,6 @@ func (s ContractDecl) empty() bool {
if s.DryRun != nil && strings.TrimSpace(s.DryRun.PreviewKind) != "" {
return false
}
if s.Wait != nil && strings.TrimSpace(s.Wait.Mode) != "" {
return false
}
if s.Result != nil {
return false
}
@@ -200,13 +200,6 @@ func TestFrameworkContractFinalDeepCopyAndSafetyConflicts(t *testing.T) {
Parameters: []contract.ParamDecl{{Name: "mode", Enum: []string{"a"}, Required: boolPointer(true)}},
Safety: &contract.SafetySpec{Effect: " read ", EffectSource: " source ", Risk: " low ", Confirmation: " not_required ", Idempotency: " idempotent "},
DryRun: &contract.DryRunSpec{PreviewKind: "plan"},
Wait: &contract.WaitSpec{
Mode: contract.WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "result.status",
Terminal: map[string]contract.ResultOutcome{"COMPLETED": contract.ResultOutcomeSuccess, "REJECTED": contract.ResultOutcomeFailure},
PendingValues: []string{"NEW"},
},
Result: &contract.ResultSpec{
Outcomes: []contract.ResultOutcome{contract.ResultOutcomeSuccess},
DataSchema: []byte(`{"type":"object"}`), SensitivePaths: []string{"token"},
@@ -225,8 +218,6 @@ func TestFrameworkContractFinalDeepCopyAndSafetyConflicts(t *testing.T) {
if !ok || got.Result == payload.Result || got.Pagination == payload.Pagination || got.Interface == payload.Interface || got.Selection == payload.Selection || got.Identity == payload.Identity {
t.Fatalf("payload not deeply cloned: %#v", got)
}
payload.Wait.Terminal["COMPLETED"] = contract.ResultOutcomeFailure
payload.Wait.PendingValues[0] = "mutated"
payload.Parameters[0].Enum[0] = "changed"
*payload.Parameters[0].Required = false
*payload.Selection.ExampleDispositions[0].Index = 9
@@ -235,10 +226,6 @@ func TestFrameworkContractFinalDeepCopyAndSafetyConflicts(t *testing.T) {
if again.Parameters[0].Enum[0] != "a" || !*again.Parameters[0].Required || *again.Selection.ExampleDispositions[0].Index != 1 || !*again.Selection.Reviewed {
t.Fatalf("stored payload aliased input: %#v", again)
}
if again.Wait == payload.Wait || again.Wait.Terminal["COMPLETED"] != contract.ResultOutcomeSuccess || again.Wait.PendingValues[0] != "NEW" {
t.Fatalf("wait spec aliased input: %#v", again.Wait)
t.Fatalf("stored payload aliased input: %#v", again)
}
matching := &cobra.Command{Use: "matching"}
t.Cleanup(func() { ClearRuntimeContractFinalForTest(matching) })
-9
View File
@@ -74,15 +74,6 @@ func cloneContractFinalPayload(in contract.ContractFinalPayload) contract.Contra
value := *in.DryRun
out.DryRun = &value
}
if in.Wait != nil {
value := *in.Wait
value.Terminal = make(map[string]contract.ResultOutcome, len(in.Wait.Terminal))
for status, outcome := range in.Wait.Terminal {
value.Terminal[status] = outcome
}
value.PendingValues = cloneSlice(in.Wait.PendingValues)
out.Wait = &value
}
if in.Result != nil {
value := *in.Result
value.Outcomes = cloneSlice(in.Result.Outcomes)
-330
View File
@@ -49,16 +49,12 @@ package corecmd
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"os"
"strconv"
"strings"
"time"
"github.com/mattn/go-isatty"
"github.com/spf13/cobra"
@@ -68,7 +64,6 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/runtimeannotate"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/wait"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
)
@@ -287,21 +282,6 @@ type Spec struct {
// Orchestrate executes a multi-step command; it assembles whatever payloads
// it needs from the Ctx.
Orchestrate func(c *Ctx) error
// WaitPoll executes one poll of the declared Contract.Wait capability.
// Exactly one poll is one call; cadence, status extraction, and outcome
// mapping belong to the framework wait phase. Required for poll/auto
// declarations — a declared capability without a runtime implementation
// can never honor --wait, so New rejects the pairing at construction.
// ctx is the wait-phase deadline (--wait-timeout); leaf I/O must honor
// it so a blocked poll cannot outlive the declared timeout.
WaitPoll func(ctx context.Context, c *Ctx) (wait.PollDoc, error)
// WaitEvents opens the push subscription of the declared Contract.Wait
// capability (event/auto modes). The framework owns correlation and
// status mapping; the leaf owns the transport. Auto mode falls back to
// WaitPoll when the stream ends before a terminal status. ctx is the
// same wait-phase deadline as WaitPoll; subscription setup must honor
// it so --wait-timeout can cancel a blocked subscribe.
WaitEvents func(ctx context.Context, c *Ctx) (wait.EventStream, error)
}
// Ctx is the framework-neutral execution context handed to Invoke/Orchestrate.
@@ -375,15 +355,6 @@ func (c *Ctx) Changed(name string) bool { return c.cmd.Flags().Changed(name) }
// DryRun reports the effective global --dry-run.
func (c *Ctx) DryRun() bool { return BoolFlag(c.cmd, "dry-run") }
// Wait reports the effective --wait flag. It is false on commands that did
// not declare the capability: the flag is not registered there, so passing it
// is an unknown-flag error rather than a silently ignored value.
func (c *Ctx) Wait() bool { return BoolFlag(c.cmd, waitFlagName) }
// WaitTimeoutSecs reports the effective --wait-timeout in seconds (flag
// value, then the declared default, then the framework default).
func (c *Ctx) WaitTimeoutSecs() int { return waitTimeoutSecs(c.cmd) }
// Yes reports the effective global --yes.
func (c *Ctx) Yes() bool { return BoolFlag(c.cmd, "yes") }
@@ -403,8 +374,6 @@ func New(spec Spec) *cobra.Command {
validateDispatchDecl(spec)
validateSafetySpec(spec)
validateContractDecl(spec)
normalizeWaitDecl(&spec)
validateWaitDecl(spec)
validateInputSpecs(spec.Use, spec.Flags)
// Help prose inherits the declaration when not authored separately:
// Selection.Examples (already contract-validated against the real flags)
@@ -421,7 +390,6 @@ func New(spec Spec) *cobra.Command {
Hidden: spec.Hidden,
}
RegisterFlags(cmd, spec.Flags)
registerWaitFlags(cmd, spec)
ValidateConstraintDecls(spec.Use, spec.Flags, spec.Constraints)
embedContractIntoSchema(cmd, spec)
AnnotateConstraints(cmd, spec.Constraints)
@@ -487,10 +455,6 @@ func New(spec Spec) *cobra.Command {
if err != nil {
return err
}
result, err = runDeclaredWaitPhase(cmd, args, spec, result)
if err != nil {
return err
}
return output.StoreResult(cmd.Context(), result)
}
return spec.Invoke(ctx, toolArgs)
@@ -542,293 +506,6 @@ func runDeclaredPreflight(cmd *cobra.Command, args []string, spec Spec) error {
return nil
}
// Wait-phase framework flags. They are registered natively on the leaf (never
// as FlagSpec) so they cannot leak into toolArgs / MCP payloads: --wait is a
// client-side execution modifier, not a backend parameter. On commands that
// did not declare Contract.Wait the flags do not exist, so passing --wait
// fails as an unknown flag instead of being silently ignored.
const (
waitFlagName = "wait"
waitTimeoutFlagName = "wait-timeout"
defaultWaitTimeoutS = 300
)
// DefaultWaitTimeoutSecs is the framework default for --wait-timeout when the
// declaration carries no reviewed default.
const DefaultWaitTimeoutSecs = defaultWaitTimeoutS
// normalizeWaitDecl rewrites the spec's declared Wait in place with its
// canonical form (contract.NormalizeWaitSpec): trimmed status values,
// duplicate/conflict rejection, defensive copy. The runtime wait phase and
// AttachContract both read spec.Contract.Wait, so normalizing once at
// construction guarantees the wait engine, the Schema wire, and the
// registered ContractFinal payload all see identical status tables — a
// padded declaration can no longer publish a Schema its own runtime treats
// as unknown statuses. An invalid declaration panics here, next to the
// authoring mistake, with the same message Validate reports.
func normalizeWaitDecl(spec *Spec) {
decl := spec.Contract.Wait
if decl == nil || strings.TrimSpace(decl.Mode) == "" {
return
}
normalized, err := contract.NormalizeWaitSpec(decl, spec.Contract.Identity.CanonicalPath)
if err != nil {
panic(fmt.Sprintf("command %q has invalid Contract.Wait: %v", spec.Use, err))
}
spec.Contract.Wait = normalized
}
// validateWaitDecl enforces the declaration ⇄ implementation pairing at build
// time: a declared Contract.Wait without a WaitPoll hook is a capability the
// command can never honor, and a WaitPoll hook without the declaration has no
// flags or Schema capability to serve. The declaration also requires the
// ResultInvoke dispatcher: only the unified-result envelope can be closed
// into the terminal outcome (error.type "wait", exit code 8) and the timed-out
// pending form — legacy Invoke/Orchestrate/RunE paths emit their own output
// and would observe a failure terminal while still exiting 0. All three
// mismatches are programming errors.
func validateWaitDecl(spec Spec) {
decl := spec.Contract.Wait
declared := decl != nil && strings.TrimSpace(decl.Mode) != ""
if !declared {
if spec.WaitPoll != nil || spec.WaitEvents != nil {
panic(fmt.Sprintf(
"command %q sets a wait hook without declaring Contract.Wait: the wait flags and Schema capability come from the declaration",
spec.Use))
}
return
}
if spec.ResultInvoke == nil {
panic(fmt.Sprintf(
"command %q declares Contract.Wait without ResultInvoke: wait closes the unified-result envelope, which legacy Invoke/Orchestrate/RunE paths cannot rewrite",
spec.Use))
}
mode := strings.TrimSpace(decl.Mode)
needsPoll := mode == contract.WaitModePoll || mode == contract.WaitModeAuto
needsEvent := mode == contract.WaitModeEvent || mode == contract.WaitModeAuto
if needsPoll && spec.WaitPoll == nil {
panic(fmt.Sprintf(
"command %q declares wait mode %s but sets no WaitPoll: a declared wait capability must carry its runtime poll implementation",
spec.Use, mode))
}
if needsEvent && spec.WaitEvents == nil {
panic(fmt.Sprintf(
"command %q declares wait mode %s but sets no WaitEvents: a declared event wait must carry its runtime subscription",
spec.Use, mode))
}
if !needsPoll && spec.WaitPoll != nil {
panic(fmt.Sprintf(
"command %q declares wait mode %s but sets WaitPoll: the declaration decides which hooks run",
spec.Use, mode))
}
if !needsEvent && spec.WaitEvents != nil {
panic(fmt.Sprintf(
"command %q declares wait mode %s but sets WaitEvents: the declaration decides which hooks run",
spec.Use, mode))
}
}
// registerWaitFlags adds --wait / --wait-timeout to a leaf that declared
// Contract.Wait. The timeout default is the reviewed declaration, falling
// back to DefaultWaitTimeoutSecs.
func registerWaitFlags(cmd *cobra.Command, spec Spec) {
decl := spec.Contract.Wait
if decl == nil || strings.TrimSpace(decl.Mode) == "" {
return
}
cmd.Flags().Bool(waitFlagName, false,
"等待到达命令声明的终态(如审批完成、导出结束)后再返回;未声明该能力的命令不接受此 flag")
timeoutDefault := decl.DefaultTimeoutSecs
if timeoutDefault <= 0 {
timeoutDefault = DefaultWaitTimeoutSecs
}
cmd.Flags().Int(waitTimeoutFlagName, timeoutDefault,
"等待超时秒数;超时以 pending 结束(异步受理不是失败)")
}
// runDeclaredWaitPhase runs the declared wait loop after a successful
// ResultInvoke dispatch and closes the accepted unified envelope into the
// wait outcome (validateWaitDecl guarantees the ResultInvoke pairing).
// Only a pending accepted result is waitable: success / failure / partial
// are already terminal and must be returned unchanged. Waiting on a
// business failure would let WithOutcome(..., success) overwrite it into
// an illegal success-with-error envelope.
func runDeclaredWaitPhase(cmd *cobra.Command, args []string, spec Spec, result output.CommandResult) (output.CommandResult, error) {
if !BoolFlag(cmd, waitFlagName) {
return result, nil
}
if result == nil || result.Outcome() != output.OutcomePending {
return result, nil
}
decl := spec.Contract.Wait
timeout, err := waitTimeoutDuration(int64(waitTimeoutSecs(cmd)))
if err != nil {
return result, err
}
ctx := newCtx(cmd, args, spec.Flags)
outcome, err := runWaitLoop(cmd.Context(), decl, timeout, spec, ctx, result)
if err != nil {
return result, err
}
if outcome.TimedOut {
cmd.PrintErrf("等待超时(%s):当前状态 %q,未到达终态,以 pending 结束\n", timeout, outcome.Status)
return output.WithOutcome(result, output.OutcomePending,
output.WithOperationTimedOut(outcome.Status)), nil
}
if outcome.Outcome == contract.ResultOutcomeFailure {
return output.WithOutcome(result, output.OutcomeFailure,
output.WithOperationTerminalState(outcome.Status),
output.WithErrorInfo(&output.ErrorInfo{
Type: "wait",
Subtype: "terminal_failure",
Message: fmt.Sprintf("等待到达失败终态:%s", outcome.Status),
})), nil
}
return output.WithOutcome(result, output.OutcomeSuccess,
output.WithOperationTerminalState(outcome.Status)), nil
}
// runWaitLoop executes the declared wait mode. One deadline spans the event
// phase and an auto-mode poll fallback (the inner loops run without their
// own timeouts and inherit this context's deadline). The deadline is
// forwarded to WaitPoll / WaitEvents and bound onto the cobra command so
// leaf I/O that reads either the hook ctx or Command().Context() is
// cancelled when --wait-timeout expires.
func runWaitLoop(parent context.Context, decl *contract.WaitSpec, timeout time.Duration, spec Spec, ctx *Ctx, result output.CommandResult) (wait.Outcome, error) {
loopCtx := parent
if timeout > 0 {
var cancel context.CancelFunc
loopCtx, cancel = context.WithTimeout(parent, timeout)
defer cancel()
}
if ctx != nil && ctx.cmd != nil {
prev := ctx.cmd.Context()
ctx.cmd.SetContext(loopCtx)
defer ctx.cmd.SetContext(prev)
}
mode := strings.TrimSpace(decl.Mode)
if mode == contract.WaitModePoll {
return wait.Run(loopCtx, wait.LoopSpec{
StatusQuery: decl.StatusQuery,
Terminal: decl.Terminal,
Pending: decl.PendingValues,
}, func(pollCtx context.Context) (wait.PollDoc, error) {
return spec.WaitPoll(pollCtx, ctx)
})
}
resource, err := waitResource(decl, result)
if err != nil {
return wait.Outcome{}, err
}
stream, err := spec.WaitEvents(loopCtx, ctx)
if err != nil {
if loopCtx.Err() != nil {
// Subscribe blocked until the wait deadline: same contract as a
// cancelled poll — close as timed-out pending, do not surface
// ctx.Err() as a subscription failure (and do not poll-fallback
// in auto mode; the shared deadline is already exhausted).
return wait.Outcome{Outcome: contract.ResultOutcomePending, TimedOut: true}, nil
}
if mode == contract.WaitModeAuto {
// Subscription failed before any event: fall back to polling.
return pollWithSpec(loopCtx, decl, spec, ctx)
}
return wait.Outcome{}, fmt.Errorf("wait: event subscription failed: %w", err)
}
eventSpec := wait.EventLoopSpec{
StatusQuery: decl.StatusQuery,
MatchField: decl.MatchField,
Terminal: decl.Terminal,
Pending: decl.PendingValues,
}
outcome, err := wait.RunEvent(loopCtx, eventSpec, resource, stream)
if err == nil {
return outcome, nil
}
if mode == contract.WaitModeAuto && errors.Is(err, wait.ErrEventStreamEnded) {
// Stream ended before a terminal status: fall back to polling under
// the same deadline.
return pollWithSpec(loopCtx, decl, spec, ctx)
}
return outcome, err
}
// pollWithSpec runs the poll loop for an auto-mode fallback.
func pollWithSpec(loopCtx context.Context, decl *contract.WaitSpec, spec Spec, ctx *Ctx) (wait.Outcome, error) {
return wait.Run(loopCtx, wait.LoopSpec{
StatusQuery: decl.StatusQuery,
Terminal: decl.Terminal,
Pending: decl.PendingValues,
}, func(pollCtx context.Context) (wait.PollDoc, error) {
return spec.WaitPoll(pollCtx, ctx)
})
}
// waitResource resolves the resource identifier an event stream correlates
// against, from the accepted result data via the declared ResourceQuery.
// result.Data() returns any deep-copied business payload, which may be a
// map[string]any, struct, or struct pointer. We normalize via JSON round-trip
// to support all valid result types uniformly.
func waitResource(decl *contract.WaitSpec, result output.CommandResult) (string, error) {
raw := result.Data()
if raw == nil {
return "", fmt.Errorf("wait: accepted result data is nil; cannot resolve resource %q", decl.ResourceQuery)
}
// Fast path: already a map.
if data, ok := raw.(map[string]any); ok {
resource, ok := wait.ExtractStatus(wait.PollDoc(data), decl.ResourceQuery)
if !ok || strings.TrimSpace(resource) == "" {
return "", fmt.Errorf("wait: resource query %q not found in accepted result data", decl.ResourceQuery)
}
return resource, nil
}
// Slow path: struct or struct pointer. Normalize via JSON round-trip.
jsonBytes, err := json.Marshal(raw)
if err != nil {
return "", fmt.Errorf("wait: accepted result data cannot be serialized to JSON: %w", err)
}
var data map[string]any
if err := json.Unmarshal(jsonBytes, &data); err != nil {
return "", fmt.Errorf("wait: accepted result data is not an object; cannot resolve resource %q", decl.ResourceQuery)
}
resource, ok := wait.ExtractStatus(wait.PollDoc(data), decl.ResourceQuery)
if !ok || strings.TrimSpace(resource) == "" {
return "", fmt.Errorf("wait: resource query %q not found in accepted result data", decl.ResourceQuery)
}
return resource, nil
}
// waitTimeoutSecs resolves the effective timeout. The flag is registered
// with the reviewed declaration default (or the framework default), so the
// flag value is authoritative; a non-positive explicit value falls back to
// the framework default.
func waitTimeoutSecs(cmd *cobra.Command) int {
if value, err := cmd.Flags().GetInt(waitTimeoutFlagName); err == nil && value > 0 {
return value
}
return DefaultWaitTimeoutSecs
}
// maxWaitTimeoutSecs is the largest second count that still fits in a
// time.Duration. Multiplying a larger int by time.Second overflows to a
// non-positive duration, which would skip the deadline and wait forever.
const maxWaitTimeoutSecs = math.MaxInt64 / int64(time.Second)
// waitTimeoutDuration converts a resolved second count into the wait-phase
// deadline. Values that cannot be represented as a positive time.Duration
// are rejected as validation errors instead of silently disabling timeout.
func waitTimeoutDuration(secs int64) (time.Duration, error) {
if secs <= 0 {
secs = DefaultWaitTimeoutSecs
}
if secs > maxWaitTimeoutSecs {
return 0, apperrors.NewValidation(fmt.Sprintf(
"参数 --%s 取值 %d 超出可表示范围(最大 %d 秒)",
waitTimeoutFlagName, secs, maxWaitTimeoutSecs))
}
return time.Duration(secs) * time.Second, nil
}
// validateDispatchDecl enforces "exactly one dispatcher" at build time. Like
// ValidateConstraintDecls this panics: a spec with no runnable body (or with two
// competing ones) is a programming error that every test and startup path should
@@ -1731,13 +1408,6 @@ func AttachContract(cmd *cobra.Command, safety contract.SafetySpec, decl Contrac
d.PreviewKind = strings.TrimSpace(d.PreviewKind)
payload.DryRun = &d
}
if decl.Wait != nil && strings.TrimSpace(decl.Wait.Mode) != "" {
waitSpec, err := contract.NormalizeWaitSpec(decl.Wait, decl.Identity.CanonicalPath)
if err != nil {
panic(fmt.Sprintf("command %q has invalid Contract.Wait: %v", cmd.Name(), err))
}
payload.Wait = waitSpec
}
if decl.Result != nil {
result, err := contract.NormalizeResultSpec(decl.Result, decl.Identity.CanonicalPath)
if err != nil {
-109
View File
@@ -26,7 +26,6 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/contractfinal"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/runtimeannotate"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/testseam"
"github.com/spf13/cobra"
)
@@ -2121,111 +2120,3 @@ func TestCrossPlatformCoverageEmbedContractSkipsBlankAndHiddenFlags(t *testing.T
}
}
}
// TestWaitResourceStronglyTypedDTOs verifies waitResource handles struct and
// struct pointer results via JSON normalization (P1 fix for auto-CR).
func TestWaitResourceStronglyTypedDTOs(t *testing.T) {
type TaskDTO struct {
TaskID string `json:"task_id"`
Status string `json:"status"`
}
type NestedDTO struct {
Meta struct {
ResourceID string `json:"resource_id"`
} `json:"meta"`
}
tests := []struct {
name string
data any
query string
wantResource string
wantErrSubstr string
}{
{
name: "map[string]any fast path",
data: map[string]any{"task_id": "abc123"},
query: "task_id",
wantResource: "abc123",
},
{
name: "struct value",
data: TaskDTO{TaskID: "struct-456", Status: "running"},
query: "task_id",
wantResource: "struct-456",
},
{
name: "struct pointer",
data: &TaskDTO{TaskID: "ptr-789", Status: "pending"},
query: "task_id",
wantResource: "ptr-789",
},
{
name: "nested struct dotted query",
data: NestedDTO{},
query: "meta.resource_id",
wantResource: "",
wantErrSubstr: "not found",
},
{
name: "nested struct with value",
data: func() NestedDTO {
var d NestedDTO
d.Meta.ResourceID = "nested-xyz"
return d
}(),
query: "meta.resource_id",
wantResource: "nested-xyz",
},
{
name: "nil data",
data: nil,
query: "task_id",
wantErrSubstr: "nil",
},
{
name: "non-object data (string)",
data: "not-an-object",
query: "task_id",
wantErrSubstr: "not an object",
},
{
name: "non-object data (slice)",
data: []string{"a", "b"},
query: "task_id",
wantErrSubstr: "not an object",
},
{
name: "missing query field in struct",
data: TaskDTO{TaskID: "abc"},
query: "nonexistent",
wantErrSubstr: "not found",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
decl := &contract.WaitSpec{
ResourceQuery: tt.query,
}
result := output.Success(tt.data)
resource, err := waitResource(decl, result)
if tt.wantErrSubstr != "" {
if err == nil {
t.Fatalf("expected error containing %q, got nil", tt.wantErrSubstr)
}
if !strings.Contains(err.Error(), tt.wantErrSubstr) {
t.Errorf("error = %q; want substring %q", err.Error(), tt.wantErrSubstr)
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if resource != tt.wantResource {
t.Errorf("resource = %q; want %q", resource, tt.wantResource)
}
})
}
}
-902
View File
@@ -1,902 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package corecmd
import (
"bytes"
"context"
"errors"
"io"
"math"
"strconv"
"strings"
"testing"
"time"
"github.com/spf13/cobra"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/contract"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/contractfinal"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/wait"
)
func waitTestDecl() ContractDecl {
return ContractDecl{
Title: "Wait Title",
Description: "Wait Desc",
Wait: &contract.WaitSpec{
Mode: contract.WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "result.status",
Terminal: map[string]contract.ResultOutcome{"COMPLETED": contract.ResultOutcomeSuccess, "REJECTED": contract.ResultOutcomeFailure},
PendingValues: []string{"NEW", "RUNNING"},
DefaultTimeoutSecs: 60,
},
Interface: &contract.InterfaceSpec{Mode: "local", Availability: "available"},
Selection: contract.SelectionSpec{
AgentSummary: "summary",
UseWhen: []string{"when wait"},
AvoidWhen: []string{"when nowait"},
Examples: []string{"dws wait-sample --wait"},
},
Identity: contract.ToolIdentitySpec{ProductID: "sample", Name: "waitsample", CanonicalPath: "sample.waitsample", CLIPath: "wait-sample", PrimaryCLIPath: "wait-sample"},
}
}
func baseWaitSpec(decl ContractDecl, poll func(context.Context, *Ctx) (wait.PollDoc, error)) Spec {
return Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: decl,
WaitPoll: poll,
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending(map[string]any{"id": "job-1"}, &output.OperationInfo{
ID: "job-1",
State: "NEW",
NextCommand: "dws wait-sample --id job-1",
}), nil
},
}
}
func TestWaitFlagsOnlyRegisteredWhenDeclared(t *testing.T) {
declared := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
if flag := declared.Flags().Lookup(waitFlagName); flag == nil {
t.Fatal("declared command missing --wait flag")
}
if flag := declared.Flags().Lookup(waitTimeoutFlagName); flag == nil {
t.Fatal("declared command missing --wait-timeout flag")
}
undeclared := New(Spec{
Use: "nowait",
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Invoke: func(*Ctx, map[string]any) error { return nil },
})
if flag := undeclared.Flags().Lookup(waitFlagName); flag != nil {
t.Fatal("undeclared command registered --wait")
}
undeclared.SetArgs([]string{"--wait"})
if err := undeclared.Execute(); err == nil || !strings.Contains(err.Error(), "unknown flag") {
t.Fatalf("err=%v want unknown-flag", err)
}
}
func TestValidateWaitDeclPairsDeclarationWithImplementation(t *testing.T) {
decl := waitTestDecl()
spec := baseWaitSpec(decl, nil)
expectPanic(t, func() { New(spec) }, "WaitPoll")
spec.WaitPoll = func(context.Context, *Ctx) (wait.PollDoc, error) { return nil, nil }
expectPanic(t, func() {
New(Spec{
Use: "hook-only",
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Invoke: func(*Ctx, map[string]any) error { return nil },
WaitPoll: func(context.Context, *Ctx) (wait.PollDoc, error) { return nil, nil },
})
}, "Contract.Wait")
}
func expectPanic(t *testing.T, fn func(), want string) {
t.Helper()
defer func() {
recovered := recover()
if recovered == nil {
t.Fatalf("expected panic containing %q", want)
}
if message, ok := recovered.(string); !ok || !strings.Contains(message, want) {
t.Fatalf("panic=%v want containing %q", recovered, want)
}
}()
fn()
}
func TestWaitTimeoutFlagDefaultsComeFromDeclaration(t *testing.T) {
stub := func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}
declared := New(baseWaitSpec(waitTestDecl(), stub))
if value, err := declared.Flags().GetInt(waitTimeoutFlagName); err != nil || value != 60 {
t.Fatalf("declared default=%d/%v, want reviewed 60", value, err)
}
decl := waitTestDecl()
decl.Wait.DefaultTimeoutSecs = 0
fallback := New(baseWaitSpec(decl, stub))
if value, err := fallback.Flags().GetInt(waitTimeoutFlagName); err != nil || value != DefaultWaitTimeoutSecs {
t.Fatalf("fallback default=%d/%v, want framework %d", value, err, DefaultWaitTimeoutSecs)
}
// A non-positive explicit value falls back to the framework default.
fallback.SetArgs([]string{"--wait-timeout", "0", "--wait"})
if err := fallback.Flags().Set(waitTimeoutFlagName, "0"); err != nil {
t.Fatal(err)
}
if got := waitTimeoutSecs(fallback); got != DefaultWaitTimeoutSecs {
t.Fatalf("waitTimeoutSecs=%d, want %d", got, DefaultWaitTimeoutSecs)
}
}
func TestWaitTimeoutDurationRejectsOverflowingSeconds(t *testing.T) {
// math.MaxInt64 (9223372036854775807) is a legal pflag int on 64-bit
// platforms and overflows time.Duration(secs)*time.Second to a negative
// value, which would disable the wait deadline.
if _, err := waitTimeoutDuration(math.MaxInt64); err == nil || !strings.Contains(err.Error(), "超出可表示范围") {
t.Fatalf("err=%v, want overflow validation", err)
}
d, err := waitTimeoutDuration(maxWaitTimeoutSecs)
if err != nil {
t.Fatal(err)
}
if d <= 0 || d != time.Duration(maxWaitTimeoutSecs)*time.Second {
t.Fatalf("duration=%d, want the largest representable timeout", d)
}
// Non-positive second counts fall back to the framework default instead
// of disabling the deadline (waitTimeoutSecs already maps a zero/negative
// flag to the default; this keeps the conversion itself fail-safe).
if got, err := waitTimeoutDuration(0); err != nil || got != time.Duration(DefaultWaitTimeoutSecs)*time.Second {
t.Fatalf("duration/err=%d/%v, want framework default", got, err)
}
if got, err := waitTimeoutDuration(-5); err != nil || got != time.Duration(DefaultWaitTimeoutSecs)*time.Second {
t.Fatalf("duration/err=%d/%v, want framework default", got, err)
}
}
func TestResultInvokeWaitTimeoutOverflowIsValidationError(t *testing.T) {
if int64(math.MaxInt) <= maxWaitTimeoutSecs {
t.Skip("platform int cannot overflow time.Duration")
}
polled := false
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
polled = true
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
cmd.SetArgs([]string{"--wait", "--wait-timeout", strconv.Itoa(math.MaxInt)})
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "超出可表示范围") {
t.Fatalf("err=%v, want overflow validation", err)
}
if polled {
t.Fatal("overflowing --wait-timeout must not start the wait loop")
}
}
func TestResultInvokeWaitPollErrorFailsTheCommand(t *testing.T) {
boom := errors.New("rpc down")
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
return nil, boom
}))
cmd.SetArgs([]string{"--wait"})
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "rpc down") {
t.Fatalf("err=%v, want poll error surfaced", err)
}
}
func TestResultInvokeWaitUnknownStatusFailsClosed(t *testing.T) {
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "Mystery"}}, nil
}))
cmd.SetArgs([]string{"--wait"})
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
err := cmd.Execute()
if err == nil || !wait.IsUnknownStatus(err) {
t.Fatalf("err=%v, want unknown-status", err)
}
}
func TestWaitCtxAccessorsExposeDeclaredCapability(t *testing.T) {
var gotWait bool
var gotTimeout int
cmd := New(baseWaitSpec(waitTestDecl(), func(_ context.Context, c *Ctx) (wait.PollDoc, error) {
gotWait = c.Wait()
gotTimeout = c.WaitTimeoutSecs()
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
cmd.SetArgs([]string{"--wait", "--wait-timeout", "90"})
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
if err := cmd.Execute(); err != nil {
t.Fatal(err)
}
if !gotWait || gotTimeout != 90 {
t.Fatalf("ctx accessors=%v/%d", gotWait, gotTimeout)
}
}
func eventTestDecl(mode string) ContractDecl {
decl := waitTestDecl()
decl.Wait.Mode = mode
decl.Wait.EventKey = "bpms_instance_change"
decl.Wait.MatchField = "process_instance_id"
decl.Wait.ResourceQuery = "id"
return decl
}
type scriptedStream struct {
events []wait.PollDoc
err error
}
func (s *scriptedStream) Recv(context.Context) (wait.PollDoc, error) {
if len(s.events) > 0 {
doc := s.events[0]
s.events = s.events[1:]
return doc, nil
}
if s.err != nil {
return nil, s.err
}
return nil, io.EOF
}
func TestValidateWaitDeclPairsModeWithHooks(t *testing.T) {
poll := func(context.Context, *Ctx) (wait.PollDoc, error) { return nil, nil }
events := func(context.Context, *Ctx) (wait.EventStream, error) { return nil, nil }
cases := []struct {
name string
mode string
waitPoll bool
waitEvents bool
want string
}{
{"event without WaitEvents", contract.WaitModeEvent, false, false, "WaitEvents"},
{"auto without WaitPoll", contract.WaitModeAuto, false, true, "WaitPoll"},
{"poll with WaitEvents", contract.WaitModePoll, true, true, "WaitEvents"},
{"event with WaitPoll", contract.WaitModeEvent, true, true, "WaitPoll"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
expectPanic(t, func() {
New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: eventTestDecl(tc.mode),
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending(map[string]any{}, nil), nil
},
WaitPoll: hookOrNil(tc.waitPoll, poll),
WaitEvents: eventHookOrNil(tc.waitEvents, events),
})
}, tc.want)
})
}
}
func hookOrNil(set bool, hook func(context.Context, *Ctx) (wait.PollDoc, error)) func(context.Context, *Ctx) (wait.PollDoc, error) {
if !set {
return nil
}
return hook
}
func eventHookOrNil(set bool, hook func(context.Context, *Ctx) (wait.EventStream, error)) func(context.Context, *Ctx) (wait.EventStream, error) {
if !set {
return nil
}
return hook
}
func runWaitModeCommand(t *testing.T, decl ContractDecl, poll func(context.Context, *Ctx) (wait.PollDoc, error), events func(context.Context, *Ctx) (wait.EventStream, error), args ...string) (string, error) {
t.Helper()
cmd := New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: decl,
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending(map[string]any{"id": "job-1"}, &output.OperationInfo{
ID: "job-1", State: "NEW", NextCommand: "dws wait-sample --id job-1",
}), nil
},
WaitPoll: poll,
WaitEvents: events,
})
cmd.SetArgs(append([]string{"--wait"}, args...))
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
err := cmd.Execute()
return stdout.String(), err
}
func TestEventModeClosesEnvelopeFromCorrelatedEvent(t *testing.T) {
stream := &scriptedStream{events: []wait.PollDoc{
{"process_instance_id": "other", "result": map[string]any{"status": "COMPLETED"}},
{"process_instance_id": "job-1", "result": map[string]any{"status": "REJECTED"}},
}}
stdout, err := runWaitModeCommand(t, eventTestDecl(contract.WaitModeEvent), nil, func(context.Context, *Ctx) (wait.EventStream, error) {
return stream, nil
})
if err != nil {
t.Fatal(err)
}
if !strings.Contains(stdout, `"outcome": "failure"`) || !strings.Contains(stdout, `"type": "wait"`) {
t.Fatalf("stdout=%s", stdout)
}
}
func TestEventModeSurfacesStreamEndAsError(t *testing.T) {
_, err := runWaitModeCommand(t, eventTestDecl(contract.WaitModeEvent), nil, func(context.Context, *Ctx) (wait.EventStream, error) {
return &scriptedStream{}, nil
})
if err == nil || !errors.Is(err, wait.ErrEventStreamEnded) {
t.Fatalf("err=%v, want stream-ended", err)
}
}
func TestEventModeRejectsUnresolvableResource(t *testing.T) {
decl := eventTestDecl(contract.WaitModeEvent)
decl.Wait.ResourceQuery = "missing"
_, err := runWaitModeCommand(t, decl, nil, func(context.Context, *Ctx) (wait.EventStream, error) {
return &scriptedStream{}, nil
})
if err == nil || !strings.Contains(err.Error(), "resource query") {
t.Fatalf("err=%v", err)
}
}
func TestAutoModeFallsBackToPollOnStreamEnd(t *testing.T) {
polled := false
stdout, err := runWaitModeCommand(t, eventTestDecl(contract.WaitModeAuto),
func(context.Context, *Ctx) (wait.PollDoc, error) {
polled = true
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
},
func(context.Context, *Ctx) (wait.EventStream, error) {
return &scriptedStream{}, nil // ends immediately
})
if err != nil {
t.Fatal(err)
}
if !polled {
t.Fatal("auto mode did not fall back to polling")
}
if !strings.Contains(stdout, `"outcome": "success"`) {
t.Fatalf("stdout=%s", stdout)
}
}
func TestAutoModeFallsBackToPollOnSubscriptionFailure(t *testing.T) {
polled := false
_, err := runWaitModeCommand(t, eventTestDecl(contract.WaitModeAuto),
func(context.Context, *Ctx) (wait.PollDoc, error) {
polled = true
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
},
func(context.Context, *Ctx) (wait.EventStream, error) {
return nil, errors.New("no subscriber credential")
})
if err != nil {
t.Fatal(err)
}
if !polled {
t.Fatal("auto mode did not fall back to polling on subscription failure")
}
}
func TestResultInvokeWaitClosesEnvelopeOutcome(t *testing.T) {
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "REJECTED"}}, nil
}))
cmd.SetArgs([]string{"--wait"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
if err := cmd.Execute(); err != nil {
t.Fatal(err)
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 8 {
t.Fatalf("stored code/emitted=%d/%v, want dedicated wait-terminal-failure code 8", code, emitted)
}
if !strings.Contains(stdout.String(), `"type": "wait"`) {
t.Fatalf("stdout=%s, want error.type wait", stdout.String())
}
if !strings.Contains(stdout.String(), `"outcome": "failure"`) {
t.Fatalf("stdout=%s", stdout.String())
}
// The final emitted envelope must carry the observed terminal status in
// meta.operation.state — the acceptance-phase state ("NEW") must not
// survive the close (P1 regression guard).
if !strings.Contains(stdout.String(), `"state": "REJECTED"`) {
t.Fatalf("stdout=%s, want operation.state synced to the terminal status", stdout.String())
}
if strings.Contains(stdout.String(), `"state": "NEW"`) {
t.Fatalf("stdout=%s, acceptance-phase operation.state leaked into the terminal envelope", stdout.String())
}
if strings.Contains(stdout.String(), `"timed_out": true`) {
t.Fatalf("stdout=%s, terminal close must not claim timed_out", stdout.String())
}
}
func TestResultInvokeWaitSuccessCloseSyncsOperationState(t *testing.T) {
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
cmd.SetArgs([]string{"--wait"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
if err := cmd.Execute(); err != nil {
t.Fatal(err)
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v, want success exit 0", code, emitted)
}
if !strings.Contains(stdout.String(), `"outcome": "success"`) {
t.Fatalf("stdout=%s", stdout.String())
}
// Success close must publish the terminal status, never the stale
// acceptance-phase state (no outcome=success with state=processing/NEW).
if !strings.Contains(stdout.String(), `"state": "COMPLETED"`) {
t.Fatalf("stdout=%s, want operation.state synced to the terminal status", stdout.String())
}
if strings.Contains(stdout.String(), `"state": "NEW"`) {
t.Fatalf("stdout=%s, acceptance-phase operation.state leaked into the success envelope", stdout.String())
}
// Operation identity (id / next_command) survives the terminal close.
if !strings.Contains(stdout.String(), `"id": "job-1"`) || !strings.Contains(stdout.String(), `"next_command"`) {
t.Fatalf("stdout=%s, want operation id/next_command preserved", stdout.String())
}
}
func TestResultInvokeWaitTimeoutKeepsPending(t *testing.T) {
polls := 0
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
polls++
return wait.PollDoc{"result": map[string]any{"status": "RUNNING"}}, nil
}))
cmd.SetArgs([]string{"--wait", "--wait-timeout", "1"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
if err := cmd.Execute(); err != nil {
t.Fatalf("timeout wait must exit 0 (pending is not failure): %v", err)
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v", code, emitted)
}
if !strings.Contains(stdout.String(), `"outcome": "pending"`) {
t.Fatalf("stdout=%s", stdout.String())
}
if polls == 0 {
t.Fatal("wait phase never polled")
}
}
func TestWaitDeclRequiresResultInvokeDispatcher(t *testing.T) {
// A declared wait on the legacy Invoke path would observe a failure
// terminal while still exiting 0 — construction must reject it.
expectPanic(t, func() {
New(Spec{
Use: "wait-sample",
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: waitTestDecl(),
WaitPoll: func(context.Context, *Ctx) (wait.PollDoc, error) { return nil, nil },
Invoke: func(*Ctx, map[string]any) error { return nil },
})
}, "ResultInvoke")
}
func TestResultInvokeWithoutWaitFlagSkipsPhase(t *testing.T) {
polled := false
cmd := New(baseWaitSpec(waitTestDecl(), func(context.Context, *Ctx) (wait.PollDoc, error) {
polled = true
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
cmd.SetArgs(nil)
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
if err := cmd.Execute(); err != nil {
t.Fatal(err)
}
if polled {
t.Fatal("wait phase ran without --wait")
}
}
func TestAttachContractPanicsOnInvalidWaitDeclaration(t *testing.T) {
decl := waitTestDecl()
decl.Wait.Mode = "event" // not implemented
defer func() {
recovered := recover()
if recovered == nil {
t.Fatal("expected panic on invalid Contract.Wait")
}
if message, ok := recovered.(string); !ok || !strings.Contains(message, "Contract.Wait") {
t.Fatalf("panic=%v", recovered)
}
}()
AttachContract(&cobra.Command{Use: "wait-sample"}, contract.SafetySpec{
Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent",
}, decl, "", "")
}
func TestWaitDeclPaddedStatusValuesAreNormalized(t *testing.T) {
// A declaration whose status values carry surrounding whitespace must be
// canonicalized at construction so the runtime wait engine and the
// published Schema agree: the backend returns "COMPLETED" verbatim, and
// a padded terminal key would fail closed as an unknown status.
decl := waitTestDecl()
decl.Wait.Terminal = map[string]contract.ResultOutcome{
" COMPLETED ": contract.ResultOutcomeSuccess,
"\tREJECTED": contract.ResultOutcomeFailure,
}
decl.Wait.PendingValues = []string{" NEW ", "RUNNING "}
cmd := New(baseWaitSpec(decl, func(context.Context, *Ctx) (wait.PollDoc, error) {
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
}))
cmd.SetArgs([]string{"--wait"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
if err := cmd.Execute(); err != nil {
t.Fatalf("padded declaration must still reach the terminal status: %v", err)
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v, want success", code, emitted)
}
final, ok := contractfinal.RuntimeContractFinal(cmd)
if !ok || final.Wait == nil {
t.Fatal("registered ContractFinal lost the wait capability")
}
if _, ok := final.Wait.Terminal["COMPLETED"]; !ok {
t.Fatalf("registered terminal table not trimmed: %#v", final.Wait.Terminal)
}
for _, value := range final.Wait.PendingValues {
if strings.TrimSpace(value) != value {
t.Fatalf("registered pending value %q not trimmed", value)
}
}
}
func TestNewPanicsOnDuplicateOrConflictingWaitStatusesAfterTrim(t *testing.T) {
// Values that collapse onto one status after trimming are programming
// errors: silently merging them would pick one outcome for two authored
// declarations.
dupTerminal := waitTestDecl()
dupTerminal.Wait.Terminal = map[string]contract.ResultOutcome{
"COMPLETED": contract.ResultOutcomeSuccess,
" COMPLETED": contract.ResultOutcomeFailure,
}
expectPanic(t, func() { New(baseWaitSpec(dupTerminal, nil)) }, "Contract.Wait")
conflict := waitTestDecl()
conflict.Wait.PendingValues = []string{" COMPLETED "}
expectPanic(t, func() { New(baseWaitSpec(conflict, nil)) }, "Contract.Wait")
}
func TestContractDeclEmptyTreatsWaitAsAuthored(t *testing.T) {
// Only Wait is authored: empty() must report non-empty through the Wait
// branch (before validateContractDecl then fails on the missing prose).
decl := ContractDecl{Wait: &contract.WaitSpec{
Mode: contract.WaitModePoll,
PollCommand: "oa approval-instance get",
StatusQuery: "result.status",
Terminal: map[string]contract.ResultOutcome{"COMPLETED": contract.ResultOutcomeSuccess},
}}
if decl.Empty() {
t.Fatal("Wait-only declaration must count as authored")
}
defer func() {
if recover() == nil {
t.Fatal("expected validateContractDecl to reject the missing prose")
}
}()
validateContractDecl(Spec{Use: "wait-only", Contract: decl})
}
func TestEventModeRejectsNonObjectResultData(t *testing.T) {
cmd := New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: eventTestDecl(contract.WaitModeEvent),
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending([]any{"not", "an", "object"}, &output.OperationInfo{
ID: "job-1", State: "NEW", NextCommand: "dws wait-sample",
}), nil
},
WaitEvents: func(context.Context, *Ctx) (wait.EventStream, error) {
return &scriptedStream{}, nil
},
})
cmd.SetArgs([]string{"--wait"})
ctx, _ := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
err := cmd.Execute()
if err == nil || !strings.Contains(err.Error(), "not an object") {
t.Fatalf("err=%v, want non-object data rejection", err)
}
}
func TestEventModeSubscriptionFailureSurfacesInStrictMode(t *testing.T) {
decl := eventTestDecl(contract.WaitModeEvent)
_, err := runWaitModeCommand(t, decl, nil, func(context.Context, *Ctx) (wait.EventStream, error) {
return nil, errors.New("no subscriber credential")
})
if err == nil || !strings.Contains(err.Error(), "subscription failed") {
t.Fatalf("err=%v, want subscription failure surfaced", err)
}
}
func TestResultInvokeNonPendingSkipsWaitPhase(t *testing.T) {
partial, err := output.NewPartialData(2,
[]any{map[string]any{"id": "ok"}},
[]output.PartialFailedEntry{{ID: "bad", Error: &output.ErrorInfo{Type: "api", Message: "item failed"}}},
nil)
if err != nil {
t.Fatal(err)
}
cases := []struct {
name string
result output.CommandResult
want string
}{
{"failure", output.Failure(&output.ErrorInfo{Type: "api", Message: "business failed"}), "failure"},
{"success", output.Success(map[string]any{"id": "job-1"}), "success"},
{"partial", output.Partial(partial), "partial_failure"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
polled := false
subscribed := false
cmd := New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: eventTestDecl(contract.WaitModeAuto),
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return tc.result, nil
},
WaitPoll: func(context.Context, *Ctx) (wait.PollDoc, error) {
polled = true
return wait.PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
},
WaitEvents: func(context.Context, *Ctx) (wait.EventStream, error) {
subscribed = true
return &scriptedStream{}, nil
},
})
cmd.SetArgs([]string{"--wait"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
if err := cmd.Execute(); err != nil {
t.Fatal(err)
}
if polled || subscribed {
t.Fatal("wait phase must not call WaitPoll/WaitEvents for a non-pending initial result")
}
if _, emitted := output.StoredExitCode(store); !emitted {
t.Fatal("initial result was not stored")
}
if !strings.Contains(stdout.String(), `"outcome": "`+tc.want+`"`) {
t.Fatalf("stdout=%s, want outcome %s preserved", stdout.String(), tc.want)
}
if strings.Contains(stdout.String(), `"type": "wait"`) {
t.Fatalf("stdout=%s, wait phase overwrote the original envelope", stdout.String())
}
})
}
}
func TestWaitTimeoutCancelsBlockingPoll(t *testing.T) {
started := make(chan struct{})
cmd := New(baseWaitSpec(waitTestDecl(), func(ctx context.Context, c *Ctx) (wait.PollDoc, error) {
close(started)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-c.Command().Context().Done():
return nil, c.Command().Context().Err()
}
}))
cmd.SetArgs([]string{"--wait", "--wait-timeout", "1"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
done := make(chan error, 1)
go func() { done <- cmd.Execute() }()
select {
case <-started:
case <-time.After(2 * time.Second):
t.Fatal("blocking poll never started")
}
select {
case err := <-done:
if err != nil {
t.Fatalf("timeout wait must exit 0 (pending is not failure): %v", err)
}
case <-time.After(3 * time.Second):
t.Fatal("blocked poll was not cancelled by --wait-timeout")
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v", code, emitted)
}
if !strings.Contains(stdout.String(), `"outcome": "pending"`) {
t.Fatalf("stdout=%s", stdout.String())
}
}
func TestWaitTimeoutCancelsBlockingSubscribe(t *testing.T) {
started := make(chan struct{})
cmd := New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: eventTestDecl(contract.WaitModeEvent),
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending(map[string]any{"id": "job-1"}, &output.OperationInfo{
ID: "job-1", State: "NEW", NextCommand: "dws wait-sample --id job-1",
}), nil
},
WaitEvents: func(ctx context.Context, c *Ctx) (wait.EventStream, error) {
close(started)
// Leaf subscribe may wait on either the hook ctx or the cobra
// command context; both must carry the wait-timeout deadline.
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-c.Command().Context().Done():
return nil, c.Command().Context().Err()
}
},
})
cmd.SetArgs([]string{"--wait", "--wait-timeout", "1"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
done := make(chan error, 1)
go func() { done <- cmd.Execute() }()
select {
case <-started:
case <-time.After(2 * time.Second):
t.Fatal("blocking subscribe never started")
}
select {
case err := <-done:
if err != nil {
t.Fatalf("timeout wait must exit 0 (pending is not failure): %v", err)
}
case <-time.After(3 * time.Second):
t.Fatal("blocked subscribe was not cancelled by --wait-timeout")
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v", code, emitted)
}
if !strings.Contains(stdout.String(), `"outcome": "pending"`) {
t.Fatalf("stdout=%s", stdout.String())
}
}
func TestWaitTimeoutCancelsBlockingPollAfterAutoFallback(t *testing.T) {
started := make(chan struct{})
cmd := New(Spec{
Use: "wait-sample",
OutputRollout: output.RolloutUnifiedActive,
Safety: contract.SafetySpec{Effect: "read", Risk: "low", Confirmation: "not_required", Idempotency: "idempotent"},
Contract: eventTestDecl(contract.WaitModeAuto),
ResultInvoke: func(*Ctx, map[string]any) (output.CommandResult, error) {
return output.Pending(map[string]any{"id": "job-1"}, &output.OperationInfo{
ID: "job-1", State: "NEW", NextCommand: "dws wait-sample --id job-1",
}), nil
},
WaitEvents: func(context.Context, *Ctx) (wait.EventStream, error) {
return &scriptedStream{}, nil // ends immediately → poll fallback
},
WaitPoll: func(ctx context.Context, _ *Ctx) (wait.PollDoc, error) {
close(started)
<-ctx.Done()
return nil, ctx.Err()
},
})
cmd.SetArgs([]string{"--wait", "--wait-timeout", "1"})
ctx, store := output.WithResultStore(context.Background())
cmd.SetContext(ctx)
var stdout bytes.Buffer
cmd.SetOut(&stdout)
cmd.PersistentPostRunE = func(executed *cobra.Command, _ []string) error {
_, _, err := output.EmitStoredResult(executed)
return err
}
done := make(chan error, 1)
go func() { done <- cmd.Execute() }()
select {
case <-started:
case <-time.After(2 * time.Second):
t.Fatal("auto-fallback blocking poll never started")
}
select {
case err := <-done:
if err != nil {
t.Fatalf("timeout wait must exit 0 (pending is not failure): %v", err)
}
case <-time.After(3 * time.Second):
t.Fatal("blocked auto-fallback poll was not cancelled by --wait-timeout")
}
if code, emitted := output.StoredExitCode(store); !emitted || code != 0 {
t.Fatalf("stored code/emitted=%d/%v", code, emitted)
}
if !strings.Contains(stdout.String(), `"outcome": "pending"`) {
t.Fatalf("stdout=%s", stdout.String())
}
}
-2
View File
@@ -58,8 +58,6 @@ const (
// 5 internal (CategoryInternal 与兜底:非结构化错误、panic 收敛均归 5)
// 6 discovery (CategoryDiscovery)
// 7 partial_failure(部分成功专用码,见 ExitCodePartial)
// 8 wait (--wait 观察到失败终态的专用码,见 internal/output
// 的 exitCodeWait;不设 Category,仅经统一信封产出)
//
// ExitCodePartial is the partial-result exit code shared with internal/output.
// It is not returned for CategoryPartial errors because they lack the typed
-1
View File
@@ -132,7 +132,6 @@ func TestClientCreateRuleBasedSubscriptionsUsesDocumentedRuleParam(t *testing.T)
{"oa_approval_task_finished", EventOAApprovalTaskFinished, RuleOptions{}, map[string]any{}},
{"oa_approval_task_redirected", EventOAApprovalTaskRedirected, RuleOptions{}, map[string]any{}},
{"oa_approval_instance_started", EventOAApprovalInstanceStarted, RuleOptions{}, map[string]any{}},
{"oa_approval_instance_cc", EventOAApprovalInstanceCC, RuleOptions{}, map[string]any{}},
{"oa_approval_instance_terminated", EventOAApprovalInstanceTerminated, RuleOptions{}, map[string]any{}},
{"oa_approval_instance_finished", EventOAApprovalInstanceFinished, RuleOptions{}, map[string]any{}},
{"read_group", EventReadGroup, RuleOptions{GroupID: "cid-1"}, map[string]any{"openConversationId": "cid-1"}},
-29
View File
@@ -180,19 +180,6 @@ type OAApprovalInstanceStartedOutput struct {
EventTime int64 `json:"event_time" description:"审批实例事件业务时间" format:"timestamp_ms"`
}
type OAApprovalInstanceCCOutput struct {
Type string `json:"type" description:"事件类型,固定为当前 event_key"`
EventID string `json:"event_id" description:"事件 ID,可用于去重"`
Timestamp int64 `json:"timestamp" description:"事件发生时间戳" format:"timestamp_ms"`
SubscribeID string `json:"subscribe_id" description:"订阅 ID"`
ProcessInstanceID string `json:"process_instance_id" description:"审批实例 ID"`
ProcessCode string `json:"process_code" description:"审批流程模板编码"`
Title string `json:"title" description:"审批标题"`
Status string `json:"status" description:"审批实例到达抄送节点时的状态"`
CreateTime int64 `json:"create_time" description:"审批实例创建时间" format:"timestamp_ms"`
EventTime int64 `json:"event_time" description:"审批抄送事件业务时间" format:"timestamp_ms"`
}
type OAApprovalInstanceTerminatedOutput struct {
Type string `json:"type" description:"事件类型,固定为当前 event_key"`
EventID string `json:"event_id" description:"事件 ID,可用于去重"`
@@ -680,19 +667,6 @@ func projectOAApprovalEvent(ev transport.Event, base baseEventOutput, raw json.R
CreateTime: payload.Body.CreateTime,
EventTime: payload.EventTime,
}, nil
case EventOAApprovalInstanceCC:
return OAApprovalInstanceCCOutput{
Type: base.Type,
EventID: base.EventID,
Timestamp: base.Timestamp,
SubscribeID: base.SubscribeID,
ProcessInstanceID: payload.Body.ProcessInstanceID,
ProcessCode: payload.Body.ProcessCode,
Title: payload.Body.Title,
Status: payload.Body.Status,
CreateTime: payload.Body.CreateTime,
EventTime: payload.EventTime,
}, nil
case EventOAApprovalInstanceTerminated:
return OAApprovalInstanceTerminatedOutput{
Type: base.Type,
@@ -897,8 +871,6 @@ func outputTypeForEvent(eventKey string) reflect.Type {
return reflect.TypeOf(OAApprovalTaskRedirectedOutput{})
case eventKey == EventOAApprovalInstanceStarted:
return reflect.TypeOf(OAApprovalInstanceStartedOutput{})
case eventKey == EventOAApprovalInstanceCC:
return reflect.TypeOf(OAApprovalInstanceCCOutput{})
case eventKey == EventOAApprovalInstanceTerminated:
return reflect.TypeOf(OAApprovalInstanceTerminatedOutput{})
case eventKey == EventOAApprovalInstanceFinished:
@@ -934,7 +906,6 @@ func isOAEvent(eventKey string) bool {
eventKey == EventOAApprovalTaskFinished ||
eventKey == EventOAApprovalTaskRedirected ||
eventKey == EventOAApprovalInstanceStarted ||
eventKey == EventOAApprovalInstanceCC ||
eventKey == EventOAApprovalInstanceTerminated ||
eventKey == EventOAApprovalInstanceFinished
}
-18
View File
@@ -173,8 +173,6 @@ func personalOAData(eventKey string) string {
body["finishTime"] = int64(1785229199000)
case EventOAApprovalInstanceStarted:
body["status"] = "RUNNING"
case EventOAApprovalInstanceCC:
body["status"] = "RUNNING"
case EventOAApprovalInstanceTerminated:
body["status"] = "TERMINATED"
body["finishTime"] = int64(1785229199000)
@@ -523,21 +521,6 @@ func TestCrossPlatformCoverageProjectOutputOAEvents(t *testing.T) {
EventTime: 1785229199000,
},
},
{
eventKey: EventOAApprovalInstanceCC,
want: OAApprovalInstanceCCOutput{
Type: EventOAApprovalInstanceCC,
EventID: "oa-event",
Timestamp: 1785229200123,
SubscribeID: "outer-sub",
ProcessInstanceID: "process-instance-1",
ProcessCode: "PROC-TEST-1",
Title: "测试审批",
Status: "RUNNING",
CreateTime: 1785229100000,
EventTime: 1785229199000,
},
},
{
eventKey: EventOAApprovalInstanceTerminated,
want: OAApprovalInstanceTerminatedOutput{
@@ -802,7 +785,6 @@ func TestCrossPlatformCoverageProjectOutputRejectsInvalidOAPayloads(t *testing.T
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
} {
-12
View File
@@ -43,7 +43,6 @@ const (
EventOAApprovalTaskFinished = "user_oa_approval_task_finished"
EventOAApprovalTaskRedirected = "user_oa_approval_task_redirected"
EventOAApprovalInstanceStarted = "user_oa_approval_instance_started"
EventOAApprovalInstanceCC = "user_oa_approval_instance_cc"
EventOAApprovalInstanceTerminated = "user_oa_approval_instance_terminated"
EventOAApprovalInstanceFinished = "user_oa_approval_instance_finished"
)
@@ -324,17 +323,6 @@ var definitions = []Definition{
Auth: map[string]any{"identity": "user"},
Public: true,
},
{
EventKey: EventOAApprovalInstanceCC,
DisplayName: "审批单抄送",
Description: "审批实例到达抄送节点,发送给被抄送人",
Category: "oa",
RuleType: "all",
Status: StatusEnabled,
RequiredParams: nil,
Auth: map[string]any{"identity": "user"},
Public: true,
},
{
EventKey: EventOAApprovalInstanceTerminated,
DisplayName: "审批单终止",
-12
View File
@@ -50,7 +50,6 @@ func TestCatalogEnabledEvents(t *testing.T) {
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
}
@@ -66,7 +65,6 @@ func TestOAEventCatalogDefinitions(t *testing.T) {
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
}
@@ -160,7 +158,6 @@ func TestSchemaDocumentsDefaultToTransportEnvelope(t *testing.T) {
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
} {
@@ -512,13 +509,6 @@ func TestOAEventSchemaDocumentsMatchOutputDTO(t *testing.T) {
"process_code", "title", "status", "create_time", "event_time",
},
},
{
eventKey: EventOAApprovalInstanceCC,
properties: []string{
"type", "event_id", "timestamp", "subscribe_id", "process_instance_id",
"process_code", "title", "status", "create_time", "event_time",
},
},
{
eventKey: EventOAApprovalInstanceTerminated,
properties: []string{
@@ -643,7 +633,6 @@ func TestBuildRuleParamAllEvents(t *testing.T) {
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
} {
@@ -856,7 +845,6 @@ func TestSupportsMessageFilter(t *testing.T) {
EventOAApprovalTaskFinished,
EventOAApprovalTaskRedirected,
EventOAApprovalInstanceStarted,
EventOAApprovalInstanceCC,
EventOAApprovalInstanceTerminated,
EventOAApprovalInstanceFinished,
"unknown_event",
+40
View File
@@ -3,6 +3,7 @@
package keychain
import (
"bytes"
"encoding/base64"
"errors"
"os"
@@ -246,6 +247,45 @@ func TestCrossPlatformCoverageDarwinDEKKeyringEdges(t *testing.T) {
if got, err := getOrCreateDEK("generate-missing"); err != nil || len(got) != dekBytes {
t.Fatalf("generate missing = %d, %v", len(got), err)
}
existing := bytesOf(7, dekBytes)
setCalls := 0
keyringGet = func(string, string) (string, error) {
return base64.StdEncoding.EncodeToString(existing), nil
}
keyringSet = func(string, string, string) error {
setCalls++
return errors.New("duplicate write must not replace an existing DEK")
}
got, err := getOrCreateDEK("reuse-existing")
if err != nil || !bytes.Equal(got, existing) {
t.Fatalf("reuse existing = %d, %v", len(got), err)
}
if setCalls != 0 {
t.Fatalf("reuse existing set calls = %d, want 0", setCalls)
}
gets := 0
setCalls = 0
keyringGet = func(string, string) (string, error) {
gets++
if gets == 1 {
return "", keyring.ErrNotFound
}
return encoded, nil
}
keyringSet = func(string, string, string) error {
setCalls++
return errors.New("already exists")
}
got, err = getOrCreateDEK("create-race")
if err != nil || !bytes.Equal(got, valid) {
t.Fatalf("create race = %d, %v", len(got), err)
}
if setCalls != 1 {
t.Fatalf("create race set calls = %d, want 1", setCalls)
}
keyringGet = func(string, string) (string, error) { return "", keyring.ErrNotFound }
keychainRandRead = func([]byte) (int, error) { return 0, errKeychainInjected }
if _, err := getOrCreateDEK("rand"); err == nil {
t.Fatal("rand error expected")
+37 -26
View File
@@ -237,6 +237,26 @@ func getSystemDEKReadOnly(service string) ([]byte, error) {
return key, err
}
func decodeSystemDEK(encodedKey string) ([]byte, bool) {
key, err := base64.StdEncoding.DecodeString(encodedKey)
return key, err == nil && len(key) == dekBytes
}
func readStoredSystemDEK(service string, runtime darwinKeychainRuntime) ([]byte, error) {
encodedKey, err := runtime.get(service, "dek")
if err == nil {
key, ok := decodeSystemDEK(encodedKey)
if ok {
return key, nil
}
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
if errors.Is(err, keyring.ErrNotFound) {
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
}
func getSystemDEKReadOnlyWithRuntime(service string, runtime darwinKeychainRuntime) ([]byte, error, <-chan struct{}) {
if err := runtime.checkAvailable(); err != nil {
return nil, err, finishedDarwinKeychainWorker()
@@ -244,18 +264,7 @@ func getSystemDEKReadOnlyWithRuntime(service string, runtime darwinKeychainRunti
const operation = "read DEK from macOS Keychain"
worker := startDarwinKeychainWorker(operation, func() ([]byte, error) {
encodedKey, err := runtime.get(service, "dek")
if err == nil {
key, decodeErr := base64.StdEncoding.DecodeString(encodedKey)
if decodeErr == nil && len(key) == dekBytes {
return key, nil
}
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
if errors.Is(err, keyring.ErrNotFound) {
return nil, fmt.Errorf("read DEK from macOS Keychain: %w", ErrDEKMissing)
}
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
return readStoredSystemDEK(service, runtime)
})
return waitDarwinKeychainWorker(runtime.timeout, operation, worker)
@@ -276,28 +285,30 @@ func getOrCreateDEKWithRuntime(service string, runtime darwinKeychainRuntime) ([
const operation = "read or create DEK in macOS Keychain"
worker := startDarwinKeychainWorker(operation, func() ([]byte, error) {
// Try to get existing DEK from system Keychain
encodedKey, err := runtime.get(service, "dek")
key, err := readStoredSystemDEK(service, runtime)
if err == nil {
key, decodeErr := base64.StdEncoding.DecodeString(encodedKey)
if decodeErr == nil && len(key) == dekBytes {
return key, nil
}
} else if !errors.Is(err, keyring.ErrNotFound) {
return nil, NewUnavailableError("read DEK from macOS Keychain", err)
return key, nil
}
if !IsDEKMissing(err) {
return nil, err
}
// Generate new DEK if not found or invalid
key := make([]byte, dekBytes)
// Generate a candidate only when the slot is empty or unreadable.
// Concurrent writers must not replace a DEK another process just stored.
key = make([]byte, dekBytes)
if _, randErr := runtime.randRead(key); randErr != nil {
return nil, randErr
}
// Store in system Keychain
encodedKey = base64.StdEncoding.EncodeToString(key)
if setErr := runtime.set(service, "dek", encodedKey); setErr != nil {
if setErr := runtime.set(service, "dek", base64.StdEncoding.EncodeToString(key)); setErr != nil {
existing, getErr := readStoredSystemDEK(service, runtime)
if getErr == nil {
return existing, nil
}
return nil, NewUnavailableError("store DEK in macOS Keychain", setErr)
}
if existing, getErr := readStoredSystemDEK(service, runtime); getErr == nil {
return existing, nil
}
return key, nil
})
+127
View File
@@ -0,0 +1,127 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"fmt"
"sync"
"time"
)
// DefaultAuthCodeTTL is the unconsumed-cache window for a freshly minted
// vendor authCode. The portal issues codes with expiresIn=120s and they are
// one-shot, so the default stays under that server window. Prefer not wrapping
// PortalAuthCode in CachedAuthCode: mint in goProxy and discard after the key
// request.
const DefaultAuthCodeTTL = 90 * time.Second
// ErrNoAuthCode means the provider returned an empty code without an error.
var ErrNoAuthCode = errors.New("msgcrypto: auth code provider returned an empty code")
// AuthCodeProvider yields a DingTalk 免登 authCode for key-server
// authentication. DWS does not mint the code itself, so integrations inject an
// implementation. The backend calls this only from the vendor goProxy
// callback, never on every encrypt or decrypt.
//
// Implementations must be safe for concurrent use; the backend may call this
// from a CGO callback while an encrypt or decrypt call is in flight.
type AuthCodeProvider interface {
AuthCode(ctx context.Context) (string, error)
}
// CorpAuthCodeProvider mints a code for a specific organization. PortalAuthCode
// implements this so goProxy can pass the C library's corpID. Domain and
// redirectURI are never part of this call.
type CorpAuthCodeProvider interface {
AuthCodeProvider
AuthCodeForCorp(ctx context.Context, corpID string) (string, error)
}
// AuthCodeFunc adapts a function to AuthCodeProvider.
type AuthCodeFunc func(ctx context.Context) (string, error)
// AuthCode calls f.
func (f AuthCodeFunc) AuthCode(ctx context.Context) (string, error) { return f(ctx) }
// StaticAuthCode returns a provider that always yields code. It is meant for
// tests and manual integration runs; a static code stops working once the
// server-side five-minute window closes.
func StaticAuthCode(code string) AuthCodeProvider {
return AuthCodeFunc(func(context.Context) (string, error) {
if code == "" {
return "", ErrNoAuthCode
}
return code, nil
})
}
// CachedAuthCode memoises an AuthCodeProvider for a TTL so a burst of key
// requests does not trigger one upstream call each.
type CachedAuthCode struct {
provider AuthCodeProvider
ttl time.Duration
now func() time.Time
mu sync.Mutex
code string
expiresAt time.Time
}
// NewCachedAuthCode wraps provider with a TTL cache. A ttl of zero or less
// selects DefaultAuthCodeTTL.
func NewCachedAuthCode(provider AuthCodeProvider, ttl time.Duration) *CachedAuthCode {
if ttl <= 0 {
ttl = DefaultAuthCodeTTL
}
return &CachedAuthCode{provider: provider, ttl: ttl, now: time.Now}
}
// AuthCode returns the cached code when it is still fresh, otherwise fetches a
// new one. A failed fetch leaves no stale value behind.
func (c *CachedAuthCode) AuthCode(ctx context.Context) (string, error) {
if c.provider == nil {
return "", ErrNoAuthCodeProvider
}
c.mu.Lock()
defer c.mu.Unlock()
if c.code != "" && c.now().Before(c.expiresAt) {
return c.code, nil
}
code, err := c.provider.AuthCode(ctx)
if err != nil {
c.code, c.expiresAt = "", time.Time{}
return "", fmt.Errorf("msgcrypto: fetch auth code: %w", err)
}
if code == "" {
c.code, c.expiresAt = "", time.Time{}
return "", ErrNoAuthCode
}
c.code = code
c.expiresAt = c.now().Add(c.ttl)
return code, nil
}
// Invalidate drops the cached code so the next AuthCode call refetches. The
// backend calls this after the key server rejects a code.
func (c *CachedAuthCode) Invalidate() {
c.mu.Lock()
c.code, c.expiresAt = "", time.Time{}
c.mu.Unlock()
}
+228
View File
@@ -0,0 +1,228 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync"
"testing"
"time"
)
// countingProvider hands out a fresh code per call and records how often it was
// asked, so cache behaviour can be asserted.
type countingProvider struct {
mu sync.Mutex
calls int
code string
err error
}
func (p *countingProvider) AuthCode(context.Context) (string, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.calls++
if p.err != nil {
return "", p.err
}
if p.code != "" {
return p.code, nil
}
return "code-" + string(rune('a'+p.calls-1)), nil
}
// callCount reports the number of upstream fetches.
func (p *countingProvider) callCount() int {
p.mu.Lock()
defer p.mu.Unlock()
return p.calls
}
func TestAuthCodeFuncAdaptsFunction(t *testing.T) {
provider := AuthCodeFunc(func(context.Context) (string, error) { return "abc", nil })
code, err := provider.AuthCode(context.Background())
if err != nil || code != "abc" {
t.Fatalf("AuthCode() = %q, %v; want abc, nil", code, err)
}
}
func TestStaticAuthCodeReturnsCode(t *testing.T) {
code, err := StaticAuthCode("fixed").AuthCode(context.Background())
if err != nil || code != "fixed" {
t.Fatalf("AuthCode() = %q, %v; want fixed, nil", code, err)
}
}
func TestStaticAuthCodeRejectsEmptyCode(t *testing.T) {
_, err := StaticAuthCode("").AuthCode(context.Background())
if !errors.Is(err, ErrNoAuthCode) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCode", err)
}
}
func TestCachedAuthCodeReusesCodeWithinTTL(t *testing.T) {
provider := &countingProvider{code: "same"}
cache := NewCachedAuthCode(provider, time.Minute)
for i := 0; i < 5; i++ {
code, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("AuthCode() #%d = %v", i+1, err)
}
if code != "same" {
t.Fatalf("AuthCode() #%d = %q, want same", i+1, code)
}
}
if got := provider.callCount(); got != 1 {
t.Fatalf("upstream called %d times, want 1 (the code must be cached)", got)
}
}
func TestCachedAuthCodeRefetchesAfterTTL(t *testing.T) {
provider := &countingProvider{}
cache := NewCachedAuthCode(provider, time.Minute)
now := time.Now()
cache.now = func() time.Time { return now }
first, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("first AuthCode() = %v", err)
}
// Move past the TTL. The DingTalk code expires server-side, so a stale
// one must not be reused.
now = now.Add(time.Minute + time.Second)
second, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("second AuthCode() = %v", err)
}
if first == second {
t.Fatalf("AuthCode() returned the same code %q after the TTL expired", first)
}
if got := provider.callCount(); got != 2 {
t.Fatalf("upstream called %d times, want 2", got)
}
}
func TestCachedAuthCodeDefaultTTLIsUnderServerWindow(t *testing.T) {
// Portal vendorAuthCode expiresIn is 120s and the code is one-shot.
// The unconsumed-cache window must stay under that server lifetime.
if DefaultAuthCodeTTL >= 120*time.Second {
t.Fatalf("DefaultAuthCodeTTL = %v, want less than the 120s portal expiresIn", DefaultAuthCodeTTL)
}
cache := NewCachedAuthCode(&countingProvider{}, 0)
if cache.ttl != DefaultAuthCodeTTL {
t.Fatalf("ttl = %v, want DefaultAuthCodeTTL %v", cache.ttl, DefaultAuthCodeTTL)
}
}
func TestCachedAuthCodeNegativeTTLFallsBackToDefault(t *testing.T) {
cache := NewCachedAuthCode(&countingProvider{}, -time.Second)
if cache.ttl != DefaultAuthCodeTTL {
t.Fatalf("ttl = %v, want DefaultAuthCodeTTL %v", cache.ttl, DefaultAuthCodeTTL)
}
}
func TestCachedAuthCodePropagatesUpstreamError(t *testing.T) {
wantErr := errors.New("token service down")
cache := NewCachedAuthCode(&countingProvider{err: wantErr}, time.Minute)
_, err := cache.AuthCode(context.Background())
if !errors.Is(err, wantErr) {
t.Fatalf("AuthCode() = %v, want it to wrap %v", err, wantErr)
}
}
func TestCachedAuthCodeDoesNotCacheFailures(t *testing.T) {
provider := &countingProvider{err: errors.New("transient")}
cache := NewCachedAuthCode(provider, time.Minute)
if _, err := cache.AuthCode(context.Background()); err == nil {
t.Fatal("AuthCode() = nil error, want failure")
}
provider.mu.Lock()
provider.err = nil
provider.code = "recovered"
provider.mu.Unlock()
code, err := cache.AuthCode(context.Background())
if err != nil {
t.Fatalf("AuthCode() after recovery = %v", err)
}
if code != "recovered" {
t.Fatalf("AuthCode() = %q, want recovered (a failure must not be cached)", code)
}
}
func TestCachedAuthCodeRejectsEmptyUpstreamCode(t *testing.T) {
// A provider that reports success with no code is a bug upstream; the
// cache must surface it instead of caching an unusable value.
cache := NewCachedAuthCode(AuthCodeFunc(func(context.Context) (string, error) {
return "", nil
}), time.Minute)
if _, err := cache.AuthCode(context.Background()); !errors.Is(err, ErrNoAuthCode) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCode", err)
}
}
func TestCachedAuthCodeInvalidateForcesRefetch(t *testing.T) {
provider := &countingProvider{}
cache := NewCachedAuthCode(provider, time.Hour)
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("first AuthCode() = %v", err)
}
cache.Invalidate()
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("second AuthCode() = %v", err)
}
if got := provider.callCount(); got != 2 {
t.Fatalf("upstream called %d times, want 2 after Invalidate", got)
}
}
func TestCachedAuthCodeWithoutProviderReportsMissingProvider(t *testing.T) {
cache := NewCachedAuthCode(nil, time.Minute)
if _, err := cache.AuthCode(context.Background()); !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("AuthCode() = %v, want ErrNoAuthCodeProvider", err)
}
}
func TestCachedAuthCodeIsSafeForConcurrentUse(t *testing.T) {
// The backend may ask for a code from a CGO callback while another
// operation is in flight, so concurrent access must not race.
provider := &countingProvider{code: "shared"}
cache := NewCachedAuthCode(provider, time.Hour)
var wg sync.WaitGroup
for i := 0; i < 32; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if code, err := cache.AuthCode(context.Background()); err != nil || code != "shared" {
t.Errorf("AuthCode() = %q, %v; want shared, nil", code, err)
}
}()
}
wg.Wait()
if got := provider.callCount(); got != 1 {
t.Fatalf("upstream called %d times, want 1", got)
}
}
+159
View File
@@ -0,0 +1,159 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// The constraint below must stay in sync with cipher_stub.go, which negates it
// verbatim. It encodes the platforms the vendor ships a libsafechat.a for:
// darwin and linux on amd64/arm64, plus windows/amd64. windows/arm64 is
// deliberately excluded because the vendor has not delivered that static
// library, and DWS does release that target.
//go:build safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64))
package msgcrypto
import (
"context"
"errors"
"fmt"
"sync"
safechat "safechat-go-sdk"
)
// BackendVersion identifies the compiled-in vendor SDK.
const BackendVersion = "safechat " + safechat.Version
// Available reports that this binary carries the SafeChat backend.
func Available() bool { return true }
// safechatCipher adapts the vendor client to Cipher.
//
// The vendor client serialises its own C calls internally, so this type adds no
// further locking. Auth codes are minted only from AuthCodeHook, which the
// vendor SDK calls inside goProxy when a key is actually missing.
type safechatCipher struct {
client *safechat.Client
codes AuthCodeProvider
allowedHost string
mu sync.Mutex
lastCodeErr error
}
// newBackend starts the vendor client against cfg's keystore.
//
// A warm keystore serves encrypt and decrypt without a key request, so no
// authCode is fetched at open time. The hook runs only if goProxy fires.
func newBackend(ctx context.Context, cfg Config) (Cipher, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
var logf func(string, ...any)
if cfg.Debug {
logf = cfg.Logf
}
c := &safechatCipher{codes: cfg.AuthCode, allowedHost: cfg.AllowedRedirectHost}
client, err := safechat.New(safechat.Config{
DataPath: cfg.KeystoreDir,
UserID: cfg.UserID,
KeyServer: cfg.KeyServer,
MaxRetry: cfg.MaxRetry,
HTTPTimeout: cfg.HTTPTimeout,
Logger: newRedactingLogger(logf),
AuthCodeHook: c.authCodeHook,
})
if err != nil {
if errors.Is(err, safechat.ErrAlreadyInitialized) {
return nil, ErrAlreadyOpen
}
return nil, fmt.Errorf("msgcrypto: start safechat backend: %w", err)
}
c.client = client
return c, nil
}
// EncryptMessage encrypts plaintext and returns the vendor ciphertext.
func (c *safechatCipher) EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
c.setLastCodeErr(nil)
out, err := c.client.EncryptMsg(corpID, staffID, plaintext)
if err != nil {
return nil, c.explain("encrypt", corpID, err)
}
return out, nil
}
// DecryptMessage decrypts a vendor ciphertext.
func (c *safechatCipher) DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
c.setLastCodeErr(nil)
out, err := c.client.DecryptMsg(corpID, staffID, ciphertext)
if err != nil {
return nil, c.explain("decrypt", corpID, err)
}
return out, nil
}
// Close releases the vendor client.
func (c *safechatCipher) Close() error {
c.client.Close()
return nil
}
// authCodeHook is invoked from goProxy immediately before the key request.
// domain is compared locally and never forwarded to portal. The returned
// code is used once by the SDK and is not stored on the client.
func (c *safechatCipher) authCodeHook(corpID, domain string) (string, error) {
code, err := mintAuthCodeForProxy(c.codes, c.allowedHost, corpID, domain)
c.setLastCodeErr(err)
return code, err
}
func (c *safechatCipher) setLastCodeErr(err error) {
c.mu.Lock()
c.lastCodeErr = err
c.mu.Unlock()
}
func (c *safechatCipher) lastAuthCodeErr() error {
c.mu.Lock()
defer c.mu.Unlock()
return c.lastCodeErr
}
// explain turns a vendor error into an actionable one, folding in a failed
// goProxy authCode mint and the admin-restricted case.
func (c *safechatCipher) explain(op, corpID string, opErr error) error {
if c.client.IsBlocked(corpID) {
return fmt.Errorf("msgcrypto: %s blocked: the organization's key is restricted by its administrator: %w", op, opErr)
}
// A key fetch was needed but we had no usable code: that is the real
// cause, so report both.
if codeErr := c.lastAuthCodeErr(); codeErr != nil {
return fmt.Errorf("msgcrypto: %s failed and no usable auth code was available: %w (auth code error: %v)", op, opErr, codeErr)
}
if errors.Is(opErr, safechat.ErrMaxRetryExceeded) {
invalidateAuthCode(c.codes)
return fmt.Errorf("msgcrypto: %s failed: key material never became available: %w", op, opErr)
}
return fmt.Errorf("msgcrypto: %s failed: %w", op, opErr)
}
+161
View File
@@ -0,0 +1,161 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Keep this constraint in sync with cipher_safechat.go.
//go:build safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64))
package msgcrypto
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
)
// These tests exercise the real vendor backend, which means they link
// libsafechat.a and initialise the C library. They never reach the key server:
// an authCode is only spent when a key is actually fetched, and asserting that
// encryption fails without one is exactly the behaviour we want pinned.
func TestBackendIsReportedAvailable(t *testing.T) {
if !Available() {
t.Fatal("Available() = false in a safechat build")
}
if BackendVersion == "" {
t.Fatal("BackendVersion is empty in a safechat build")
}
if !strings.Contains(BackendVersion, "safechat") {
t.Fatalf("BackendVersion = %q, want it to name the vendor SDK", BackendVersion)
}
}
func TestOpenInitialisesCLibraryAndClosesCleanly(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
cipher, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v, want the C library to initialise", err)
}
if info, statErr := os.Stat(dir); statErr != nil || !info.IsDir() {
t.Fatalf("Open did not prepare the keystore dir: %v", statErr)
}
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
// The slot must be free again, otherwise a second Open in the same
// process would be refused forever.
second, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("second Open() after Close = %v, want success", err)
}
if err := second.Close(); err != nil {
t.Fatalf("second Close() = %v", err)
}
}
func TestOpenRefusesConcurrentSecondCipher(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
first, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer first.Close()
_, err = Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrAlreadyOpen) {
t.Fatalf("second Open() = %v, want ErrAlreadyOpen (the C library keeps global state)", err)
}
}
func TestOpenHonoursCancelledContext(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
_, err := Open(ctx, Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, context.Canceled) {
t.Fatalf("Open() with a cancelled context = %v, want context.Canceled", err)
}
}
func TestEncryptWithoutUsableKeyReportsAuthCodeCause(t *testing.T) {
// A cold keystore forces a key request. With no reachable key server the
// operation must fail with a message that names the auth code, rather
// than a bare vendor return code.
cipher, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: AuthCodeFunc(func(context.Context) (string, error) {
return "", errors.New("no code available in test")
}),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer cipher.Close()
_, err = cipher.EncryptMessage(context.Background(), "test-corp", "test-staff", []byte("hello"))
if err == nil {
t.Skip("the environment served a key without an auth code; nothing to assert")
}
if !strings.Contains(err.Error(), "auth code") {
t.Fatalf("EncryptMessage() = %v, want the error to name the auth code cause", err)
}
}
func TestCipherRejectsBadArgumentsBeforeCallingC(t *testing.T) {
cipher, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("placeholder-code"),
KeyServer: "https://key.example.test",
})
if err != nil {
t.Fatalf("Open() = %v", err)
}
defer cipher.Close()
if _, err := cipher.EncryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("EncryptMessage() with no corpID = %v, want ErrNoCorpID", err)
}
// The vendor SDK dereferences the first byte of the payload, so an empty
// slice must never reach it.
if _, err := cipher.DecryptMessage(context.Background(), "corp", "staff", nil); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("DecryptMessage() with no payload = %v, want ErrEmptyPayload", err)
}
}
+34
View File
@@ -0,0 +1,34 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// The constraint below is the exact negation of the one in cipher_safechat.go;
// change both together. This file covers every default DWS build: the release
// binaries are cross-compiled with CGO_ENABLED=0, and windows/arm64 has no
// vendor static library even when the tag is set.
//go:build !(safechat && cgo && (((darwin || linux) && (amd64 || arm64)) || (windows && amd64)))
package msgcrypto
import "context"
// BackendVersion is empty because no backend is compiled in.
const BackendVersion = ""
// Available reports that this binary has no SafeChat backend, so callers should
// not offer message encryption.
func Available() bool { return false }
// newBackend always fails here. Open checks Available first, so this exists to
// keep the package compiling and to fail safe if that check is ever bypassed.
func newBackend(context.Context, Config) (Cipher, error) { return nil, ErrUnavailable }
+54
View File
@@ -0,0 +1,54 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import "context"
// mintAuthCodeForProxy is the goProxy-only mint path. domain is compared
// locally and never forwarded. A successful mint invalidates any unconsumed
// cache so a one-shot code cannot be reused.
func mintAuthCodeForProxy(codes AuthCodeProvider, allowedHost, corpID, domain string) (string, error) {
if err := matchRedirectHost(domain, allowedHost); err != nil {
return "", err
}
if codes == nil {
return "", ErrNoAuthCodeProvider
}
var (
code string
err error
)
if provider, ok := codes.(CorpAuthCodeProvider); ok {
code, err = provider.AuthCodeForCorp(context.Background(), corpID)
} else {
code, err = codes.AuthCode(context.Background())
}
if err != nil {
invalidateAuthCode(codes)
return "", err
}
if code == "" {
invalidateAuthCode(codes)
return "", ErrNoAuthCode
}
invalidateAuthCode(codes)
return code, nil
}
func invalidateAuthCode(codes AuthCodeProvider) {
if invalidator, ok := codes.(interface{ Invalidate() }); ok {
invalidator.Invalidate()
}
}
@@ -0,0 +1,99 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync"
"testing"
"time"
)
type corpCountingProvider struct {
mu sync.Mutex
calls []string
code string
err error
}
func (p *corpCountingProvider) AuthCode(context.Context) (string, error) {
return p.AuthCodeForCorp(context.Background(), "")
}
func (p *corpCountingProvider) AuthCodeForCorp(_ context.Context, corpID string) (string, error) {
p.mu.Lock()
defer p.mu.Unlock()
p.calls = append(p.calls, corpID)
if p.err != nil {
return "", p.err
}
if p.code != "" {
return p.code, nil
}
return "code-for-" + corpID, nil
}
func TestMintAuthCodeForProxyUsesCorpProvider(t *testing.T) {
inner := &corpCountingProvider{}
code, err := mintAuthCodeForProxy(inner, "sso.anhei.test", "ding_corp", "https://sso.anhei.test/login")
if err != nil || code != "code-for-ding_corp" {
t.Fatalf("mint = %q, %v; want code-for-ding_corp, nil", code, err)
}
if len(inner.calls) != 1 || inner.calls[0] != "ding_corp" {
t.Fatalf("corpIDs = %v, want [ding_corp]", inner.calls)
}
}
func TestMintAuthCodeForProxyInvalidatesUnconsumedCache(t *testing.T) {
inner := &countingProvider{code: "once"}
cache := NewCachedAuthCode(inner, time.Hour)
if _, err := cache.AuthCode(context.Background()); err != nil {
t.Fatalf("seed cache: %v", err)
}
if got := inner.callCount(); got != 1 {
t.Fatalf("seed fetches = %d, want 1", got)
}
code, err := mintAuthCodeForProxy(cache, "sso.anhei.test", "ding_corp", "https://sso.anhei.test/login")
if err != nil || code != "once" {
t.Fatalf("mint = %q, %v; want once, nil", code, err)
}
// Cache was invalidated after spend; next mint hits upstream again.
if _, err := mintAuthCodeForProxy(cache, "sso.anhei.test", "ding_corp", "sso.anhei.test"); err != nil {
t.Fatalf("second mint: %v", err)
}
if got := inner.callCount(); got != 2 {
t.Fatalf("upstream calls = %d, want 2 (seed reused once, then refetch)", got)
}
}
func TestMintAuthCodeForProxyRejectsDomainMismatchWithoutFetching(t *testing.T) {
inner := &corpCountingProvider{code: "once"}
_, err := mintAuthCodeForProxy(inner, "sso.anhei.test", "ding_corp", "evil.example.test")
if !errors.Is(err, ErrRedirectHostMismatch) {
t.Fatalf("mint = %v, want ErrRedirectHostMismatch", err)
}
if got := len(inner.calls); got != 0 {
t.Fatalf("upstream called %d times on domain mismatch, want 0", got)
}
}
func TestMintAuthCodeForProxyRequiresProvider(t *testing.T) {
if _, err := mintAuthCodeForProxy(nil, "", "ding_corp", ""); !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("mint = %v, want ErrNoAuthCodeProvider", err)
}
}
+83
View File
@@ -0,0 +1,83 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"fmt"
"net"
"net/url"
"strings"
)
// validateKeyServer requires an HTTPS URL with a host so the vendor C
// library cannot pick the key-request destination.
func validateKeyServer(raw string) error {
raw = strings.TrimSpace(raw)
if raw == "" {
return ErrNoKeyServer
}
u, err := url.Parse(raw)
if err != nil || u.Host == "" || u.Scheme == "" {
return fmt.Errorf("%w: %q", ErrInvalidKeyServer, raw)
}
if !strings.EqualFold(u.Scheme, "https") {
return fmt.Errorf("%w: %q", ErrKeyServerNotHTTPS, raw)
}
if hostnameOf(raw) == "" {
return fmt.Errorf("%w: %q", ErrInvalidKeyServer, raw)
}
return nil
}
// matchRedirectHost compares the goProxy domain to AllowedRedirectHost.
// Both sides are reduced to a hostname. An empty domain or an empty
// allowed host skips the check; the domain is never sent to portal.
func matchRedirectHost(domain, allowed string) error {
domain = strings.TrimSpace(domain)
allowed = strings.TrimSpace(allowed)
if domain == "" || allowed == "" {
return nil
}
got := hostnameOf(domain)
want := hostnameOf(allowed)
if got == "" || want == "" || got != want {
return fmt.Errorf("%w: got %q, want %q", ErrRedirectHostMismatch, got, want)
}
return nil
}
// hostnameOf returns the lower-cased hostname of a URL, host:port, or bare
// host. Path, query, userinfo and port are ignored.
func hostnameOf(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if strings.Contains(raw, "://") {
u, err := url.Parse(raw)
if err == nil {
if host := strings.ToLower(u.Hostname()); host != "" {
return host
}
}
}
candidate := raw
if i := strings.IndexAny(candidate, "/?"); i >= 0 {
candidate = candidate[:i]
}
if host, _, err := net.SplitHostPort(candidate); err == nil {
return strings.ToLower(host)
}
return strings.ToLower(candidate)
}
+77
View File
@@ -0,0 +1,77 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"errors"
"testing"
)
func TestValidateKeyServerAcceptsHTTPS(t *testing.T) {
if err := validateKeyServer("https://key.example.test/v1"); err != nil {
t.Fatalf("validateKeyServer() = %v, want nil", err)
}
}
func TestValidateKeyServerRejectsEmpty(t *testing.T) {
if err := validateKeyServer(" "); !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("validateKeyServer() = %v, want ErrNoKeyServer", err)
}
}
func TestValidateKeyServerRejectsHTTP(t *testing.T) {
if err := validateKeyServer("http://key.example.test"); !errors.Is(err, ErrKeyServerNotHTTPS) {
t.Fatalf("validateKeyServer() = %v, want ErrKeyServerNotHTTPS", err)
}
}
func TestValidateKeyServerRejectsBareHost(t *testing.T) {
if err := validateKeyServer("key.example.test"); !errors.Is(err, ErrInvalidKeyServer) {
t.Fatalf("validateKeyServer() = %v, want ErrInvalidKeyServer", err)
}
}
func TestMatchRedirectHostComparesHostOnly(t *testing.T) {
if err := matchRedirectHost("https://sso.anhei.test:443/login", "https://sso.anhei.test/path"); err != nil {
t.Fatalf("matchRedirectHost() = %v, want nil", err)
}
if err := matchRedirectHost("sso.anhei.test", "https://sso.anhei.test"); err != nil {
t.Fatalf("bare host match = %v, want nil", err)
}
}
func TestMatchRedirectHostSkipsWhenEitherSideEmpty(t *testing.T) {
if err := matchRedirectHost("", "https://sso.anhei.test"); err != nil {
t.Fatalf("empty domain = %v, want nil", err)
}
if err := matchRedirectHost("sso.anhei.test", ""); err != nil {
t.Fatalf("empty allowed host = %v, want nil", err)
}
}
func TestMatchRedirectHostRejectsMismatch(t *testing.T) {
err := matchRedirectHost("evil.example.test", "https://sso.anhei.test")
if !errors.Is(err, ErrRedirectHostMismatch) {
t.Fatalf("matchRedirectHost() = %v, want ErrRedirectHostMismatch", err)
}
}
func TestHostnameOfStripsPortAndPath(t *testing.T) {
if got := hostnameOf("https://SSO.Example.TEST:8443/login?x=1"); got != "sso.example.test" {
t.Fatalf("hostnameOf(url) = %q", got)
}
if got := hostnameOf("SSO.Example.TEST:8443/login"); got != "sso.example.test" {
t.Fatalf("hostnameOf(hostport) = %q", got)
}
}
+96
View File
@@ -0,0 +1,96 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import "fmt"
// redactingLogger satisfies the vendor SDK's logger interface while making it
// impossible for the SDK to leak secrets into DWS output.
//
// This matters because the vendor SDK logs, at debug level, both the authCode
// it sends to the key server and the raw key-server response body, which
// carries key material. Rather than trying to enumerate and pattern-match
// every sensitive field, we drop every string argument and keep only its
// length. Numeric and boolean arguments pass through, so operators still get
// the useful diagnostics: HTTP status, payload sizes, timings and return
// codes. The format string is preserved so it stays clear which field was
// elided.
type redactingLogger struct {
logf func(format string, args ...any)
}
// newRedactingLogger returns a logger that forwards to logf, or nil when logf
// is nil so the SDK skips logging entirely.
func newRedactingLogger(logf func(format string, args ...any)) *redactingLogger {
if logf == nil {
return nil
}
return &redactingLogger{logf: logf}
}
// Debug forwards a redacted debug line.
func (l *redactingLogger) Debug(msg string, args ...interface{}) { l.emit("debug", msg, args) }
// Info forwards a redacted info line.
func (l *redactingLogger) Info(msg string, args ...interface{}) { l.emit("info", msg, args) }
// Error forwards a redacted error line.
func (l *redactingLogger) Error(msg string, args ...interface{}) { l.emit("error", msg, args) }
// emit rewrites args so no string value survives, then forwards the line.
func (l *redactingLogger) emit(level, msg string, args []interface{}) {
if l == nil || l.logf == nil {
return
}
l.logf("safechat[%s] "+msg, append([]any{level}, redactArgs(args)...)...)
}
// redactArgs replaces every string-like argument with a length marker and
// leaves other kinds intact.
func redactArgs(args []interface{}) []any {
out := make([]any, 0, len(args))
for _, arg := range args {
out = append(out, redactArg(arg))
}
return out
}
// redactArg elides a single argument's contents when it could carry a secret.
// Strings and byte slices are reduced to their length; errors are reduced to
// their type so a wrapped body preview cannot slip through; everything else
// (numbers, booleans, durations) is kept because it cannot carry key material.
func redactArg(arg any) any {
switch v := arg.(type) {
case string:
return redactedValue(len(v))
case []byte:
return redactedValue(len(v))
case fmt.Stringer:
// time.Duration and friends are Stringers, but so are opaque types
// that may embed a payload. Keep durations, elide the rest.
if _, ok := arg.(interface{ Nanoseconds() int64 }); ok {
return v.String()
}
return redactedValue(len(v.String()))
case error:
return fmt.Sprintf("<%T redacted>", v)
default:
return v
}
}
// redactedValue renders the placeholder used in place of elided content.
func redactedValue(n int) string {
return fmt.Sprintf("<redacted len=%d>", n)
}
+157
View File
@@ -0,0 +1,157 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"errors"
"fmt"
"strings"
"testing"
"time"
)
// captureLogf collects emitted lines for assertions.
func captureLogf(lines *[]string) func(string, ...any) {
return func(format string, args ...any) {
*lines = append(*lines, fmt.Sprintf(format, args...))
}
}
func TestNewRedactingLoggerReturnsNilWhenSinkIsNil(t *testing.T) {
// A nil logger makes the vendor SDK skip logging entirely, which is the
// safe default because it logs the authCode at debug level.
if got := newRedactingLogger(nil); got != nil {
t.Fatalf("newRedactingLogger(nil) = %v, want nil", got)
}
}
func TestRedactingLoggerHidesAuthCode(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
// This mirrors the vendor SDK's own debug line, which prints the code.
const secret = "abc123authcode"
logger.Debug("Code (auth_token, length=%d): %s", len(secret), secret)
if len(lines) != 1 {
t.Fatalf("got %d lines, want 1", len(lines))
}
if strings.Contains(lines[0], secret) {
t.Fatalf("log line leaked the auth code: %q", lines[0])
}
if !strings.Contains(lines[0], "redacted") {
t.Fatalf("log line = %q, want a redaction marker", lines[0])
}
}
func TestRedactingLoggerHidesKeyServerResponseBody(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
body := `{"key":"BASE64KEYMATERIAL==","keyVersion":3}`
logger.Debug("Body (length=%d): %s", len(body), body)
if strings.Contains(lines[0], "BASE64KEYMATERIAL") {
t.Fatalf("log line leaked key material: %q", lines[0])
}
}
func TestRedactingLoggerHidesByteSlices(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Info("payload=%s", []byte("plaintext-message"))
if strings.Contains(lines[0], "plaintext-message") {
t.Fatalf("log line leaked a byte payload: %q", lines[0])
}
}
func TestRedactingLoggerHidesErrorText(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
// Vendor error strings can embed a response body preview.
logger.Error("request failed: %v", errors.New(`server said {"key":"LEAKED"}`))
if strings.Contains(lines[0], "LEAKED") {
t.Fatalf("log line leaked error contents: %q", lines[0])
}
}
func TestRedactingLoggerKeepsDiagnosticNumbers(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Debug("Status: %d (took %s), size=%d", 503, 1500*time.Millisecond, 4096)
line := lines[0]
for _, want := range []string{"503", "1.5s", "4096"} {
if !strings.Contains(line, want) {
t.Fatalf("log line = %q, want it to keep %q for diagnostics", line, want)
}
}
}
func TestRedactingLoggerLabelsLevel(t *testing.T) {
var lines []string
logger := newRedactingLogger(captureLogf(&lines))
logger.Debug("d")
logger.Info("i")
logger.Error("e")
if len(lines) != 3 {
t.Fatalf("got %d lines, want 3", len(lines))
}
for i, want := range []string{"debug", "info", "error"} {
if !strings.Contains(lines[i], want) {
t.Fatalf("line %d = %q, want level %q", i, lines[i], want)
}
}
}
func TestRedactingLoggerToleratesNilSinkAtEmit(t *testing.T) {
// Guard against a partially constructed logger being used.
var logger *redactingLogger
logger.Debug("must not panic %s", "value")
}
func TestRedactArgKeepsDurations(t *testing.T) {
if got := redactArg(2 * time.Second); got != "2s" {
t.Fatalf("redactArg(2s) = %v, want 2s", got)
}
}
func TestRedactArgElidesStrings(t *testing.T) {
got, ok := redactArg("secret").(string)
if !ok {
t.Fatalf("redactArg returned %T, want string", got)
}
if strings.Contains(got, "secret") {
t.Fatalf("redactArg = %q, want the value elided", got)
}
if !strings.Contains(got, "len=6") {
t.Fatalf("redactArg = %q, want the length preserved", got)
}
}
func TestRedactArgPassesThroughNumbers(t *testing.T) {
if got := redactArg(42); got != 42 {
t.Fatalf("redactArg(42) = %v, want 42", got)
}
if got := redactArg(true); got != true {
t.Fatalf("redactArg(true) = %v, want true", got)
}
}
+297
View File
@@ -0,0 +1,297 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Package msgcrypto wraps the third-party SafeChat SDK so DWS can decrypt
// DingTalk messages that an organization has encrypted with its own key
// material, and encrypt outbound ones.
//
// The SafeChat backend links a prebuilt C static library and therefore needs
// CGO. Because DWS ships CGO-free cross-compiled release binaries, the backend
// is compiled only under the "safechat" build tag:
//
// CGO_ENABLED=1 go build -tags safechat ./cmd
//
// Every other build gets a stub whose constructor fails with ErrUnavailable,
// so callers must always handle that error rather than assume the capability
// exists. Use Available to branch before offering the feature to a user.
//
// Key material is fetched from the vendor key server on demand, which requires
// a DingTalk 免登 authCode supplied through AuthCodeProvider. DWS does not mint
// that code itself; the caller injects a provider.
package msgcrypto
import (
"context"
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// Errors reported by this package. Callers are expected to test for
// ErrUnavailable explicitly, because it is the normal outcome on every build
// that does not enable the safechat tag.
var (
// ErrUnavailable means this binary was built without the SafeChat
// backend, or for a platform the vendor does not ship a static library
// for (notably windows/arm64).
ErrUnavailable = errors.New("msgcrypto: SafeChat backend not built into this binary")
// ErrAlreadyOpen means a Cipher is already open. The underlying C
// library keeps global state, so only one may exist per process.
ErrAlreadyOpen = errors.New("msgcrypto: a cipher is already open in this process")
// ErrClosed is returned by a Cipher whose Close has already run.
ErrClosed = errors.New("msgcrypto: cipher is closed")
// ErrNoAuthCodeProvider means Config.AuthCode was nil. Without it the
// backend cannot fetch or rotate key material.
ErrNoAuthCodeProvider = errors.New("msgcrypto: config.AuthCode is required")
// ErrEmptyPayload means an encrypt or decrypt call got no bytes. The
// vendor SDK rejects empty input, so we reject it earlier with a
// clearer message.
ErrEmptyPayload = errors.New("msgcrypto: payload is empty")
// ErrNoCorpID means the caller omitted the organization id, which
// selects the key and therefore cannot be defaulted.
ErrNoCorpID = errors.New("msgcrypto: corpID is required")
// ErrNoKeyServer means Config.KeyServer was empty. The vendor C
// library would otherwise pick the key-request destination.
ErrNoKeyServer = errors.New("msgcrypto: config.KeyServer is required")
// ErrInvalidKeyServer means Config.KeyServer is not a usable URL.
ErrInvalidKeyServer = errors.New("msgcrypto: config.KeyServer is not a valid URL")
// ErrKeyServerNotHTTPS means Config.KeyServer is not https.
ErrKeyServerNotHTTPS = errors.New("msgcrypto: config.KeyServer must be an https URL")
// ErrRedirectHostMismatch means the domain goProxy received does not
// match Config.AllowedRedirectHost. The domain is never sent to portal.
ErrRedirectHostMismatch = errors.New("msgcrypto: goProxy domain host does not match AllowedRedirectHost")
)
// keystoreDirPerm keeps the key cache owner-only. The directory holds
// organization key material, so it must not be group- or world-readable.
const keystoreDirPerm fs.FileMode = 0o700
// DefaultKeystoreDir returns the default key cache directory,
// ~/.dws/safechat/keystore, honouring DWS_CONFIG_DIR like the rest of DWS.
func DefaultKeystoreDir() string {
return filepath.Join(config.DefaultConfigDir(), "safechat", "keystore")
}
// Cipher encrypts and decrypts message payloads for one organization at a
// time. Implementations are safe for concurrent use.
type Cipher interface {
// EncryptMessage encrypts plaintext for corpID/staffID and returns the
// vendor ciphertext envelope.
EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error)
// DecryptMessage decrypts a vendor ciphertext envelope.
DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error)
// Close releases the backend and frees the process-wide slot so a later
// Open can succeed. Calling it twice is safe.
Close() error
}
// Config parameterises Open.
type Config struct {
// KeystoreDir is where fetched keys are cached. Defaults to
// DefaultKeystoreDir. It is created with 0700 if missing.
KeystoreDir string
// UserID is an opaque local identifier. The vendor SDK stores it but
// does not use it for key selection; leave empty to let the SDK
// generate one.
UserID string
// AuthCode supplies the DingTalk 免登 authCode used to authenticate key
// requests. Required. The backend calls it only from the vendor goProxy
// callback (cold keystore or key-version rotation), never on every
// encrypt/decrypt.
AuthCode AuthCodeProvider
// KeyServer is the HTTPS URL of the vendor key service. Required: it
// replaces the host the closed-source C library would otherwise pick.
KeyServer string
// AllowedRedirectHost, when set, is compared to the domain the C
// library passes into goProxy. A mismatch fails that key fetch. It is
// a local check only; the domain is never sent to portal.
AllowedRedirectHost string
// MaxRetry bounds retries while a key is still being fetched. Zero
// selects the vendor default.
MaxRetry int
// HTTPTimeout bounds each key request. Zero selects the vendor default.
HTTPTimeout time.Duration
// Debug enables backend logging through a redacting logger. Off by
// default because the vendor SDK logs the authCode and raw key-server
// responses at debug level.
Debug bool
// Logf receives already-redacted backend log lines when Debug is set.
// Nil discards them.
Logf func(format string, args ...any)
}
// withDefaults returns cfg with empty optional fields filled in.
func (cfg Config) withDefaults() Config {
if cfg.KeystoreDir == "" {
cfg.KeystoreDir = DefaultKeystoreDir()
}
return cfg
}
// validate reports whether cfg carries everything the backend needs.
func (cfg Config) validate() error {
if cfg.AuthCode == nil {
return ErrNoAuthCodeProvider
}
if cfg.KeystoreDir == "" {
return errors.New("msgcrypto: config.KeystoreDir resolved to an empty path")
}
if err := validateKeyServer(cfg.KeyServer); err != nil {
return err
}
return nil
}
// prepareKeystore creates dir if needed and makes sure it is owner-only.
// An existing directory with looser bits is tightened, because it caches key
// material.
func prepareKeystore(dir string) error {
if err := os.MkdirAll(dir, keystoreDirPerm); err != nil {
return fmt.Errorf("msgcrypto: create keystore dir: %w", err)
}
info, err := os.Stat(dir)
if err != nil {
return fmt.Errorf("msgcrypto: stat keystore dir: %w", err)
}
if !info.IsDir() {
return fmt.Errorf("msgcrypto: keystore path %s is not a directory", dir)
}
// Windows does not model POSIX bits, so only tighten where they apply.
if runtimeSupportsPOSIXPerm && info.Mode().Perm() != keystoreDirPerm {
if err := os.Chmod(dir, keystoreDirPerm); err != nil {
return fmt.Errorf("msgcrypto: restrict keystore dir permissions: %w", err)
}
}
return nil
}
// process holds the single-instance guard. The vendor C library keeps global
// state, so a second concurrent Cipher would corrupt it.
var process struct {
mu sync.Mutex
open bool
}
// Open validates cfg, prepares the keystore and starts the backend.
//
// It returns ErrUnavailable when the backend was not compiled in, so callers
// can degrade gracefully. Only one Cipher may be open per process; Close frees
// the slot.
func Open(ctx context.Context, cfg Config) (Cipher, error) {
cfg = cfg.withDefaults()
if err := cfg.validate(); err != nil {
return nil, err
}
if !Available() {
return nil, ErrUnavailable
}
if err := prepareKeystore(cfg.KeystoreDir); err != nil {
return nil, err
}
process.mu.Lock()
defer process.mu.Unlock()
if process.open {
return nil, ErrAlreadyOpen
}
backend, err := newBackend(ctx, cfg)
if err != nil {
return nil, err
}
process.open = true
return &trackedCipher{backend: backend}, nil
}
// trackedCipher releases the process-wide slot when the wrapped backend closes.
type trackedCipher struct {
mu sync.Mutex
backend Cipher
}
// EncryptMessage validates the payload and delegates to the backend.
func (c *trackedCipher) EncryptMessage(ctx context.Context, corpID, staffID string, plaintext []byte) ([]byte, error) {
backend, err := c.live(corpID, plaintext)
if err != nil {
return nil, err
}
return backend.EncryptMessage(ctx, corpID, staffID, plaintext)
}
// DecryptMessage validates the payload and delegates to the backend.
func (c *trackedCipher) DecryptMessage(ctx context.Context, corpID, staffID string, ciphertext []byte) ([]byte, error) {
backend, err := c.live(corpID, ciphertext)
if err != nil {
return nil, err
}
return backend.DecryptMessage(ctx, corpID, staffID, ciphertext)
}
// live returns the backend after checking the cipher is open and the arguments
// are usable.
func (c *trackedCipher) live(corpID string, payload []byte) (Cipher, error) {
if corpID == "" {
return nil, ErrNoCorpID
}
if len(payload) == 0 {
return nil, ErrEmptyPayload
}
c.mu.Lock()
defer c.mu.Unlock()
if c.backend == nil {
return nil, ErrClosed
}
return c.backend, nil
}
// Close closes the backend once and releases the process-wide slot.
func (c *trackedCipher) Close() error {
c.mu.Lock()
backend := c.backend
c.backend = nil
c.mu.Unlock()
if backend == nil {
return nil
}
err := backend.Close()
process.mu.Lock()
process.open = false
process.mu.Unlock()
return err
}
+345
View File
@@ -0,0 +1,345 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"testing"
)
// fakeCipher stands in for a backend so the wrapper logic can be tested on
// every platform, with or without the safechat tag.
type fakeCipher struct {
mu sync.Mutex
encrypted [][]byte
decrypted [][]byte
closeCount int
closeErr error
}
func (f *fakeCipher) EncryptMessage(_ context.Context, _, _ string, plaintext []byte) ([]byte, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.encrypted = append(f.encrypted, plaintext)
return []byte("cipher:" + string(plaintext)), nil
}
func (f *fakeCipher) DecryptMessage(_ context.Context, _, _ string, ciphertext []byte) ([]byte, error) {
f.mu.Lock()
defer f.mu.Unlock()
f.decrypted = append(f.decrypted, ciphertext)
return []byte(strings.TrimPrefix(string(ciphertext), "cipher:")), nil
}
func (f *fakeCipher) Close() error {
f.mu.Lock()
defer f.mu.Unlock()
f.closeCount++
return f.closeErr
}
// newTrackedForTest wraps backend and claims the process slot the same way
// Open does, so slot release can be asserted.
func newTrackedForTest(t *testing.T, backend Cipher) *trackedCipher {
t.Helper()
process.mu.Lock()
process.open = true
process.mu.Unlock()
t.Cleanup(func() {
process.mu.Lock()
process.open = false
process.mu.Unlock()
})
return &trackedCipher{backend: backend}
}
func TestConfigWithDefaultsFillsKeystoreDir(t *testing.T) {
cfg := Config{}.withDefaults()
if cfg.KeystoreDir == "" {
t.Fatal("withDefaults left KeystoreDir empty")
}
if want := DefaultKeystoreDir(); cfg.KeystoreDir != want {
t.Fatalf("KeystoreDir = %q, want %q", cfg.KeystoreDir, want)
}
}
func TestConfigWithDefaultsKeepsExplicitKeystoreDir(t *testing.T) {
cfg := Config{KeystoreDir: "/custom/keys"}.withDefaults()
if cfg.KeystoreDir != "/custom/keys" {
t.Fatalf("KeystoreDir = %q, want /custom/keys", cfg.KeystoreDir)
}
}
func TestDefaultKeystoreDirHonoursConfigDirOverride(t *testing.T) {
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
got := DefaultKeystoreDir()
if !strings.HasSuffix(got, filepath.Join("safechat", "keystore")) {
t.Fatalf("DefaultKeystoreDir() = %q, want it to end with safechat/keystore", got)
}
if !strings.HasPrefix(got, os.Getenv("DWS_CONFIG_DIR")) {
t.Fatalf("DefaultKeystoreDir() = %q, want it under DWS_CONFIG_DIR", got)
}
}
func TestConfigValidateRequiresAuthCodeProvider(t *testing.T) {
err := Config{KeystoreDir: "/tmp/keys"}.validate()
if !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("validate() = %v, want ErrNoAuthCodeProvider", err)
}
}
func TestConfigValidateAcceptsCompleteConfig(t *testing.T) {
cfg := Config{
KeystoreDir: "/tmp/keys",
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
}
if err := cfg.validate(); err != nil {
t.Fatalf("validate() = %v, want nil", err)
}
}
func TestConfigValidateRequiresKeyServer(t *testing.T) {
err := Config{KeystoreDir: "/tmp/keys", AuthCode: StaticAuthCode("code")}.validate()
if !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("validate() = %v, want ErrNoKeyServer", err)
}
}
func TestConfigValidateRejectsHTTPKeyServer(t *testing.T) {
err := Config{
KeystoreDir: "/tmp/keys",
AuthCode: StaticAuthCode("code"),
KeyServer: "http://key.example.test",
}.validate()
if !errors.Is(err, ErrKeyServerNotHTTPS) {
t.Fatalf("validate() = %v, want ErrKeyServerNotHTTPS", err)
}
}
func TestPrepareKeystoreCreatesOwnerOnlyDir(t *testing.T) {
dir := filepath.Join(t.TempDir(), "nested", "keystore")
if err := prepareKeystore(dir); err != nil {
t.Fatalf("prepareKeystore() = %v", err)
}
info, err := os.Stat(dir)
if err != nil {
t.Fatalf("stat: %v", err)
}
if !info.IsDir() {
t.Fatal("prepareKeystore did not create a directory")
}
if runtime.GOOS == "windows" {
return
}
if got := info.Mode().Perm(); got != keystoreDirPerm {
t.Fatalf("perm = %#o, want %#o", got, keystoreDirPerm)
}
}
func TestPrepareKeystoreTightensLoosePermissions(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("POSIX permission bits are not modelled on Windows")
}
dir := filepath.Join(t.TempDir(), "keystore")
if err := os.MkdirAll(dir, fs.FileMode(0o755)); err != nil {
t.Fatalf("mkdir: %v", err)
}
if err := prepareKeystore(dir); err != nil {
t.Fatalf("prepareKeystore() = %v", err)
}
info, err := os.Stat(dir)
if err != nil {
t.Fatalf("stat: %v", err)
}
if got := info.Mode().Perm(); got != keystoreDirPerm {
t.Fatalf("perm = %#o, want %#o (key material must stay owner-only)", got, keystoreDirPerm)
}
}
func TestPrepareKeystoreRejectsFilePath(t *testing.T) {
path := filepath.Join(t.TempDir(), "keystore")
if err := os.WriteFile(path, []byte("not a dir"), 0o600); err != nil {
t.Fatalf("write: %v", err)
}
err := prepareKeystore(path)
if err == nil {
t.Fatal("prepareKeystore() = nil, want an error for a non-directory path")
}
}
func TestOpenRejectsMissingAuthCodeProviderBeforeTouchingDisk(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
_, err := Open(context.Background(), Config{KeystoreDir: dir})
if !errors.Is(err, ErrNoAuthCodeProvider) {
t.Fatalf("Open() = %v, want ErrNoAuthCodeProvider", err)
}
if _, statErr := os.Stat(dir); !os.IsNotExist(statErr) {
t.Fatal("Open created the keystore dir despite an invalid config")
}
}
func TestOpenRejectsMissingKeyServerBeforeTouchingDisk(t *testing.T) {
dir := filepath.Join(t.TempDir(), "keystore")
_, err := Open(context.Background(), Config{
KeystoreDir: dir,
AuthCode: StaticAuthCode("code"),
})
if !errors.Is(err, ErrNoKeyServer) {
t.Fatalf("Open() = %v, want ErrNoKeyServer", err)
}
if _, statErr := os.Stat(dir); !os.IsNotExist(statErr) {
t.Fatal("Open created the keystore dir despite a missing KeyServer")
}
}
func TestOpenWithoutBackendReportsUnavailable(t *testing.T) {
if Available() {
t.Skip("this binary has the safechat backend compiled in")
}
_, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrUnavailable) {
t.Fatalf("Open() = %v, want ErrUnavailable", err)
}
}
func TestAvailableAgreesWithBackendVersion(t *testing.T) {
if Available() != (BackendVersion != "") {
t.Fatalf("Available() = %v but BackendVersion = %q; they must agree", Available(), BackendVersion)
}
}
func TestTrackedCipherRoundTripsThroughBackend(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
ciphertext, err := cipher.EncryptMessage(context.Background(), "corp", "staff", []byte("hello"))
if err != nil {
t.Fatalf("EncryptMessage() = %v", err)
}
plaintext, err := cipher.DecryptMessage(context.Background(), "corp", "staff", ciphertext)
if err != nil {
t.Fatalf("DecryptMessage() = %v", err)
}
if string(plaintext) != "hello" {
t.Fatalf("round trip = %q, want hello", plaintext)
}
}
func TestTrackedCipherRejectsEmptyCorpID(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if _, err := cipher.EncryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("EncryptMessage() = %v, want ErrNoCorpID", err)
}
if _, err := cipher.DecryptMessage(context.Background(), "", "staff", []byte("x")); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("DecryptMessage() = %v, want ErrNoCorpID", err)
}
}
func TestTrackedCipherRejectsEmptyPayload(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if _, err := cipher.EncryptMessage(context.Background(), "corp", "staff", nil); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("EncryptMessage() = %v, want ErrEmptyPayload", err)
}
if _, err := cipher.DecryptMessage(context.Background(), "corp", "staff", []byte{}); !errors.Is(err, ErrEmptyPayload) {
t.Fatalf("DecryptMessage() = %v, want ErrEmptyPayload", err)
}
}
func TestTrackedCipherRejectsUseAfterClose(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
if _, err := cipher.EncryptMessage(context.Background(), "corp", "staff", []byte("x")); !errors.Is(err, ErrClosed) {
t.Fatalf("EncryptMessage() after Close = %v, want ErrClosed", err)
}
}
func TestTrackedCipherCloseIsIdempotentAndClosesBackendOnce(t *testing.T) {
backend := &fakeCipher{}
cipher := newTrackedForTest(t, backend)
for i := 0; i < 3; i++ {
if err := cipher.Close(); err != nil {
t.Fatalf("Close() #%d = %v", i+1, err)
}
}
if backend.closeCount != 1 {
t.Fatalf("backend Close called %d times, want exactly 1", backend.closeCount)
}
}
func TestTrackedCipherCloseReleasesProcessSlot(t *testing.T) {
cipher := newTrackedForTest(t, &fakeCipher{})
if err := cipher.Close(); err != nil {
t.Fatalf("Close() = %v", err)
}
process.mu.Lock()
open := process.open
process.mu.Unlock()
if open {
t.Fatal("Close did not release the process slot, so a later Open would fail")
}
}
func TestTrackedCipherClosePropagatesBackendError(t *testing.T) {
wantErr := errors.New("backend close failed")
cipher := newTrackedForTest(t, &fakeCipher{closeErr: wantErr})
if err := cipher.Close(); !errors.Is(err, wantErr) {
t.Fatalf("Close() = %v, want %v", err, wantErr)
}
process.mu.Lock()
open := process.open
process.mu.Unlock()
if open {
t.Fatal("a failing backend Close must still release the process slot")
}
}
func TestOpenRefusesSecondCipherWhileOneIsOpen(t *testing.T) {
// The vendor C library keeps global state, so a second concurrent
// cipher must be refused rather than silently corrupt it.
process.mu.Lock()
process.open = true
process.mu.Unlock()
t.Cleanup(func() {
process.mu.Lock()
process.open = false
process.mu.Unlock()
})
if !Available() {
t.Skip("Open reports ErrUnavailable before reaching the single-instance guard")
}
_, err := Open(context.Background(), Config{
KeystoreDir: filepath.Join(t.TempDir(), "keystore"),
AuthCode: StaticAuthCode("code"),
KeyServer: "https://key.example.test",
})
if !errors.Is(err, ErrAlreadyOpen) {
t.Fatalf("Open() = %v, want ErrAlreadyOpen", err)
}
}
+20
View File
@@ -0,0 +1,20 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//go:build !windows
package msgcrypto
// runtimeSupportsPOSIXPerm reports that the keystore directory's permission
// bits are meaningful here, so prepareKeystore enforces owner-only access.
const runtimeSupportsPOSIXPerm = true
+21
View File
@@ -0,0 +1,21 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//go:build windows
package msgcrypto
// runtimeSupportsPOSIXPerm reports that Windows does not model POSIX
// permission bits, so prepareKeystore leaves the directory mode alone and
// relies on the user profile ACL instead.
const runtimeSupportsPOSIXPerm = false
+170
View File
@@ -0,0 +1,170 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"fmt"
"net/http"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
)
// VendorSafeChat is the first vendorAuthCode vendor.
const VendorSafeChat = "safechat"
// PortalAuthCode mints a one-shot 免登 authCode from portal
// POST /oauth2/vendorAuthCode. It does not cache the code: goProxy spends it
// immediately. Do not wrap this in CachedAuthCode.
type PortalAuthCode struct {
ConfigDir string
Vendor string
CLIVersion string
HTTPClient *http.Client
clientID func() string
snapshot func(ctx context.Context, configDir string) (*auth.TokenData, error)
refresh func(ctx context.Context, configDir, rejected string) (string, error)
fetch func(ctx context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error)
}
// NewPortalAuthCode returns a provider that talks to portal with the current
// login. cliVersion is sent as x-dws-cli-version; leave empty only when the
// caller cannot know the CLI version.
func NewPortalAuthCode(configDir, cliVersion string) *PortalAuthCode {
return &PortalAuthCode{
ConfigDir: configDir,
Vendor: VendorSafeChat,
CLIVersion: cliVersion,
clientID: auth.ClientID,
snapshot: func(ctx context.Context, dir string) (*auth.TokenData, error) {
return auth.NewOAuthProvider(dir, nil).GetTokenSnapshot(ctx)
},
refresh: func(ctx context.Context, dir, rejected string) (string, error) {
return auth.NewOAuthProvider(dir, nil).ForceRefreshRejectedToken(ctx, rejected)
},
fetch: auth.FetchVendorAuthCode,
}
}
// AuthCode mints a code for the logged-in organization. Prefer AuthCodeForCorp
// when goProxy already has a corpID.
func (p *PortalAuthCode) AuthCode(ctx context.Context) (string, error) {
snap, err := p.loadSnapshot(ctx)
if err != nil {
return "", err
}
return p.AuthCodeForCorp(ctx, snap.CorpID)
}
// AuthCodeForCorp mints a one-shot code for corpID. The request body is only
// vendor + corpId.
func (p *PortalAuthCode) AuthCodeForCorp(ctx context.Context, corpID string) (string, error) {
corpID = strings.TrimSpace(corpID)
if corpID == "" {
return "", ErrNoCorpID
}
snap, err := p.loadSnapshot(ctx)
if err != nil {
return "", err
}
token := strings.TrimSpace(snap.AccessToken)
if token == "" {
return "", fmt.Errorf("msgcrypto: access token is empty")
}
result, err := p.fetchOnce(ctx, snap, token, corpID)
if err == nil {
return result.AuthCode, nil
}
var verr *auth.VendorAuthCodeError
if errors.As(err, &verr) && verr.Code == auth.VendorAuthCodeTokenInvalid {
refreshed, rerr := p.doRefresh(ctx, token)
if rerr == nil && refreshed != "" && refreshed != token {
result, err = p.fetchOnce(ctx, snap, refreshed, corpID)
if err == nil {
return result.AuthCode, nil
}
}
return "", err
}
if retryableVendorAuthCode(err) {
result, err = p.fetchOnce(ctx, snap, token, corpID)
if err == nil {
return result.AuthCode, nil
}
}
return "", err
}
func (p *PortalAuthCode) loadSnapshot(ctx context.Context) (*auth.TokenData, error) {
if p.snapshot == nil {
return nil, fmt.Errorf("msgcrypto: portal authCode snapshot loader is not configured")
}
snap, err := p.snapshot(ctx, p.ConfigDir)
if err != nil {
return nil, fmt.Errorf("msgcrypto: load access token: %w", err)
}
if snap == nil {
return nil, fmt.Errorf("msgcrypto: load access token: empty snapshot")
}
return snap, nil
}
func (p *PortalAuthCode) fetchOnce(ctx context.Context, snap *auth.TokenData, token, corpID string) (*auth.VendorAuthCodeResult, error) {
fetch := p.fetch
if fetch == nil {
fetch = auth.FetchVendorAuthCode
}
vendor := strings.TrimSpace(p.Vendor)
if vendor == "" {
vendor = VendorSafeChat
}
clientID := strings.TrimSpace(snap.ClientID)
if clientID == "" && p.clientID != nil {
clientID = strings.TrimSpace(p.clientID())
}
return fetch(ctx, auth.VendorAuthCodeInput{
AccessToken: token,
ClientID: clientID,
CLIVersion: p.CLIVersion,
LoginRegion: auth.LoginRegion(strings.TrimSpace(snap.LoginRegion)),
Vendor: vendor,
CorpID: corpID,
HTTPClient: p.HTTPClient,
})
}
func (p *PortalAuthCode) doRefresh(ctx context.Context, rejected string) (string, error) {
if p.refresh == nil {
return "", fmt.Errorf("msgcrypto: token refresh is not configured")
}
return p.refresh(ctx, p.ConfigDir, rejected)
}
func retryableVendorAuthCode(err error) bool {
var verr *auth.VendorAuthCodeError
if errors.As(err, &verr) {
return verr.Retryable()
}
var statusErr *auth.HTTPStatusError
if errors.As(err, &statusErr) && statusErr != nil {
return statusErr.StatusCode == http.StatusTooManyRequests ||
statusErr.StatusCode >= http.StatusInternalServerError
}
return false
}
+121
View File
@@ -0,0 +1,121 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package msgcrypto
import (
"context"
"errors"
"sync/atomic"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
)
func TestPortalAuthCodePostsVendorAndCorpOnly(t *testing.T) {
var got auth.VendorAuthCodeInput
p := &PortalAuthCode{
ConfigDir: t.TempDir(),
Vendor: VendorSafeChat,
CLIVersion: "1.2.3",
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{
AccessToken: "user-token",
ClientID: "dws-client",
CorpID: "ding_login",
LoginRegion: "",
}, nil
},
fetch: func(_ context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
got = in
return &auth.VendorAuthCodeResult{AuthCode: "tmp-code", ExpiresIn: 120}, nil
},
}
code, err := p.AuthCodeForCorp(context.Background(), "ding_target")
if err != nil || code != "tmp-code" {
t.Fatalf("AuthCodeForCorp() = %q, %v", code, err)
}
if got.Vendor != VendorSafeChat || got.CorpID != "ding_target" {
t.Fatalf("body vendor/corpId = %q/%q", got.Vendor, got.CorpID)
}
if got.AccessToken != "user-token" || got.ClientID != "dws-client" || got.CLIVersion != "1.2.3" {
t.Fatalf("headers token/client/version = %q/%q/%q", got.AccessToken, got.ClientID, got.CLIVersion)
}
}
func TestPortalAuthCodeRetriesOnceOnTokenInvalidAfterRefresh(t *testing.T) {
var calls atomic.Int32
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "stale", ClientID: "dws-client", CorpID: "ding"}, nil
},
refresh: func(context.Context, string, string) (string, error) {
return "fresh", nil
},
fetch: func(_ context.Context, in auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
n := calls.Add(1)
if in.AccessToken == "stale" {
return nil, &auth.VendorAuthCodeError{Code: auth.VendorAuthCodeTokenInvalid, Message: "expired"}
}
if in.AccessToken != "fresh" {
t.Fatalf("retry token = %q, want fresh", in.AccessToken)
}
if n != 2 {
t.Fatalf("fetch calls = %d, want 2", n)
}
return &auth.VendorAuthCodeResult{AuthCode: "new-code", ExpiresIn: 120}, nil
},
}
code, err := p.AuthCodeForCorp(context.Background(), "ding")
if err != nil || code != "new-code" {
t.Fatalf("AuthCodeForCorp() = %q, %v", code, err)
}
if calls.Load() != 2 {
t.Fatalf("fetch called %d times, want 2", calls.Load())
}
}
func TestPortalAuthCodeDoesNotRetryOrgMismatch(t *testing.T) {
var calls atomic.Int32
want := &auth.VendorAuthCodeError{Code: auth.VendorAuthCodeOrgMismatch, Message: "mismatch"}
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "tok", ClientID: "dws-client", CorpID: "ding"}, nil
},
fetch: func(context.Context, auth.VendorAuthCodeInput) (*auth.VendorAuthCodeResult, error) {
calls.Add(1)
return nil, want
},
}
_, err := p.AuthCodeForCorp(context.Background(), "other")
if !errors.Is(err, want) {
t.Fatalf("AuthCodeForCorp() = %v, want %v", err, want)
}
if calls.Load() != 1 {
t.Fatalf("fetch called %d times, want 1", calls.Load())
}
}
func TestPortalAuthCodeRequiresCorpID(t *testing.T) {
p := &PortalAuthCode{
snapshot: func(context.Context, string) (*auth.TokenData, error) {
return &auth.TokenData{AccessToken: "tok", ClientID: "id"}, nil
},
}
if _, err := p.AuthCodeForCorp(context.Background(), " "); !errors.Is(err, ErrNoCorpID) {
t.Fatalf("AuthCodeForCorp() = %v, want ErrNoCorpID", err)
}
}
-3
View File
@@ -435,7 +435,6 @@ const (
exitCodeInternal = 5
exitCodeDiscovery = 6
exitCodePartial = 7 // partial_failure 专用(契约 §4;规划 WS2 第4项;B142 将在 errors 侧补同源常量)
exitCodeWait = 8 // wait 终态失败专用(--wait 观察到失败终态;与 partial 同为"仅新增专用码")
)
// subtypeConfirmationRequired 是门禁拦截的 failure 子类标记(契约规范 §2.4),
@@ -492,8 +491,6 @@ func exitCodeForErrorInfo(info *ErrorInfo) int {
return exitCodePermission
case "discovery":
return exitCodeDiscovery
case "wait":
return exitCodeWait
default:
return exitCodeInternal
}
+2 -4
View File
@@ -283,11 +283,9 @@ func (e *ErrorInfo) Validate() error {
}
// error.type is a wire-stable Agent branch key, not an open-ended label.
// Keep this set aligned with exitCodeForErrorInfo. "permission" is the
// compatibility projection for PAT failures (rc=4); "wait" is the wait
// phase's terminal-failure projection (rc=8, e.g. an approval observed
// REJECTED after --wait).
// compatibility projection for PAT failures (rc=4).
switch errorType {
case "api", "auth", "validation", "permission", "discovery", "internal", "wait":
case "api", "auth", "validation", "permission", "discovery", "internal":
default:
return fmt.Errorf("output: unsupported failure error.type %q", e.Type)
}
@@ -19,7 +19,6 @@ type forgedResult struct {
func (r forgedResult) Outcome() Outcome { return r.env.Outcome }
func (r forgedResult) ExitCode() int { return r.exit }
func (r forgedResult) Data() any { return r.env.Data }
func (r forgedResult) envelope() *Envelope { copy := r.env; return &copy }
type cloneNode struct {
-90
View File
@@ -19,10 +19,6 @@ import (
type CommandResult interface {
Outcome() Outcome
ExitCode() int
// Data returns the accepted payload (already deep-copied). The wait
// phase reads it to resolve the resource identifier an event stream
// correlates against.
Data() any
envelope() *Envelope
}
@@ -34,7 +30,6 @@ type commandResult struct {
func (r *commandResult) Outcome() Outcome { return r.env.Outcome }
func (r *commandResult) ExitCode() int { return r.exitCode }
func (r *commandResult) Data() any { return cloneResultData(r.env.Data) }
func (r *commandResult) envelope() *Envelope {
copy := cloneEnvelope(r.env)
return &copy
@@ -76,91 +71,6 @@ func Partial(data *PartialData, opts ...ResultOption) CommandResult {
return newCommandResult(OutcomePartialFailure, data, nil, opts...)
}
// WithMeta(meta *Meta) ResultOption was declared above; the two options below
// exist for the wait phase.
// WithErrorInfo replaces the error info of the envelope. Used when a wait
// phase closes an accepted result into failure: envelope invariant I3
// requires an error iff the outcome is failure.
func WithErrorInfo(info *ErrorInfo) ResultOption {
return ResultOption{apply: func(env *Envelope) {
if info != nil {
info = cloneErrorInfo(info)
}
env.Error = info
}}
}
// WithOperationTimedOut marks the envelope's async operation as timed out at
// the last observed state, preserving the declared id / next_command resume
// facts (契约规范 §2.2: 超时必须保持 State 真实值并置 TimedOut:true). A result
// without operation info keeps nil — the pending envelope invariant then fails
// at emission, surfacing the leaf bug instead of synthesizing fake resume
// facts.
func WithOperationTimedOut(state string) ResultOption {
return ResultOption{apply: func(env *Envelope) {
if env.Meta == nil || env.Meta.Operation == nil {
return
}
operation := *env.Meta.Operation
// A subscribe/poll that never observed a status still times out
// against the accepted pending result: keep the original state
// rather than wiping it to empty (pending requires operation.state).
if strings.TrimSpace(state) != "" {
operation.State = state
}
operation.TimedOut = true
env.Meta.Operation = &operation
}}
}
// WithOperationTerminalState closes the envelope's async operation at the observed
// terminal status (契约规范 §2.2: 终态封装必须同步 operation.state — a success or
// failure close that kept the acceptance-phase state would emit a
// self-contradicting envelope such as outcome=success with
// operation.state=processing). The declared id / next_command facts are kept
// as the operation identity, and timed_out is cleared: the §2.2 anti-spoof
// rule forbids a timed-out claim on an operation that reached a terminal
// state. A result without operation info is left untouched.
func WithOperationTerminalState(state string) ResultOption {
return ResultOption{apply: func(env *Envelope) {
if env.Meta == nil || env.Meta.Operation == nil {
return
}
operation := *env.Meta.Operation
if strings.TrimSpace(state) != "" {
operation.State = state
}
operation.TimedOut = false
env.Meta.Operation = &operation
}}
}
// WithOutcome rewraps an existing result with a new outcome, preserving data,
// meta, identity, and any error info (subject to the opts). The corecmd wait
// phase uses it to close an accepted result into its terminal (or timed-out
// pending) outcome; the exit code is re-derived from the new envelope.
func WithOutcome(result CommandResult, outcome Outcome, opts ...ResultOption) CommandResult {
env := *result.envelope()
for _, opt := range opts {
if opt.apply != nil {
opt.apply(&env)
}
}
env.Outcome = outcome
env.OK = outcome == OutcomeSuccess || outcome == OutcomePending
if outcome == OutcomeFailure {
// Invariant I3: data and error are mutually exclusive. Closing into
// failure replaces the accepted data with the failure error.
env.Data = nil
}
exitCode := ExitCodeForEnvelope(&env)
if env.Error != nil {
env.Error.ExitCode = exitCode
}
return &commandResult{env: env, exitCode: exitCode}
}
// Failure constructs an immutable typed failure result.
func Failure(info *ErrorInfo, opts ...ResultOption) CommandResult {
if info != nil {
-165
View File
@@ -1,165 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package output
import (
"strings"
"testing"
)
func pendingAcceptedResult() CommandResult {
return Pending(map[string]any{"id": "job-1"}, &OperationInfo{
ID: "job-1",
State: "NEW",
NextCommand: "dws wait-sample --id job-1",
})
}
func TestWithOutcomeClosesSuccessPreservingDataAndMeta(t *testing.T) {
result := WithOutcome(pendingAcceptedResult(), OutcomeSuccess)
if result.Outcome() != OutcomeSuccess || result.ExitCode() != 0 {
t.Fatalf("outcome=%s exit=%d", result.Outcome(), result.ExitCode())
}
env := result.envelope()
if env.Meta == nil || env.Meta.Operation == nil {
t.Fatal("operation info lost on success close")
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
}
func TestWithOutcomeFailureDropsDataAndCarriesError(t *testing.T) {
result := WithOutcome(pendingAcceptedResult(), OutcomeFailure, WithErrorInfo(&ErrorInfo{
Type: "wait",
Subtype: "terminal_failure",
Message: "等待到达失败终态:REJECTED",
}))
if result.Outcome() != OutcomeFailure || result.ExitCode() != 8 {
t.Fatalf("outcome=%s exit=%d, want failure/8", result.Outcome(), result.ExitCode())
}
env := result.envelope()
if env.Data != nil {
t.Fatal("failure close must drop data (I3)")
}
if env.Error == nil || env.Error.Type != "wait" {
t.Fatal("failure close must carry error info (I3)")
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
}
func TestWithOperationTimedOutMarksStateAndKeepsResumeFacts(t *testing.T) {
result := WithOutcome(pendingAcceptedResult(), OutcomePending, WithOperationTimedOut("RUNNING"))
env := result.envelope()
op := env.Meta.Operation
if op.State != "RUNNING" || !op.TimedOut || op.ID != "job-1" || op.NextCommand != "dws wait-sample --id job-1" {
t.Fatalf("operation=%+v", op)
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
}
func TestWithOperationTimedOutEmptyStatePreservesExisting(t *testing.T) {
result := WithOutcome(pendingAcceptedResult(), OutcomePending, WithOperationTimedOut(""))
op := result.envelope().Meta.Operation
if op.State != "NEW" || !op.TimedOut {
t.Fatalf("operation=%+v, want original state kept and timed_out set", op)
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
}
func TestWithOperationTimedOutWithoutOperationInfoLeavesEnvelopeUntouched(t *testing.T) {
// A result without operation info keeps nil — ValidateResult must then
// reject the pending envelope instead of the option synthesizing fake
// resume facts.
result := WithOutcome(Success(map[string]any{"ok": true}), OutcomePending, WithOperationTimedOut("RUNNING"))
env := result.envelope()
if env.Meta != nil && env.Meta.Operation != nil {
t.Fatalf("operation=%+v, want untouched", env.Meta.Operation)
}
err := ValidateResult(result)
if err == nil || !strings.Contains(err.Error(), "meta.operation") {
t.Fatalf("err=%v, want pending-requires-operation rejection", err)
}
}
func TestDataAccessorReturnsDeepCopy(t *testing.T) {
result := pendingAcceptedResult()
data, ok := result.Data().(map[string]any)
if !ok {
t.Fatalf("data=%#v", result.Data())
}
data["id"] = "mutated"
again := result.Data().(map[string]any)
if again["id"] != "job-1" {
t.Fatalf("Data() aliased internal state: %#v", again)
}
}
func TestWithOperationTerminalStateSyncsStateAndClearsTimedOut(t *testing.T) {
for _, tc := range []struct {
name string
outcome Outcome
}{
{"success", OutcomeSuccess},
{"failure", OutcomeFailure},
} {
t.Run(tc.name, func(t *testing.T) {
opts := []ResultOption{WithOperationTerminalState("COMPLETED")}
if tc.outcome == OutcomeFailure {
opts = append(opts, WithErrorInfo(&ErrorInfo{
Type: "wait", Subtype: "terminal_failure", Message: "等待到达失败终态:COMPLETED",
}))
}
result := WithOutcome(pendingAcceptedResult(), tc.outcome, opts...)
op := result.envelope().Meta.Operation
if op.State != "COMPLETED" || op.TimedOut {
t.Fatalf("operation=%+v, want terminal state synced and timed_out cleared", op)
}
if op.ID != "job-1" || op.NextCommand != "dws wait-sample --id job-1" {
t.Fatalf("operation=%+v, want resume identity facts preserved", op)
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
})
}
}
func TestWithOperationTerminalStateBlankKeepsObservedState(t *testing.T) {
// A terminal observation that carries no status must keep the last known
// state rather than wipe it: operation.state may never be emptied by a
// close transition.
result := WithOutcome(pendingAcceptedResult(), OutcomeSuccess, WithOperationTerminalState(" "))
op := result.envelope().Meta.Operation
if op.State != "NEW" || op.TimedOut {
t.Fatalf("operation=%+v, want original state kept and timed_out cleared", op)
}
if err := ValidateResult(result); err != nil {
t.Fatal(err)
}
}
func TestWithOperationTerminalStateWithoutOperationInfoLeavesEnvelopeUntouched(t *testing.T) {
result := WithOutcome(Success(map[string]any{"ok": true}), OutcomeSuccess, WithOperationTerminalState("COMPLETED"))
env := result.envelope()
if env.Meta != nil && env.Meta.Operation != nil {
t.Fatalf("operation=%+v, want untouched", env.Meta.Operation)
}
}
-273
View File
@@ -1,273 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Package wait is the framework terminal-state wait engine behind the
// reviewed contract.WaitSpec capability. It owns polling cadence, status
// extraction, and status→outcome mapping; it knows nothing about Cobra, MCP,
// or any product backend. How one poll executes is supplied by the leaf's
// WaitPoll hook (corecmd), so "poll = an existing read command" stays a leaf
// decision rather than a framework assumption.
package wait
import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/contract"
)
// DefaultPollInterval is the cadence between polls when LoopSpec.Interval is
// zero. The first poll runs immediately so an already-terminal resource does
// not pay a sleep tax.
const DefaultPollInterval = 2 * time.Second
// MaxPollInterval caps the exponential backoff growth between polls so a long
// wait cannot degenerate into effectively-blind polling.
const MaxPollInterval = 30 * time.Second
// PollDoc is one decoded poll response document (typically the unified-output
// envelope data of the poll command).
type PollDoc map[string]any
// Poller executes one poll. Returning an error fails the wait phase; the
// engine never retries a poller error because read commands failing is a real
// failure, not a "not yet" signal.
type Poller func(ctx context.Context) (PollDoc, error)
// LoopSpec is the runtime-resolved projection of contract.WaitSpec plus the
// caller-provided timeout.
type LoopSpec struct {
StatusQuery string
Terminal map[string]contract.ResultOutcome
Pending []string
Timeout time.Duration
Interval time.Duration
}
// Outcome is the closed result of a wait loop. TimedOut reports deadline
// exhaustion (Outcome is then pending — an accepted-but-not-terminal state is
// not a process failure per the exit-code contract); Status is the last
// observed status value.
type Outcome struct {
Status string
Outcome contract.ResultOutcome
Attempts int
TimedOut bool
}
// ErrUnknownStatus reports a status value that is neither declared terminal
// nor declared pending. Unknown fails closed: mapping it to pending could
// hide a real state change until timeout, mapping it to success is worse.
type ErrUnknownStatus struct {
Status string
Query string
}
func (e *ErrUnknownStatus) Error() string {
return fmt.Sprintf("wait: status %q (from %q) is neither terminal nor pending", e.Status, e.Query)
}
// Run polls poller until a declared terminal status, deadline exhaustion, or
// a poller error. The first poll is immediate; subsequent polls back off
// exponentially (×1.5) from Interval, capped at MaxPollInterval. Deadline
// exhaustion anywhere — before a poll, during a poll (a context-aware poller
// returns ctx.Err()), or during the wait between polls — always closes as
// timed-out pending with the last observed status, never as a poll failure.
func Run(ctx context.Context, spec LoopSpec, poll Poller) (Outcome, error) {
if spec.Interval <= 0 {
spec.Interval = DefaultPollInterval
}
if spec.Timeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, spec.Timeout)
defer cancel()
}
pending := make(map[string]bool, len(spec.Pending))
for _, value := range spec.Pending {
pending[value] = true
}
timedOut := func(status string, attempts int) Outcome {
return Outcome{Status: status, Outcome: contract.ResultOutcomePending, Attempts: attempts, TimedOut: true}
}
interval := spec.Interval
attempts := 0
lastStatus := ""
for {
if ctx.Err() != nil {
return timedOut(lastStatus, attempts), nil
}
doc, err := poll(ctx)
if err != nil {
if ctx.Err() != nil {
return timedOut(lastStatus, attempts), nil
}
return Outcome{Attempts: attempts}, fmt.Errorf("wait: poll failed: %w", err)
}
attempts++
status, ok := ExtractStatus(doc, spec.StatusQuery)
if !ok {
return Outcome{Attempts: attempts}, fmt.Errorf(
"wait: status query %q not found in poll result", spec.StatusQuery)
}
lastStatus = status
if outcome, ok := spec.Terminal[status]; ok {
return Outcome{Status: status, Outcome: outcome, Attempts: attempts}, nil
}
if !pending[status] {
return Outcome{Status: status, Attempts: attempts}, &ErrUnknownStatus{Status: status, Query: spec.StatusQuery}
}
timer := time.NewTimer(interval)
select {
case <-ctx.Done():
timer.Stop()
return timedOut(status, attempts), nil
case <-timer.C:
}
interval = nextInterval(interval)
}
}
func nextInterval(current time.Duration) time.Duration {
next := current * 3 / 2
if next > MaxPollInterval {
next = MaxPollInterval
}
return next
}
// ExtractStatus resolves a dotted status query against a poll document. Each
// segment walks one map level; array indexes are not supported because wait
// targets a single resource. Numeric segments are stringified, so a document
// decoded with json.Number keys still resolves.
func ExtractStatus(doc PollDoc, query string) (string, bool) {
query = strings.TrimSpace(query)
if query == "" {
return "", false
}
// PollDoc is a defined type, so its dynamic type does not satisfy a
// map[string]any assertion — convert once at the boundary; nested values
// from JSON decoding are plain maps.
var current any = map[string]any(doc)
for _, segment := range strings.Split(query, ".") {
segment = strings.TrimSpace(segment)
if segment == "" {
return "", false
}
node, ok := current.(map[string]any)
if !ok {
return "", false
}
value, ok := node[segment]
if !ok {
return "", false
}
current = value
}
switch value := current.(type) {
case string:
return value, true
case fmt.Stringer:
return value.String(), true
case bool:
return strconv.FormatBool(value), true
case int:
return strconv.Itoa(value), true
case int64:
return strconv.FormatInt(value, 10), true
case float64:
return strconv.FormatFloat(value, 'f', -1, 64), true
default:
return "", false
}
}
// IsUnknownStatus reports whether err is the closed fail-on-unknown error.
func IsUnknownStatus(err error) bool {
var unknown *ErrUnknownStatus
return errors.As(err, &unknown)
}
// EventStream is the leaf-owned push subscription consumed by the event
// phase (the WaitEvents hook in corecmd). Recv delivers the next decoded
// event document; it returns an error or io.EOF-style termination when the
// stream ends — the engine treats non-terminal termination as a stream
// failure the caller (auto mode) may fall back from.
type EventStream interface {
Recv(ctx context.Context) (PollDoc, error)
}
// EventLoopSpec is the event-phase projection of contract.WaitSpec.
type EventLoopSpec struct {
StatusQuery string
MatchField string
Terminal map[string]contract.ResultOutcome
Pending []string
Timeout time.Duration
}
// ErrEventStreamEnded reports a stream that terminated before a terminal
// status. Auto mode uses it to fall back to polling; strict event mode
// surfaces it as a wait failure.
var ErrEventStreamEnded = errors.New("wait: event stream ended before a terminal status")
// RunEvent consumes stream until a correlated event reaches a declared
// terminal status, the deadline exhausts, or the stream ends. Events whose
// MatchField value does not equal resource are ignored (other resources on
// the same channel); a correlated event with an unknown status fails closed
// exactly like a poll would.
func RunEvent(ctx context.Context, spec EventLoopSpec, resource string, stream EventStream) (Outcome, error) {
if spec.Timeout > 0 {
var cancel context.CancelFunc
ctx, cancel = context.WithTimeout(ctx, spec.Timeout)
defer cancel()
}
pending := make(map[string]bool, len(spec.Pending))
for _, value := range spec.Pending {
pending[value] = true
}
attempts := 0
lastStatus := ""
for {
doc, err := stream.Recv(ctx)
if err != nil {
if ctx.Err() != nil {
return Outcome{Status: lastStatus, Outcome: contract.ResultOutcomePending, Attempts: attempts, TimedOut: true}, nil
}
// Wrap with ErrEventStreamEnded so auto mode can fall back to
// polling while correlated-status failures (unknown status,
// missing status query) stay non-recoverable.
return Outcome{Attempts: attempts}, fmt.Errorf("%w: %v", ErrEventStreamEnded, err)
}
attempts++
correlated, ok := ExtractStatus(doc, spec.MatchField)
if !ok || correlated != resource {
continue
}
status, ok := ExtractStatus(doc, spec.StatusQuery)
if !ok {
return Outcome{Attempts: attempts}, fmt.Errorf(
"wait: status query %q not found in event document", spec.StatusQuery)
}
lastStatus = status
if outcome, ok := spec.Terminal[status]; ok {
return Outcome{Status: status, Outcome: outcome, Attempts: attempts}, nil
}
if !pending[status] {
return Outcome{Status: status, Attempts: attempts}, &ErrUnknownStatus{Status: status, Query: spec.StatusQuery}
}
}
}
-365
View File
@@ -1,365 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package wait
import (
"context"
"errors"
"io"
"strings"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/corecmd/contract"
)
func loopSpec() LoopSpec {
return LoopSpec{
StatusQuery: "result.status",
Terminal: map[string]contract.ResultOutcome{
"COMPLETED": contract.ResultOutcomeSuccess,
"REJECTED": contract.ResultOutcomeFailure,
},
Pending: []string{"NEW", "RUNNING"},
Interval: time.Millisecond,
}
}
func TestExtractStatusResolvesDottedPaths(t *testing.T) {
doc := PollDoc{
"result": map[string]any{
"instance": map[string]any{"status": "RUNNING"},
"count": float64(3),
},
}
if status, ok := ExtractStatus(doc, "result.instance.status"); !ok || status != "RUNNING" {
t.Fatalf("status=%q ok=%v", status, ok)
}
if status, ok := ExtractStatus(doc, "result.count"); !ok || status != "3" {
t.Fatalf("numeric status=%q ok=%v", status, ok)
}
if _, ok := ExtractStatus(doc, "result.missing"); ok {
t.Fatal("missing path resolved")
}
if _, ok := ExtractStatus(doc, "result.instance.status.deep"); ok {
t.Fatal("descending into a scalar resolved")
}
if _, ok := ExtractStatus(doc, ""); ok {
t.Fatal("empty query resolved")
}
}
func TestRunReturnsTerminalOnFirstPoll(t *testing.T) {
polls := 0
outcome, err := Run(context.Background(), loopSpec(), func(context.Context) (PollDoc, error) {
polls++
return PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
})
if err != nil {
t.Fatal(err)
}
if polls != 1 || outcome.Attempts != 1 {
t.Fatalf("polls=%d attempts=%d", polls, outcome.Attempts)
}
if outcome.Outcome != contract.ResultOutcomeSuccess || outcome.Status != "COMPLETED" {
t.Fatalf("outcome=%s status=%s", outcome.Outcome, outcome.Status)
}
}
func TestRunPollsUntilTerminal(t *testing.T) {
seen := []string{"NEW", "RUNNING", "RUNNING", "COMPLETED"}
index := 0
outcome, err := Run(context.Background(), loopSpec(), func(context.Context) (PollDoc, error) {
status := seen[index]
index++
return PollDoc{"result": map[string]any{"status": status}}, nil
})
if err != nil {
t.Fatal(err)
}
if outcome.Attempts != len(seen) || outcome.Outcome != contract.ResultOutcomeSuccess {
t.Fatalf("attempts=%d outcome=%s", outcome.Attempts, outcome.Outcome)
}
}
func TestRunTimesOutAsPendingDuringWait(t *testing.T) {
spec := loopSpec()
spec.Timeout = 5 * time.Millisecond
polls := 0
outcome, err := Run(context.Background(), spec, func(context.Context) (PollDoc, error) {
polls++
return PollDoc{"result": map[string]any{"status": "RUNNING"}}, nil
})
if err != nil {
t.Fatal(err)
}
if !outcome.TimedOut || outcome.Outcome != contract.ResultOutcomePending {
t.Fatalf("timedOut=%v outcome=%s", outcome.TimedOut, outcome.Outcome)
}
if outcome.Status != "RUNNING" {
t.Fatalf("status=%q, want last observed", outcome.Status)
}
if polls == 0 {
t.Fatal("timeout during wait must still have polled at least once")
}
}
func TestRunTimesOutAsPendingWhenPollerRespectsDeadline(t *testing.T) {
spec := loopSpec()
spec.Timeout = 5 * time.Millisecond
polls := 0
// A context-aware poller: blocks until the deadline, then reports the
// cancellation as an error — the loop must close it as timed-out pending,
// never as a poll failure.
outcome, err := Run(context.Background(), spec, func(ctx context.Context) (PollDoc, error) {
polls++
<-ctx.Done()
return nil, ctx.Err()
})
if err != nil {
t.Fatalf("deadline during poll closed as error: %v", err)
}
if !outcome.TimedOut || outcome.Outcome != contract.ResultOutcomePending {
t.Fatalf("timedOut=%v outcome=%s", outcome.TimedOut, outcome.Outcome)
}
}
func TestRunTimesOutBeforeFirstPoll(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
outcome, err := Run(ctx, loopSpec(), func(context.Context) (PollDoc, error) {
t.Fatal("poller ran on a pre-cancelled context")
return nil, nil
})
if err != nil {
t.Fatal(err)
}
if !outcome.TimedOut || outcome.Attempts != 0 || outcome.Outcome != contract.ResultOutcomePending {
t.Fatalf("timedOut=%v attempts=%d outcome=%s", outcome.TimedOut, outcome.Attempts, outcome.Outcome)
}
}
func TestRunFailsClosedOnUnknownStatus(t *testing.T) {
_, err := Run(context.Background(), loopSpec(), func(context.Context) (PollDoc, error) {
return PollDoc{"result": map[string]any{"status": "Mystery"}}, nil
})
if !IsUnknownStatus(err) {
t.Fatalf("err=%v want unknown-status", err)
}
}
func TestRunFailsOnMissingStatusQuery(t *testing.T) {
_, err := Run(context.Background(), loopSpec(), func(context.Context) (PollDoc, error) {
return PollDoc{"unexpected": true}, nil
})
if err == nil || !errors.Is(err, err) {
t.Fatalf("err=%v", err)
}
}
func TestRunPropagatesPollerError(t *testing.T) {
boom := errors.New("rpc down")
_, err := Run(context.Background(), loopSpec(), func(context.Context) (PollDoc, error) {
return nil, boom
})
if !errors.Is(err, boom) {
t.Fatalf("err=%v", err)
}
}
func TestUnknownStatusErrorCarriesStatusAndQuery(t *testing.T) {
err := &ErrUnknownStatus{Status: "Mystery", Query: "result.status"}
message := err.Error()
if !strings.Contains(message, "Mystery") || !strings.Contains(message, "result.status") {
t.Fatalf("message=%q", message)
}
}
func TestRunAppliesDefaultIntervalWhenUnset(t *testing.T) {
spec := loopSpec()
spec.Interval = 0
polls := 0
// Terminal on the second poll forces one interval wait; with Interval=0
// the loop must still work using DefaultPollInterval (not spin/panic).
_, err := Run(context.Background(), spec, func(context.Context) (PollDoc, error) {
polls++
if polls == 1 {
return PollDoc{"result": map[string]any{"status": "NEW"}}, nil
}
return PollDoc{"result": map[string]any{"status": "COMPLETED"}}, nil
})
if err != nil {
t.Fatal(err)
}
if polls != 2 {
t.Fatalf("polls=%d", polls)
}
}
func TestRunTimesOutDuringWaitBetweenPolls(t *testing.T) {
spec := loopSpec()
spec.Interval = time.Hour // the deadline wins long before the next poll
spec.Timeout = 5 * time.Millisecond
outcome, err := Run(context.Background(), spec, func(context.Context) (PollDoc, error) {
return PollDoc{"result": map[string]any{"status": "RUNNING"}}, nil
})
if err != nil {
t.Fatal(err)
}
if !outcome.TimedOut || outcome.Outcome != contract.ResultOutcomePending || outcome.Status != "RUNNING" {
t.Fatalf("outcome=%+v", outcome)
}
}
type stringStatus string
func (s stringStatus) String() string { return string(s) }
func TestExtractStatusCoversScalarShapes(t *testing.T) {
doc := PollDoc{
"result": map[string]any{
"flag": true,
"small": 7,
"big": int64(9007199254740993),
"fraction": 1.5,
"custom": stringStatus("CUSTOM"),
"nested": map[string]any{"deep": "x"},
},
}
cases := map[string]string{
"result.flag": "true",
"result.small": "7",
"result.big": "9007199254740993",
"result.fraction": "1.5",
"result.custom": "CUSTOM",
}
for query, want := range cases {
if got, ok := ExtractStatus(doc, query); !ok || got != want {
t.Fatalf("query=%s got=%q ok=%v want=%q", query, got, ok, want)
}
}
if _, ok := ExtractStatus(doc, "result.nested"); ok {
t.Fatal("non-scalar nested map must not resolve")
}
if _, ok := ExtractStatus(doc, "result..flag"); ok {
t.Fatal("empty segment must not resolve")
}
}
func TestNextIntervalCapsAtMax(t *testing.T) {
if got := nextInterval(MaxPollInterval); got != MaxPollInterval {
t.Fatalf("nextInterval(max)=%s", got)
}
if got := nextInterval(10 * time.Millisecond); got != 15*time.Millisecond {
t.Fatalf("nextInterval(10ms)=%s", got)
}
}
type fakeEventStream struct {
events []PollDoc
err error // returned after events are exhausted (nil = clean end)
block bool // hold until the context deadline
}
func (f *fakeEventStream) Recv(ctx context.Context) (PollDoc, error) {
if f.block {
<-ctx.Done()
return nil, ctx.Err()
}
if len(f.events) > 0 {
doc := f.events[0]
f.events = f.events[1:]
return doc, nil
}
if f.err != nil {
return nil, f.err
}
return nil, io.EOF
}
func eventLoopSpec() EventLoopSpec {
return EventLoopSpec{
StatusQuery: "result.status",
MatchField: "process_instance_id",
Terminal: map[string]contract.ResultOutcome{
"COMPLETED": contract.ResultOutcomeSuccess,
"REJECTED": contract.ResultOutcomeFailure,
},
Pending: []string{"RUNNING"},
}
}
func approvalEvent(instance, status string) PollDoc {
return PollDoc{"process_instance_id": instance, "result": map[string]any{"status": status}}
}
func TestRunEventReturnsCorrelatedTerminal(t *testing.T) {
stream := &fakeEventStream{events: []PollDoc{
approvalEvent("other-instance", "COMPLETED"), // other resource: ignored
approvalEvent("job-1", "RUNNING"), // correlated pending: kept waiting
approvalEvent("job-1", "COMPLETED"),
}}
outcome, err := RunEvent(context.Background(), eventLoopSpec(), "job-1", stream)
if err != nil {
t.Fatal(err)
}
if outcome.Outcome != contract.ResultOutcomeSuccess || outcome.Status != "COMPLETED" {
t.Fatalf("outcome=%+v", outcome)
}
}
func TestRunEventFailsClosedOnUnknownCorrelatedStatus(t *testing.T) {
stream := &fakeEventStream{events: []PollDoc{approvalEvent("job-1", "Mystery")}}
_, err := RunEvent(context.Background(), eventLoopSpec(), "job-1", stream)
if !IsUnknownStatus(err) {
t.Fatalf("err=%v", err)
}
}
func TestRunEventRejectsCorrelatedEventWithoutStatus(t *testing.T) {
stream := &fakeEventStream{events: []PollDoc{
{"process_instance_id": "job-1"}, // correlated but no status document
}}
_, err := RunEvent(context.Background(), eventLoopSpec(), "job-1", stream)
if err == nil || !strings.Contains(err.Error(), "status query") {
t.Fatalf("err=%v", err)
}
}
func TestRunEventStreamEndSurfacesFallbackSentinel(t *testing.T) {
stream := &fakeEventStream{events: []PollDoc{approvalEvent("job-1", "RUNNING")}}
_, err := RunEvent(context.Background(), eventLoopSpec(), "job-1", stream)
if !errors.Is(err, ErrEventStreamEnded) {
t.Fatalf("err=%v, want ErrEventStreamEnded", err)
}
failing := &fakeEventStream{err: errors.New("transport reset")}
_, err = RunEvent(context.Background(), eventLoopSpec(), "job-1", failing)
if !errors.Is(err, ErrEventStreamEnded) {
t.Fatalf("err=%v, want ErrEventStreamEnded wrapping the transport error", err)
}
}
func TestRunEventTimesOutAsPendingWhileBlocked(t *testing.T) {
spec := eventLoopSpec()
spec.Timeout = 5 * time.Millisecond
stream := &fakeEventStream{block: true}
outcome, err := RunEvent(context.Background(), spec, "job-1", stream)
if err != nil {
t.Fatal(err)
}
if !outcome.TimedOut || outcome.Outcome != contract.ResultOutcomePending {
t.Fatalf("outcome=%+v", outcome)
}
}
+1 -1
View File
@@ -1,6 +1,6 @@
---
name: dws
description: 管理钉钉产品能力(AI表格/AI搜问/日历/通讯录/群聊与机器人/待办/审批/考勤/日志/DING消息/开放平台文档/钉钉文档/钉钉云盘/原生Markdown文件/AI听记/邮箱/在线电子表格/知识库等)。当用户需要操作表格数据、管理日程会议、模糊找人/查谁负责某事项、查询通讯录、管理群聊、机器人发消息、创建待办、提交审批、查看考勤、提交日报周报(钉钉日志模版)、读写钉钉文档、上传下载云盘文件、读取或修改原生.md文件、查询听记纪要、收发邮件、读写在线电子表格(axls)、管理钉钉知识库,或订阅个人 IM 事件或 OA 审批事件、实时监听群成员加入、群成员退出、群改名和群解散、审批实例发起/抄送/终止/完成,以及审批任务创建/完成/转交时使用。
description: 管理钉钉产品能力(AI表格/AI搜问/日历/通讯录/群聊与机器人/待办/审批/考勤/日志/DING消息/开放平台文档/钉钉文档/钉钉云盘/原生Markdown文件/AI听记/邮箱/在线电子表格/知识库等)。当用户需要操作表格数据、管理日程会议、模糊找人/查谁负责某事项、查询通讯录、管理群聊、机器人发消息、创建待办、提交审批、查看考勤、提交日报周报(钉钉日志模版)、读写钉钉文档、上传下载云盘文件、读取或修改原生.md文件、查询听记纪要、收发邮件、读写在线电子表格(axls)、管理钉钉知识库,或订阅个人 IM 事件或 OA 审批事件、实时监听群成员加入、群成员退出、群改名和群解散、审批实例发起/终止/完成,以及审批任务创建/完成/转交时使用。
cli_version: ">=1.0.15"
---
+4 -9
View File
@@ -47,11 +47,10 @@
| `user_oa_approval_task_finished` | 审批任务已完成 | 无 |
| `user_oa_approval_task_redirected` | 审批任务已转交 | 无 |
| `user_oa_approval_instance_started` | 审批实例已发起 | 无 |
| `user_oa_approval_instance_cc` | 审批实例到达抄送节点,发送给被抄送人 | 无 |
| `user_oa_approval_instance_terminated` | 审批实例已终止 | 无 |
| `user_oa_approval_instance_finished` | 审批实例完成,发送给审批单发起人 | 无 |
只承认上表 23 个事件码。默认身份就是当前用户,使用当前用户 OAuth 登录态,不要额外加身份切换 flag。七个 OA 事件订阅当前用户相关的全部审批事件,规则均为 `all`、空 `filterRule`,不需要目标参数。
只承认上表 22 个事件码。默认身份就是当前用户,使用当前用户 OAuth 登录态,不要额外加身份切换 flag。六个 OA 事件订阅当前用户相关的全部审批事件,规则均为 `all`、空 `filterRule`,不需要目标参数。
## Intent mapping
@@ -77,10 +76,9 @@
| "审批任务完成时通知我" | `event consume`,事件码 `user_oa_approval_task_finished`,参数 `--flatten -f ndjson` |
| "审批任务被转交时通知我" | `event consume`,事件码 `user_oa_approval_task_redirected`,参数 `--flatten -f ndjson` |
| "有审批单发起时通知我" | `event consume`,事件码 `user_oa_approval_instance_started`,参数 `--flatten -f ndjson` |
| "有审批抄送给我时通知我" | `event consume`,事件码 `user_oa_approval_instance_cc`,参数 `--flatten -f ndjson` |
| "有审批单终止时通知我" | `event consume`,事件码 `user_oa_approval_instance_terminated`,参数 `--flatten -f ndjson` |
| "监听我发起的审批何时完成" / "审批实例完成时通知我" | `event consume`,事件码 `user_oa_approval_instance_finished`,参数 `--flatten -f ndjson` |
| "同时监听全部已公开 OA 事件" | 一个 consume 放入七个 OA event key,不加目标或消息过滤参数 |
| "同时监听全部已公开 OA 事件" | 一个 consume 放入六个 OA event key,不加目标或消息过滤参数 |
| "查看个人事件 schema" | `dws event schema <event_key> --flatten` |
| "看个人事件订阅状态" | `dws event status --event <event_key>` |
| "停止这个个人事件订阅" | `dws event stop <subscribe_id> --dry-run`,确认后改用 `--yes` |
@@ -132,7 +130,6 @@ dws event schema user_oa_approval_task_created --flatten
dws event schema user_oa_approval_task_finished --flatten
dws event schema user_oa_approval_task_redirected --flatten
dws event schema user_oa_approval_instance_started --flatten
dws event schema user_oa_approval_instance_cc --flatten
dws event schema user_oa_approval_instance_terminated --flatten
dws event schema user_oa_approval_instance_finished --flatten
```
@@ -160,7 +157,6 @@ dws event consume user_oa_approval_task_created --flatten -f ndjson
dws event consume user_oa_approval_task_finished --flatten -f ndjson
dws event consume user_oa_approval_task_redirected --flatten -f ndjson
dws event consume user_oa_approval_instance_started --flatten -f ndjson
dws event consume user_oa_approval_instance_cc --flatten -f ndjson
dws event consume user_oa_approval_instance_terminated --flatten -f ndjson
dws event consume user_oa_approval_instance_finished --flatten -f ndjson
```
@@ -189,14 +185,13 @@ dws event consume \
user_oa_approval_task_finished \
user_oa_approval_task_redirected \
user_oa_approval_instance_started \
user_oa_approval_instance_cc \
user_oa_approval_instance_terminated \
user_oa_approval_instance_finished \
--flatten \
-f ndjson
```
用户类事件共享 `--user` 或 `--open-dingtalk-id`,群类事件共享 `--group`,无目标 IM 事件可加入任一组合。用户类与群类、不同目标或不同过滤条件要拆成多个进程。七个 OA 事件可以同进程消费并共享 personal bus,但各自建立独立订阅。多事件共享 `--query` / `--filter-json` 时,所选事件必须全部是 IM 消息接收事件;OA 事件单独或组合消费都禁止使用这两个消息过滤参数。
用户类事件共享 `--user` 或 `--open-dingtalk-id`,群类事件共享 `--group`,无目标 IM 事件可加入任一组合。用户类与群类、不同目标或不同过滤条件要拆成多个进程。六个 OA 事件可以同进程消费并共享 personal bus,但各自建立独立订阅。多事件共享 `--query` / `--filter-json` 时,所选事件必须全部是 IM 消息接收事件;OA 事件单独或组合消费都禁止使用这两个消息过滤参数。
上述所有 `*_o2o` 命令和 `user_im_message_receive_user` 都可将 `--user <userId>` 替换为 `--open-dingtalk-id <openDingtalkId>`,但两个参数不能同时使用。
@@ -221,7 +216,7 @@ dws event stop --all --yes
## 订阅创建失败与重试预算
以下约束适用于上表全部 23 个公开个人事件(16 个 IM + 7 个 OA)以及多事件命令中的每一项,只治理 `[event] ready` 之前的订阅创建;ready 之后的 Stream 断线由长连接重连机制处理。
以下约束适用于上表全部 22 个公开个人事件(16 个 IM + 6 个 OA)以及多事件命令中的每一项,只治理 `[event] ready` 之前的订阅创建;ready 之后的 Stream 断线由长连接重连机制处理。
- `0/2/1` 是 **Agent/host 编排约束**,不是 CLI 持久化硬总次数上限。每次 `dws event consume` 调用对每个逻辑订阅最多发送一次订阅创建 HTTP 请求,进程内不会自动重试。CLI 本地状态只持久化 `in_flight`、`cooldown`、`terminal_hold` 三种保护状态,不持久化或计算跨调用的 Agent/host 尝试次数。
- 解析人名或群名、执行 `event consume` 以及后续 `event status/stop` 必须使用同一个 `--profile`。不得把其它 profile 下解析出的 userId、openDingtalkId 或 openConversationId 直接带入当前 profile 的订阅。
+5 -5
View File
@@ -1,6 +1,6 @@
---
name: dingtalk-event
description: 钉钉个人 IM 与 OA 审批事件长连接监听。Use when 用户说监听消息/@我/某人/某群/全部消息、已读/撤回/reaction、群成员加入/群成员退出/群状态变化,或监听审批任务创建/完成/转交、审批实例发起/抄送/终止/完成。命令前缀:dws event。
description: 钉钉个人 IM 与 OA 审批事件长连接监听。Use when 用户说监听消息/@我/某人/某群/全部消息、已读/撤回/reaction、群成员加入/群成员退出/群状态变化,或监听审批任务创建/完成/转交、审批实例发起/终止/完成。命令前缀:dws event。
metadata:
cli_version: ">=0.2.14"
category: product
@@ -45,7 +45,7 @@ metadata:
- `group` 必须且只能传 `--chat-id` 或 `--chat-query` 之一。
- `--query` 只用于纯 `message` 监听;混入 reaction/read/recall 时不得使用。
OA 事件不进入 `+listen-im`。七个公开 OA EventKey 都订阅当前 OAuth 用户相关的全部审批事件,使用 `ruleType=all`、`filterRule={}`;不接受 `--user`、`--open-dingtalk-id`、`--group`、`--query` 或 `--filter-json`。七项可放入同一个 consume,每项建立独立订阅并共享 bus。
OA 事件不进入 `+listen-im`。六个公开 OA EventKey 都订阅当前 OAuth 用户相关的全部审批事件,使用 `ruleType=all`、`filterRule={}`;不接受 `--user`、`--open-dingtalk-id`、`--group`、`--query` 或 `--filter-json`。六项可放入同一个 consume,每项建立独立订阅并共享 bus。
自然姓名和群名由 CLI 内部唯一解析:零命中或多候选返回结构化失败,在创建任何订阅前停止。`--dry-run` 走同一解析链。解析、监听、状态和停止必须使用同一个 `--profile`,不得跨组织搬运 ID。
@@ -65,7 +65,7 @@ user_im_group_member_added user_im_group_member_exited
user_im_group_disbanded
```
七个 OA EventKey 及其输出字段见 [OA 事件参考](references/event-oa.md)。
六个 OA EventKey 及其输出字段见 [OA 事件参考](references/event-oa.md)。
用户类事件传 `--user` 或 `--open-dingtalk-id`,群类事件传 `--group`。群生命周期输出可含 `operator_open_dingtalk_id` 和 `members`;成员项使用 `open_dingtalk_id`。精确组合、兼容性和 Filter 规则见 reference。
@@ -98,7 +98,7 @@ kind + events + target
- `event stop` 会取消订阅并影响本地 consumer:先 `--dry-run`,用户确认后再加 `--yes`。
- 多事件属于一次原始操作;任一订阅启动失败时 Runtime 回滚本次已创建项,不拆成新命令绕过重试预算。
- 这套 `0/2/1` 是 **Agent/host** 编排预算,适用于全部 23 个公开个人 EventKey(16 个 IM + 7 个 OA):`retryable=false` 对应 `max_additional_attempts=0`;`retryable=true` 对应 `max_additional_attempts=2`;`retryable=unknown` 对应 `max_additional_attempts=1`。它不是 CLI 持久化硬总次数上限;每次调用最多创建一次,进程内不会自动重试,CLI 也不持久化或计算跨调用的 Agent/host 尝试次数。
- 这套 `0/2/1` 是 **Agent/host** 编排预算,适用于全部 22 个公开个人 EventKey(16 个 IM + 6 个 OA):`retryable=false` 对应 `max_additional_attempts=0`;`retryable=true` 对应 `max_additional_attempts=2`;`retryable=unknown` 对应 `max_additional_attempts=1`。它不是 CLI 持久化硬总次数上限;每次调用最多创建一次,进程内不会自动重试,CLI 也不持久化或计算跨调用的 Agent/host 尝试次数。
- 重试必须遵守 `retry_after_seconds` / `next_retry_at`。遇到 `in_flight`、`cooldown`、`terminal_hold` 不并发或递归重启同一逻辑订阅,也不换 `subscribe_id` / `trace_id` 绕过保护。
- 认证、profile、订阅保护状态和 bus 排障按失败类型读取 [订阅运维](references/event-im-operations.md),不要在正常路径预加载完整运维手册。
@@ -125,4 +125,4 @@ kind + events + target
| ready、bounded consume 与退出清理 | [event-im-lifecycle.md](references/event-im-lifecycle.md) | 启动/托管/关闭 consumer |
| 扁平字段与事件到 Chat 交接 | [event-im-output.md](references/event-im-output.md) | 解析事件或自动回复 |
| Filter、status/stop、重试与排障 | [event-im-operations.md](references/event-im-operations.md) | 订阅控制或失败恢复 |
| OA 审批事件 | [event-oa.md](references/event-oa.md) | 选择七个 OA EventKey、组合消费或解析审批字段 |
| OA 审批事件 | [event-oa.md](references/event-oa.md) | 选择六个 OA EventKey、组合消费或解析审批字段 |
@@ -1,6 +1,6 @@
# OA 个人审批事件
先读事件产品入口 [SKILL.md](../SKILL.md) 的命令规则、调用流和子进程契约。本参考覆盖当前公开的七个 OA 个人事件:审批实例发起、抄送、终止和完成,以及审批任务创建、完成和转交。
先读事件产品入口 [SKILL.md](../SKILL.md) 的命令规则、调用流和子进程契约。本参考覆盖当前公开的六个 OA 个人事件:审批实例发起、终止和完成,以及审批任务创建、完成和转交。
<!-- dws-intent: event.listen.oa -->实时监听审批事件必须使用 `dws event consume` 长连接,不要轮询 OA 待办或审批实例列表来模拟事件。
@@ -22,11 +22,10 @@ dws auth login
| `user_oa_approval_task_finished` | `all` | 审批任务已完成 | 无 |
| `user_oa_approval_task_redirected` | `all` | 审批任务已转交 | 无 |
| `user_oa_approval_instance_started` | `all` | 审批实例已发起 | 无 |
| `user_oa_approval_instance_cc` | `all` | 审批实例到达抄送节点,发送给被抄送人 | 无 |
| `user_oa_approval_instance_terminated` | `all` | 审批实例已终止 | 无 |
| `user_oa_approval_instance_finished` | `all` | 审批实例完成,发送给审批单发起人 | 无 |
只承认上表 7 个 OA 事件码。CLI 为每个事件发送 `ruleType=all`、`filterRule={}` 的独立订阅请求;不要添加 `--user`、`--open-dingtalk-id`、`--group`、`--query` 或 `--filter-json`。
只承认上表 6 个 OA 事件码。CLI 为每个事件发送 `ruleType=all`、`filterRule={}` 的独立订阅请求;不要添加 `--user`、`--open-dingtalk-id`、`--group`、`--query` 或 `--filter-json`。
## Intent mapping
@@ -36,14 +35,13 @@ dws auth login
| “审批任务完成时通知我” | `dws event consume user_oa_approval_task_finished --flatten -f ndjson` |
| “审批任务被转交时通知我” | `dws event consume user_oa_approval_task_redirected --flatten -f ndjson` |
| “有审批单发起时通知我” | `dws event consume user_oa_approval_instance_started --flatten -f ndjson` |
| “有审批抄送给我时通知我” | `dws event consume user_oa_approval_instance_cc --flatten -f ndjson` |
| “有审批单终止时通知我” | `dws event consume user_oa_approval_instance_terminated --flatten -f ndjson` |
| “监听我发起的审批何时完成” / “审批实例完成时通知我” | `dws event consume user_oa_approval_instance_finished --flatten -f ndjson` |
| “同时监听全部已公开 OA 事件” | 一个 consume 放入七个 OA event key,不加目标或过滤参数 |
| “同时监听全部已公开 OA 事件” | 一个 consume 放入六个 OA event key,不加目标或过滤参数 |
| “查看 OA 事件目录” | `dws event list --category oa` |
| “查看 OA 事件输出字段” | 对对应事件运行 `dws event schema <event_key> --flatten` |
三个审批任务事件分别表达任务已创建、已完成和已转交;四个审批实例事件分别表达实例已发起、到达抄送节点、已终止和已完成。`status` 和 `result` 保留服务端原值,不把当前样本值推断为完整枚举。
三个审批任务事件分别表达任务已创建、已完成和已转交;三个审批实例事件分别表达实例已发起、已终止和已完成。扁平字段来自六类事件的预发联调样本;`status` 和 `result` 保留服务端原值,不把当前样本值推断为完整枚举。
## Commands
@@ -54,7 +52,6 @@ dws event schema user_oa_approval_task_created --flatten
dws event schema user_oa_approval_task_finished --flatten
dws event schema user_oa_approval_task_redirected --flatten
dws event schema user_oa_approval_instance_started --flatten
dws event schema user_oa_approval_instance_cc --flatten
dws event schema user_oa_approval_instance_terminated --flatten
dws event schema user_oa_approval_instance_finished --flatten
```
@@ -66,12 +63,11 @@ dws event consume user_oa_approval_task_created --flatten -f ndjson
dws event consume user_oa_approval_task_finished --flatten -f ndjson
dws event consume user_oa_approval_task_redirected --flatten -f ndjson
dws event consume user_oa_approval_instance_started --flatten -f ndjson
dws event consume user_oa_approval_instance_cc --flatten -f ndjson
dws event consume user_oa_approval_instance_terminated --flatten -f ndjson
dws event consume user_oa_approval_instance_finished --flatten -f ndjson
```
同时监听七种事件:
同时监听六种事件:
```bash
dws event consume \
@@ -79,14 +75,13 @@ dws event consume \
user_oa_approval_task_finished \
user_oa_approval_task_redirected \
user_oa_approval_instance_started \
user_oa_approval_instance_cc \
user_oa_approval_instance_terminated \
user_oa_approval_instance_finished \
--flatten \
-f ndjson
```
多事件 consume 会为七个 event key 分别创建订阅和逻辑 consumer,并共享当前组织的 personal bus、远程连接、stdout 和生命周期。不要给 OA 命令加 `--query` 或 `--filter-json`;这两个 flag 只用于兼容的 IM 消息接收事件。
多事件 consume 会为六个 event key 分别创建订阅和逻辑 consumer,并共享当前组织的 personal bus、远程连接、stdout 和生命周期。不要给 OA 命令加 `--query` 或 `--filter-json`;这两个 flag 只用于兼容的 IM 消息接收事件。
## Output contract
@@ -110,7 +105,7 @@ dws event consume \
- `type` 是当前 event key;`event_id` 可用于去重;`timestamp` 是 transport 事件发生时间;`subscribe_id` 标识对应的独立订阅。
- `process_instance_id` 是审批实例 ID,可传给 OA 审批命令的 `--instance-id`;`process_code` 是审批流程模板编码。
- `create_time`、`finish_time` 和 `event_time` 都是毫秒时间戳。`event_time` 是审批业务事件时间,`timestamp` 是 transport 事件时间。
- 七类事件的额外字段如下;具体事件始终以 `dws event schema <event_key> --flatten` 为准。
- 六类事件的额外字段如下;具体事件始终以 `dws event schema <event_key> --flatten` 为准。
| 事件 | 额外顶层字段 |
|---|---|
@@ -118,7 +113,6 @@ dws event consume \
| `user_oa_approval_task_finished` | `task_id`、`result`、`finish_time` |
| `user_oa_approval_task_redirected` | `task_id`、`result`、`finish_time` |
| `user_oa_approval_instance_started` | 无 |
| `user_oa_approval_instance_cc` | 无 |
| `user_oa_approval_instance_terminated` | `finish_time` |
| `user_oa_approval_instance_finished` | `result`、`finish_time` |
@@ -150,6 +144,6 @@ dws event consume \
## Lifecycle
- 单事件等待 `[event] ready event_key=<key> bus_pid=<pid> subscribe_id=<id>`。
- 七事件先保存七条 `[event] subscription event_key=<key> subscribe_id=<id>`,再等待 `[event] ready event_count=7 bus_pid=<pid>`。
- 六事件先保存六条 `[event] subscription event_key=<key> subscribe_id=<id>`,再等待 `[event] ready event_count=6 bus_pid=<pid>`。
- 临时验证使用 `--max-events 1` 或 `--duration 10m`;任务完成后优雅结束 consume,本次新建的订阅会自动取消。
- 外部停止已有订阅时先运行 `dws event stop <subscribe_id> --dry-run`,确认后再加 `--yes`。不要 `kill -9`,否则会跳过自动退订。
-4
View File
@@ -1190,10 +1190,6 @@ Agent 安装 dws skill 后,仅依据 skill 提供的参考文档,将自然
- Prompt: 有和我相关的审批实例发起时实时通知我
- Expected: `dws event consume user_oa_approval_instance_started --flatten -f ndjson`
**event_event_consume_oa_instance_cc_001**
- Prompt: 有审批实例抄送给我时实时通知我
- Expected: `dws event consume user_oa_approval_instance_cc --flatten -f ndjson`
**event_event_consume_oa_instance_terminated_001**
- Prompt: 和我相关的审批实例终止时实时通知我
- Expected: `dws event consume user_oa_approval_instance_terminated --flatten -f ndjson`
+2 -3
View File
@@ -280,7 +280,7 @@ func TestEventSkillFrontmatterAdvertisesGroupMemberLifecycle(t *testing.T) {
"群成员加入",
"群成员退出",
"审批任务创建/完成/转交",
"审批实例发起/抄送/终止/完成",
"审批实例发起/终止/完成",
} {
if !strings.Contains(frontmatter, required) {
t.Errorf("%s frontmatter missing event discovery trigger %q", path, required)
@@ -306,7 +306,7 @@ func TestStandaloneEventSkillOwnsAllPersonalEventContracts(t *testing.T) {
"<!-- dws-intent: event.listen.im -->",
"<!-- dws-intent: event.listen.oa -->",
"16 个 EventKey",
"23 个公开个人 EventKey",
"22 个公开个人 EventKey",
} {
if !strings.Contains(string(skillContent), required) {
t.Errorf("%s missing standalone event contract %q", skillPath, required)
@@ -360,7 +360,6 @@ func TestStandaloneEventSkillOwnsAllPersonalEventContracts(t *testing.T) {
"user_oa_approval_task_finished",
"user_oa_approval_task_redirected",
"user_oa_approval_instance_started",
"user_oa_approval_instance_cc",
"user_oa_approval_instance_terminated",
"user_oa_approval_instance_finished",
}
+148
View File
@@ -0,0 +1,148 @@
package safechat
/*
#include "csrc/safechat.h"
#include "goproxy_bridge.h"
#include <stdlib.h>
*/
import "C"
import (
"sync/atomic"
"unsafe"
)
var globalClient atomic.Value // stores *Client
// registerGlobalClient stores the client reference for CGO callbacks.
func registerGlobalClient(c *Client) {
globalClient.Store(c)
}
// getGlobalClient retrieves the active client from the atomic value.
// Returns nil if no client is registered.
func getGlobalClient() *Client {
v := globalClient.Load()
if v == nil {
return nil
}
return v.(*Client)
}
// goProxy is the CGO callback invoked by the C library when a key request needs to be sent.
//
//export goProxy
func goProxy(corpid, uid, domain, url, param, seqID *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
// Copy C strings to Go values immediately to avoid dangling pointer issues.
// The C layer may free these buffers after goProxy returns.
goCorpID := C.GoString(corpid)
goDomain := C.GoString(domain)
goURL := C.GoString(url)
goParam := C.GoString(param)
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy called: corpID=%s, url=%s", goCorpID, goURL)
client.cfg.Logger.Debug("goProxy input param (length=%d): %s", len(goParam), previewString(goParam, 1024))
}
code := client.cfg.Code
if client.cfg.AuthCodeHook != nil {
hookCode, hookErr := client.cfg.AuthCodeHook(goCorpID, goDomain)
if hookErr != nil {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: authCode hook failed for corp %s: %v", goCorpID, hookErr)
}
return -1
}
if hookCode == "" {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: authCode hook returned an empty code for corp %s", goCorpID)
}
return -1
}
code = hookCode
}
// Use kcMu to protect the HTTP request - NOT c.mu!
// c.mu is already held by the caller (EncryptMsg/DecryptMsg etc.)
client.kcMu.Lock()
resp, err := client.kc.doKeyRequest(goURL, goParam, code)
client.kcMu.Unlock()
if err != nil {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: key request failed for corp %s: %v", goCorpID, err)
}
return -1
}
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy: HTTP request succeeded, feeding response to C setResponse (corpID=%s, response length=%d): %s",
goCorpID, len(resp), previewString(resp, 1024))
}
cCorpID := C.CString(goCorpID)
cResp := C.CString(resp)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cResp))
ret := C.setResponse(
cCorpID,
cResp,
(C.block_crypto_func)(C.goBlockBridge),
)
if client.cfg.Logger != nil {
client.cfg.Logger.Debug("goProxy: C.setResponse returned %d for corpID=%s", ret, goCorpID)
}
if ret != 0 {
if client.cfg.Logger != nil {
client.cfg.Logger.Error("goProxy: C.setResponse FAILED for corpID=%s, ret=%d (key file may NOT have been generated)", goCorpID, ret)
}
return -1
}
return 0
}
// goBlock is called by C library when an enterprise key becomes restricted.
// This sets a flag that can be checked by the Go application layer.
//
//export goBlock
func goBlock(corpid *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
goCorpID := C.GoString(corpid)
if client.cfg.Logger != nil {
client.cfg.Logger.Info("goBlock: enterprise %s key restricted", goCorpID)
}
// Store blocked status - application can check via IsBlocked()
client.blockedCorps.Store(goCorpID, true)
return 0
}
// goCancelBlock is called when an enterprise key restriction is lifted.
//
//export goCancelBlock
func goCancelBlock(corpid *C.char) C.int {
client := getGlobalClient()
if client == nil {
return -1
}
goCorpID := C.GoString(corpid)
if client.cfg.Logger != nil {
client.cfg.Logger.Info("goCancelBlock: enterprise %s key unblocked", goCorpID)
}
client.blockedCorps.Delete(goCorpID)
return 0
}
+9
View File
@@ -0,0 +1,9 @@
//go:build darwin && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT
#cgo LDFLAGS: -L${SRCDIR}/lib/darwin_amd64 -lsafechat -lpthread -ldl -lm -framework Security
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build darwin && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT
#cgo LDFLAGS: -L${SRCDIR}/lib/darwin_arm64 -lsafechat -lpthread -ldl -lm -framework Security
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build linux && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_GNU_SOURCE
#cgo LDFLAGS: -L${SRCDIR}/lib/linux_amd64 -lsafechat -lpthread -ldl -lm
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build linux && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_GNU_SOURCE
#cgo LDFLAGS: -L${SRCDIR}/lib/linux_arm64 -lsafechat -lpthread -ldl -lm
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build windows && amd64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_WIN32 -DWIN32
#cgo LDFLAGS: ${SRCDIR}/lib/windows_amd64/libsafechat.a -lws2_32 -lgdi32 -lcrypt32 -ladvapi32 -luser32
*/
import "C"
+9
View File
@@ -0,0 +1,9 @@
//go:build windows && arm64
package safechat
/*
#cgo CFLAGS: -I${SRCDIR}/csrc -I${SRCDIR}/csrc/include -DDLL_EXPORT -D_WIN32 -DWIN32
#cgo LDFLAGS: -L${SRCDIR}/lib/windows_arm64 -lsafechat -lws2_32 -lgdi32 -lcrypt32 -ladvapi32 -luser32
*/
import "C"
+88
View File
@@ -0,0 +1,88 @@
package safechat
import (
"time"
"github.com/google/uuid"
)
// Logger defines the logging interface for the SDK.
// Implementations should be safe for concurrent use.
type Logger interface {
// Debug logs a debug-level message
Debug(msg string, args ...interface{})
// Info logs an info-level message
Info(msg string, args ...interface{})
// Error logs an error-level message
Error(msg string, args ...interface{})
}
// Config holds configuration parameters for the SafeChat client.
type Config struct {
// DataPath is the directory path for storing keys and related data (required).
// The directory must exist and be writable.
DataPath string
// UserID is an optional user identifier. If empty, a random UUID is
// generated automatically. The C library stores this value but does not
// use it for key operations, so any non-empty string works.
UserID string
// Code is a DingTalk 免登 authCode used for key server authentication
// when AuthCodeHook is nil. Prefer AuthCodeHook: the code is one-shot
// and should be minted only when goProxy actually needs a key.
//
// Required (via Code or AuthCodeHook) when:
// - First-time key fetch (empty keystore)
// - Server-side key version rotation
//
// If keys are already cached AND the key version has not changed,
// neither field is used (no network request is made).
Code string
// AuthCodeHook is called from goProxy immediately before the key
// request. corpID and domain come from the C library callback.
// domain is the vendor SSO host and must not be forwarded as an
// authorize redirectURI; it is only for local host checks.
// If set, it replaces Code for that request and the returned value
// is not stored on the client.
AuthCodeHook func(corpID, domain string) (string, error)
// KeyServer is an optional override for the key server URL.
// If empty, the URL provided by the C library's goProxy callback will be used.
KeyServer string
// MaxRetry is the maximum number of retry attempts when key is not yet available.
// Default: 5 (private server redirect consumes 2 retries, so 5 provides sufficient margin)
MaxRetry int
// HTTPTimeout is the timeout for HTTP key requests.
// Default: 10s
HTTPTimeout time.Duration
// Logger is an optional logger instance.
// If nil, no logging will be performed.
Logger Logger
}
// defaultConfig returns a Config with sensible defaults applied.
func defaultConfig(cfg Config) Config {
if cfg.MaxRetry <= 0 {
cfg.MaxRetry = 5
}
if cfg.HTTPTimeout <= 0 {
cfg.HTTPTimeout = 10 * time.Second
}
if cfg.UserID == "" {
cfg.UserID = uuid.New().String()
}
return cfg
}
// validate checks that required config fields are set.
func (cfg *Config) validate() error {
if cfg.DataPath == "" {
return ErrConfigDataPathEmpty
}
return nil
}
+137
View File
@@ -0,0 +1,137 @@
#ifndef SAFECHAT_H
#define SAFECHAT_H
#ifdef WIN32
#ifndef DLL_EXPORT
#define CREATEDLL_API __declspec(dllimport)
#else
#define CREATEDLL_API __declspec(dllexport)
#endif
#else
#define CREATEDLL_API __attribute__((visibility("default")))
#endif
#include <stdint.h>
#ifdef __cplusplus
extern "C" {
#endif
typedef unsigned char byte;
// using namespace std;
// #pragma comment(lib, "safechat.lib")
#define DTMsgRetOK 0 /*operate return success*/
#define WARNING_RET_CODE 0x10000000
enum { RET_CODE_SUCCESS = 0, RET_CODE_ERROR, RET_CODE_WARNING };
// 获取返回值的类型,RET_CODE_WARNING | RET_CODE_ERROR| RET_CODE_SUCCESS
// #define GET_RET_CODETYPE(x) (x == 0) ? RET_CODE_SUCCESS : ((-x) & WARNING_RET_CODE == WARNING_RET_CODE ?
// RET_CODE_WARNING : RET_CODE_ERROR)
/*
打印日志回调函数
@param:ret_type 打印日志类型,RET_CODE_WARNING | RET_CODE_ERROR| RET_CODE_SUCCESS
@param:msg_detail 打印日志的详细信息,char * UTF-8格式
@return 备用,暂且返回NULL
*/
typedef void *(*log_func)(int ret_type, char *msg_detail);
/*
代理回调函数
@param: corpid 企业id
@param: uid 当前用户id
@param: domain 请求的域名
@param: url 请求url
@param: param 请求参数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
typedef int (*call_proxy_func)(char *corpid, char *uid, char *domain, char *url, char *param,
char *seq_id);
/*
设置密钥管控标记
@param: corpid 企业id
*/
typedef int (*block_crypto_func)(char *corpid);
/*
解除密钥管控标记
@param: corpid 企业id
*/
typedef int (*cancel_block_crypto_func)(char *corpid);
/*
初始化函数
@param: path 用户数据路径, 保存key等信息
@param: my_id 当前登W录用户的id
@return: 返回执行结果
*/
CREATEDLL_API int safechatInit(char *path, char *my_id);
/*
数据加密函数
@param: corp_id 企业id
@param: staffid 员工id
@param: data_buf 加密前数据
@param: data_len data_buf长度
@param: id 对方用户id, 如果群消息设置为NULL
@param: encrypt_buf 加密后的数据指针,需要外部函数释放
@param: ret_len 加密数据的返回长度
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int encryptData(char *corp_id, char *staffid, unsigned char *data_buf, unsigned int data_len, char *id,
unsigned char **encrypt_buf, unsigned int *ret_len, char *seq_id, call_proxy_func proxy);
/*
数据解密函数
@param: corp_id 企业id
@param: staffid 员工id
@param: data_buf 解密前数据
@param: data_len data_buf长度
@param: id 对方用户id, 如果群消息设置为NULL
@param: decrypt_buf 解密后的数据指针,需要外部函数释放
@param: ret_len 解密数据的返回长度
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int decryptData(char *corp_id, char *staffid, unsigned char *data_buf, unsigned int data_len, char *id,
unsigned char **decrypt_buf, unsigned int *ret_len, char *seq_id, call_proxy_func proxy);
CREATEDLL_API int encryptFile(char *corp_id, char *staffid, char *id, char *in_file_path, char *out_file_path,
char *seq_id, call_proxy_func proxy);
CREATEDLL_API int encryptBuffer(char *corp_id, char *staffid, byte *data_buf, uint32_t data_len, char *id,
byte **encrypt_buf, uint32_t *ret_len, char *seq_id, call_proxy_func proxy);
CREATEDLL_API int decryptFile(char *corp_id, char *staffid, char *id, char *in_file_path, char *out_file_path,
char *seq_id, call_proxy_func proxy);
CREATEDLL_API int decryptBuffer(char *corp_id, char *staffid, byte *data_buf, uint32_t data_len, char *id,
byte **decrypt_buf, uint32_t *ret_len, char *seq_id, call_proxy_func proxy);
/*
服务器响应处理
@param: corp_id 企业id
@param: json_str json格式的服务器返回值
@return: 返回执行结果
*/
CREATEDLL_API int setResponse(char *corp_id, char *json_str, block_crypto_func block_func1);
/*
消息推送处理
@param: corp_id 企业id
@param: staffid 员工id
@param: push_data 发送请求的数据
@param: proxy 代理请求回调函数
@return: 返回执行结果
@param: seq_id 回调函数线程识别参数
*/
CREATEDLL_API int setPushData(char *corp_id, char *staffid, char *push_data, char *seq_id, call_proxy_func proxy,
cancel_block_crypto_func cancel_crypto_func);
#ifdef __cplusplus
}
#endif
#endif // SAFECHAT_H
+265
View File
@@ -0,0 +1,265 @@
package safechat
import (
"errors"
"fmt"
)
// SDK-level errors (Go side)
var (
// ErrConfigDataPathEmpty indicates DataPath is not set in Config.
ErrConfigDataPathEmpty = errors.New("safechat: config.DataPath is required")
// ErrConfigUserIDEmpty indicates UserID is not set in Config.
ErrConfigUserIDEmpty = errors.New("safechat: config.UserID is required")
// ErrNotInitialized indicates the client has not been initialized.
ErrNotInitialized = errors.New("safechat: client not initialized")
// ErrMaxRetryExceeded indicates max retry attempts for key acquisition exceeded.
ErrMaxRetryExceeded = errors.New("safechat: max retry exceeded, key not available")
// ErrKeyRestricted indicates the enterprise key is restricted (managed/blocked).
ErrKeyRestricted = errors.New("safechat: enterprise key is restricted by admin")
// ErrKeyRequestFailed indicates the HTTP key request failed.
ErrKeyRequestFailed = errors.New("safechat: key request HTTP call failed")
// ErrAlreadyInitialized indicates init was called more than once.
ErrAlreadyInitialized = errors.New("safechat: client already initialized")
)
// CError represents an error code returned from the C library layer.
type CError struct {
Code int
Message string
}
func (e *CError) Error() string {
return fmt.Sprintf("safechat: C library error %d: %s", e.Code, e.Message)
}
// mapCError maps a C library return code to a Go error.
// Returns nil if code == 0 (FUNCTION_OK).
func mapCError(code int) error {
if code == 0 {
return nil
}
msg, ok := cErrorMessages[code]
if !ok {
msg = "unknown error"
}
return &CError{Code: code, Message: msg}
}
// C library error code constants - mapped from native.h
const (
// Special status codes
cFunctionOK = 0
cSendRequestParamOK = -15002 // Key not found, request sent via goProxy
// Init errors (-20000 ~ -20016)
cInitParamPathNull = -20000
cInitParamMyIDNull = -20001
cInitPathUTF8Null = -20002
cInitPathNotAvailable = -20003
cInitLogInitError = -20004
cInitEncryptInitError = -20005
cInitPCNativeError = -20006
cInitReadLocalKey = -20013
cInitReadKey = -20014
// EncryptData errors (-20100 ~ -20114)
cEncryptDataCorpIDNull = -20100
cEncryptDataMsgNull = -20101
cEncryptDataBufNull = -20102
cEncryptDataRetLenNull = -20103
cEncryptDataProxyNull = -20104
cEncryptDataLenError = -20105
cEncryptDataKeyNegtive = -20111
cEncryptDataTmpKeyNull = -20109
cEncryptDataBuildReqErr = -20110
// DecryptData errors (-20200 ~ -20215)
cDecryptDataCorpIDNull = -20200
cDecryptDataMsgNull = -20201
cDecryptDataBufNull = -20202
cDecryptDataRetLenNull = -20203
cDecryptDataProxyNull = -20204
cDecryptDataLenError = -20205
cDecryptDataFormatError = -20206
cDecryptDataDecryptErr = -20211
cDecryptDataKeyNegtive = -20214
cDecryptDataTmpKeyNull = -20212
cDecryptDataBuildReqErr = -20213
// EncryptFile errors (-20300 ~ -20339)
cEncryptFileCorpIDNull = -20300
cEncryptFilePathNull = -20301
cEncryptFileTmpKeyNull = -20315
cEncryptFileBuildReqErr = -20316
// EncryptBuffer errors (-20400 ~ -20416)
cEncryptBufferCorpIDNull = -20400
cEncryptBufferTmpKeyNull = -20414
cEncryptBufferBuildReqErr = -20415
// DecryptFile errors (-20500 ~ -20538)
cDecryptFileCorpIDNull = -20500
cDecryptFilePathNull = -20501
cDecryptFileHeadError = -20505
cDecryptFileHashError = -20517
cDecryptFileTmpKeyNull = -20518
cDecryptFileBuildReqErr = -20519
// DecryptBuffer errors (-20600 ~ -20622)
cDecryptBufferCorpIDNull = -20600
cDecryptBufferHeadNull = -20608
cDecryptBufferHashError = -20618
cDecryptBufferTmpKeyNull = -20619
cDecryptBufferBuildReqErr = -20620
// SetResponse errors (-20700 ~ -20730)
cSetResponseCorpIDNull = -20700
cSetResponseJSONNull = -20701
cSetResponseKeyNagtive = -20728
cSetResponseSaveKeyErr = -20723
// SetPushData errors (-20800 ~ -20814)
cSetPushDataCorpIDNull = -20800
cSetPushDataTypeUndef = -20814
// V3 signing errors (-31001 ~ -34002)
cV3BuildReq3SM2Failed = -31006
cV3BuildReq3SignError = -31009
cV3ParseResp3Downgrade = -32012
)
// cErrorMessages maps C error codes to human-readable messages.
var cErrorMessages = map[int]string{
// Init
-20000: "init: path parameter is NULL",
-20001: "init: my_id parameter is NULL",
-20002: "init: path UTF8 conversion returned NULL",
-20003: "init: path is not available/accessible",
-20004: "init: log initialization failed",
-20005: "init: encryption engine initialization failed",
-20006: "init: pcnative initialization failed",
-20007: "init: get URL full path error",
-20013: "init: read local key failed",
-20014: "init: read key failed",
-20015: "init: my_id parameter is NULL",
-20016: "init: logFunc parameter is NULL",
// EncryptData
-20100: "encryptData: corp_id is NULL or empty",
-20101: "encryptData: message content is NULL",
-20102: "encryptData: encrypt_buf is NULL",
-20103: "encryptData: ret_len is NULL",
-20104: "encryptData: proxy function is NULL",
-20105: "encryptData: data length invalid (too small or too big)",
-20106: "encryptData: encryptMsgHelper parameter error",
-20107: "encryptData: SM4 encryption failed",
-20108: "encryptData: base64 encoding failed",
-20109: "encryptData: key request already in progress",
-20110: "encryptData: build key request failed",
-20111: "encryptData: enterprise key is restricted",
-20113: "encryptData: malloc encode buffer failed",
-20114: "encryptData: malloc encrypt buffer failed",
// DecryptData
-20200: "decryptData: corp_id is NULL or empty",
-20201: "decryptData: message content is NULL",
-20202: "decryptData: decrypt_buf is NULL",
-20203: "decryptData: ret_len is NULL",
-20204: "decryptData: proxy function is NULL",
-20205: "decryptData: data length invalid",
-20206: "decryptData: message content format error (missing ||separators)",
-20207: "decryptData: decryptMsgHelper parameter error",
-20208: "decryptData: base64 decode failed",
-20209: "decryptData: decode length error",
-20210: "decryptData: decrypt buffer malloc failed",
-20211: "decryptData: SM4 decryption failed",
-20212: "decryptData: key request already in progress",
-20213: "decryptData: build key request failed",
-20214: "decryptData: enterprise key is restricted",
// EncryptFile
-20300: "encryptFile: corp_id is NULL or empty",
-20301: "encryptFile: source or dest file path is NULL",
-20302: "encryptFile: id parameter is NULL",
-20303: "encryptFile: seq_id parameter is NULL",
-20304: "encryptFile: proxy function is NULL",
-20305: "encryptFile: key_info is NULL",
-20308: "encryptFile: get file size failed",
-20309: "encryptFile: open source or dest file failed",
-20313: "encryptFile: encrypt block failed",
-20315: "encryptFile: key request already in progress",
-20316: "encryptFile: build key request failed",
-20333: "encryptFile: create thread failed",
-20339: "encryptFile: test environment error",
// EncryptBuffer
-20400: "encryptBuffer: corp_id is NULL or empty",
-20401: "encryptBuffer: file content is NULL",
-20407: "encryptBuffer: proxy function is NULL",
-20413: "encryptBuffer: encrypted block error",
-20414: "encryptBuffer: key request already in progress",
-20415: "encryptBuffer: build key request failed",
// DecryptFile
-20500: "decryptFile: corp_id is NULL or empty",
-20501: "decryptFile: source or dest file path is NULL",
-20505: "decryptFile: file header format error",
-20509: "decryptFile: header magic error",
-20511: "decryptFile: open file failed",
-20515: "decryptFile: decrypt block failed",
-20517: "decryptFile: file hash verification failed (warning)",
-20518: "decryptFile: key request already in progress",
-20519: "decryptFile: build key request failed",
-20520: "decryptFile: file length error",
// DecryptBuffer
-20600: "decryptBuffer: corp_id is NULL or empty",
-20608: "decryptBuffer: parseHeadInfo returned NULL",
-20615: "decryptBuffer: header magic error",
-20617: "decryptBuffer: decrypt block error",
-20618: "decryptBuffer: file hash verification failed",
-20619: "decryptBuffer: key request already in progress",
-20620: "decryptBuffer: build key request failed",
-20621: "decryptBuffer: file length error",
// SetResponse
-20700: "setResponse: corp_id is NULL or empty",
-20701: "setResponse: json_str is NULL or empty",
-20703: "setResponse: find tmp corp no tmp key found",
-20704: "setResponse: JSON parse failed",
-20716: "setResponse: server request error",
-20722: "setResponse: localKeyPath is NULL",
-20723: "setResponse: save key failed",
-20728: "setResponse: enterprise key is restricted (nagtive)",
// SetPushData
-20800: "setPushData: corp_id is NULL",
-20801: "setPushData: push_data is NULL",
-20804: "setPushData: JSON parse failed",
-20805: "setPushData: type field is NULL",
-20806: "setPushData: find_create_tmp_key failed",
-20814: "setPushData: unknown push type",
// V3 Signing
-31001: "v3: build_key_request3 arg is NULL",
-31006: "v3: SM2 encrypt R2 failed",
-31009: "v3: sign step one error",
-32005: "v3: deserialize response failed",
-32010: "v3: new SDK request old private error",
-32012: "v3: public server downgrade",
-32022: "v3: MAC of response mismatch",
-32029: "v3: verify server sign failed",
// Misc
-34001: "v3: setResponse not found tmp sign",
-34002: "v3: setResponse malloc for ret buf failed",
}
+463
View File
@@ -0,0 +1,463 @@
// Package main demonstrates how to integrate the SafeChat Go SDK
// into the DingTalk Workspace CLI or any other Go application.
//
// Two usage modes:
//
// 1. Single-action mode (default)
// Run one of: encrypt-msg / decrypt-msg / encrypt-file / decrypt-file /
// encrypt-buf / decrypt-buf via -action flag.
//
// 2. Full-test mode (-test-all)
// Exercise all 3 encryption APIs (Msg / File / Buffer) in one run,
// verify round-trip (decrypt == original), and print a summary.
//
// Key cache files
//
// The SDK persists negotiated keys under -data directory:
// ahflag_256.store local key identifier
// ahkey_256.store encrypted key material
//
// Build:
//
// go build -o safechat-example ./example/
package main
import (
"bytes"
"crypto/rand"
"flag"
"fmt"
"log"
"os"
"path/filepath"
"strings"
"time"
safechat "safechat-go-sdk"
)
func main() {
var (
dataPath = flag.String("data", "./keystore", "Path to key storage directory (contains ahflag_256.store / ahkey_256.store)")
userID = flag.String("user", "", "Current login user ID")
code = flag.String("code", "", "DingTalk authCode (only needed when keys must be fetched from server)")
corpID = flag.String("corp", "", "Enterprise/corp ID (required)")
staffID = flag.String("staff", "", "Staff ID for encryption target (defaults to -user)")
action = flag.String("action", "", "Single action: encrypt-msg|decrypt-msg|encrypt-file|decrypt-file|encrypt-buf|decrypt-buf")
input = flag.String("input", "", "Input text / file path")
output = flag.String("output", "", "Output file path (for file operations)")
server = flag.String("server", "", "Key server URL override (optional)")
testAll = flag.Bool("test-all", false, "Run a full round-trip test of all 3 encryption APIs using cached keys")
verbose = flag.Bool("v", false, "Verbose logging")
)
flag.Parse()
if *corpID == "" {
fmt.Fprintln(os.Stderr, "Error: -corp flag is required")
flag.Usage()
os.Exit(1)
}
// Validate key cache files up-front so the user gets a clear message
// instead of a cryptic C-library error later.
checkKeyCache(*dataPath, *code)
cfg := safechat.Config{
DataPath: *dataPath,
UserID: *userID,
Code: *code,
MaxRetry: 5,
HTTPTimeout: 15 * time.Second,
}
if !*verbose {
cfg.Logger = &quietLogger{}
} else {
cfg.Logger = &stdLogger{}
}
if *server != "" {
cfg.KeyServer = *server
}
client, err := safechat.New(cfg)
if err != nil {
log.Fatalf("Failed to initialize SafeChat client: %v", err)
}
defer client.Close()
targetStaff := *staffID
if targetStaff == "" {
targetStaff = *userID
}
if *testAll {
runFullTest(client, *corpID, targetStaff)
return
}
if *action == "" {
fmt.Fprintln(os.Stderr, "Error: either -action <name> or -test-all must be provided")
flag.Usage()
os.Exit(1)
}
switch *action {
case "encrypt-msg":
encryptMessage(client, *corpID, targetStaff, *input)
case "decrypt-msg":
decryptMessage(client, *corpID, targetStaff, *input)
case "encrypt-file":
encryptFile(client, *corpID, targetStaff, *input, *output)
case "decrypt-file":
decryptFile(client, *corpID, targetStaff, *input, *output)
case "encrypt-buf":
encryptBuffer(client, *corpID, targetStaff, *input)
case "decrypt-buf":
decryptBuffer(client, *corpID, targetStaff, *input)
default:
fmt.Fprintf(os.Stderr, "Unknown action: %s\n", *action)
os.Exit(1)
}
}
// ---------- key-cache pre-check ----------
// checkKeyCache prints a friendly hint about which keys are available and
// whether a network round-trip to the key server should be expected.
func checkKeyCache(dataPath, code string) {
flagFile := filepath.Join(dataPath, "ahflag_256.store")
keyFile := filepath.Join(dataPath, "ahkey_256.store")
_, err1 := os.Stat(flagFile)
_, err2 := os.Stat(keyFile)
hasFlag := err1 == nil
hasKey := err2 == nil
switch {
case hasFlag && hasKey:
fmt.Printf("[key-cache] found %s + %s (existence check only; C layer still validates hash/version)\n",
filepath.Base(flagFile), filepath.Base(keyFile))
fmt.Printf("[key-cache] NOTE: store files are NOT portable across CPU architectures — " +
"calc_hash() depends on char signedness (signed on x86_64, unsigned on AArch64). " +
"A foreign store fails the hash check, gets deleted and re-generated, which triggers a server call.\n")
case hasFlag || hasKey:
fmt.Fprintf(os.Stderr,
"[key-cache] WARNING: only one of the pair exists (%s / %s); key negotiation will likely fail\n",
filepath.Base(flagFile), filepath.Base(keyFile))
default:
if code == "" {
fmt.Fprintf(os.Stderr,
"[key-cache] no cached keys in %s and -code is empty; "+
"either copy ahflag_256.store + ahkey_256.store into that dir, "+
"or provide a valid DingTalk access_token via -code\n",
dataPath)
} else {
fmt.Printf("[key-cache] no cached keys in %s; will try to fetch from server using -code\n", dataPath)
}
}
}
// ---------- full round-trip test of all 3 APIs ----------
// runFullTest exercises the 3 Go API pairs defined in safechat.go.
//
// The APIs under test are the PUBLIC Go methods on *safechat.Client:
//
// API #1 — Msg API : EncryptMsg / DecryptMsg (safechat.go L115, L154)
// API #2 — File API : EncryptFile / DecryptFile (safechat.go L188, L213)
// API #3 — Buffer API : EncryptBuffer / DecryptBuffer (safechat.go L240, L268)
//
// These are pure Go entry points: parameter validation, mutex locking,
// MaxRetry loop, and Go error translation are all done in safechat.go.
// The underlying CGO bridge (cEncryptData / cDecryptData / ...) is a
// PRIVATE implementation detail and is NOT what this test targets.
func runFullTest(client *safechat.Client, corpID, staffID string) {
fmt.Println("==========================================================")
fmt.Printf(" SafeChat Go SDK — Go API round-trip test\n")
fmt.Printf(" corpID=%s staffID=%s\n", corpID, staffID)
fmt.Println(" Target: the 3 public Go API pairs in safechat.go")
fmt.Println(" #1 Msg API : EncryptMsg / DecryptMsg")
fmt.Println(" #2 File API : EncryptFile / DecryptFile")
fmt.Println(" #3 Buffer API : EncryptBuffer / DecryptBuffer")
fmt.Println("==========================================================")
passed, failed := 0, 0
// Go API #1 — Msg
fmt.Println("\n[1/3] Go API #1 : EncryptMsg / DecryptMsg (safechat.go L115, L154)")
if runCase("Msg API", func() error { return testMsgAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
// Go API #2 — File
fmt.Println("\n[2/3] Go API #2 : EncryptFile / DecryptFile (safechat.go L188, L213)")
if runCase("File API", func() error { return testFileAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
// Go API #3 — Buffer
fmt.Println("\n[3/3] Go API #3 : EncryptBuffer / DecryptBuffer (safechat.go L240, L268)")
if runCase("Buffer API", func() error { return testBufferAPI(client, corpID, staffID) }) {
passed++
} else {
failed++
}
fmt.Println("\n==========================================================")
fmt.Printf(" RESULT: %d passed, %d failed\n", passed, failed)
fmt.Println("==========================================================")
if failed > 0 {
os.Exit(1)
}
}
func runCase(name string, fn func() error) bool {
if err := fn(); err != nil {
fmt.Printf(" ✗ %s FAILED: %v\n", name, err)
return false
}
fmt.Printf(" ✓ %s OK\n", name)
return true
}
// testMsgAPI exercises Go API #1:
//
// client.EncryptMsg(corpID, staffID, plain) → ciphertext
// client.DecryptMsg(corpID, staffID, ciphertext) → plaintext
//
// Both methods are defined in safechat.go (L115, L154). They perform
// parameter validation, mutex locking, and a MaxRetry loop before
// returning a Go []byte + error.
func testMsgAPI(c *safechat.Client, corpID, staffID string) error {
plain := make([]byte, 128)
if _, err := rand.Read(plain); err != nil {
return err
}
// Go API call — safechat.go L115
ct, err := c.EncryptMsg(corpID, staffID, plain)
if err != nil {
return fmt.Errorf("EncryptMsg (Go API): %w", err)
}
fmt.Printf(" [Go API] EncryptMsg : plain=%d bytes -> ct=%d bytes\n", len(plain), len(ct))
fmt.Printf(" [ciphertext str] %s\n", strings.ReplaceAll(string(ct), "\n", ""))
// Go API call — safechat.go L154
got, err := c.DecryptMsg(corpID, staffID, ct)
if err != nil {
return fmt.Errorf("DecryptMsg (Go API): %w", err)
}
if !bytes.Equal(got, plain) {
return fmt.Errorf("round-trip mismatch: got %d bytes, want %d", len(got), len(plain))
}
fmt.Printf(" [Go API] DecryptMsg : ct=%d bytes -> plain=%d bytes OK\n", len(ct), len(got))
return nil
}
// testFileAPI exercises Go API #2:
//
// client.EncryptFile(corpID, staffID, src, enc) error
// client.DecryptFile(corpID, staffID, enc, dec) error
//
// Both methods are defined in safechat.go (L188, L213). They operate on
// file paths and return only an error — no data crosses the Go/C boundary
// in the caller-visible API.
func testFileAPI(c *safechat.Client, corpID, staffID string) error {
dir, err := os.MkdirTemp("", "safechat-file-test-*")
if err != nil {
return err
}
defer os.RemoveAll(dir)
src := filepath.Join(dir, "plain.bin")
enc := filepath.Join(dir, "plain.bin.enc")
dec := filepath.Join(dir, "plain.bin.dec")
buf := make([]byte, 1024)
if _, err := rand.Read(buf); err != nil {
return err
}
if err := os.WriteFile(src, buf, 0644); err != nil {
return err
}
// Go API call — safechat.go L188
if err := c.EncryptFile(corpID, staffID, src, enc); err != nil {
return fmt.Errorf("EncryptFile (Go API): %w", err)
}
encSize, _ := fileSize(enc)
fmt.Printf(" [Go API] EncryptFile : src=%d bytes -> enc=%d bytes\n", len(buf), encSize)
// Go API call — safechat.go L213
if err := c.DecryptFile(corpID, staffID, enc, dec); err != nil {
return fmt.Errorf("DecryptFile (Go API): %w", err)
}
decBuf, err := os.ReadFile(dec)
if err != nil {
return err
}
if !bytes.Equal(decBuf, buf) {
return fmt.Errorf("round-trip mismatch: dec=%d bytes, want=%d", len(decBuf), len(buf))
}
fmt.Printf(" [Go API] DecryptFile : enc=%d bytes -> dec=%d bytes OK\n", encSize, len(decBuf))
return nil
}
// testBufferAPI exercises Go API #3:
//
// client.EncryptBuffer(corpID, staffID, data) → []byte, error
// client.DecryptBuffer(corpID, staffID, data) → []byte, error
//
// Both methods are defined in safechat.go (L240, L268). They operate on
// in-memory []byte buffers, suitable for binary protocols or DB blobs.
func testBufferAPI(c *safechat.Client, corpID, staffID string) error {
plain := make([]byte, 512)
if _, err := rand.Read(plain); err != nil {
return err
}
// Go API call — safechat.go L240
ct, err := c.EncryptBuffer(corpID, staffID, plain)
if err != nil {
return fmt.Errorf("EncryptBuffer (Go API): %w", err)
}
fmt.Printf(" [Go API] EncryptBuffer : plain=%d bytes -> ct=%d bytes\n", len(plain), len(ct))
// Go API call — safechat.go L268
got, err := c.DecryptBuffer(corpID, staffID, ct)
if err != nil {
return fmt.Errorf("DecryptBuffer (Go API): %w", err)
}
if !bytes.Equal(got, plain) {
return fmt.Errorf("round-trip mismatch: got %d bytes, want %d", len(got), len(plain))
}
fmt.Printf(" [Go API] DecryptBuffer : ct=%d bytes -> plain=%d bytes OK\n", len(ct), len(got))
return nil
}
func fileSize(p string) (int64, error) {
fi, err := os.Stat(p)
if err != nil {
return 0, err
}
return fi.Size(), nil
}
// ---------- single-action helpers ----------
func encryptMessage(client *safechat.Client, corpID, staffID, plaintext string) {
if plaintext == "" {
plaintext = "Hello, this is a test message from SafeChat Go SDK!"
}
fmt.Printf("Encrypting message: %q\n", plaintext)
ciphertext, err := client.EncryptMsg(corpID, staffID, []byte(plaintext))
if err != nil {
log.Fatalf("EncryptMsg failed: %v", err)
}
fmt.Printf("Encrypted (%d bytes): %s\n", len(ciphertext), string(ciphertext))
}
func decryptMessage(client *safechat.Client, corpID, staffID, ciphertext string) {
if ciphertext == "" {
log.Fatal("decrypt-msg requires -input with the ciphertext")
}
fmt.Printf("Decrypting message (%d bytes)...\n", len(ciphertext))
plaintext, err := client.DecryptMsg(corpID, staffID, []byte(ciphertext))
if err != nil {
log.Fatalf("DecryptMsg failed: %v", err)
}
fmt.Printf("Decrypted: %s\n", string(plaintext))
}
func encryptFile(client *safechat.Client, corpID, staffID, srcPath, dstPath string) {
if srcPath == "" {
log.Fatal("encrypt-file requires -input with source file path")
}
if dstPath == "" {
dstPath = srcPath + ".enc"
}
fmt.Printf("Encrypting file: %s -> %s\n", srcPath, dstPath)
if err := client.EncryptFile(corpID, staffID, srcPath, dstPath); err != nil {
log.Fatalf("EncryptFile failed: %v", err)
}
fmt.Println("File encrypted successfully")
}
func decryptFile(client *safechat.Client, corpID, staffID, srcPath, dstPath string) {
if srcPath == "" {
log.Fatal("decrypt-file requires -input with source file path")
}
if dstPath == "" {
dstPath = srcPath + ".dec"
}
fmt.Printf("Decrypting file: %s -> %s\n", srcPath, dstPath)
if err := client.DecryptFile(corpID, staffID, srcPath, dstPath); err != nil {
log.Fatalf("DecryptFile failed: %v", err)
}
fmt.Println("File decrypted successfully")
}
func encryptBuffer(client *safechat.Client, corpID, staffID, inputPath string) {
if inputPath == "" {
log.Fatal("encrypt-buf requires -input with file path containing data to encrypt")
}
data, err := os.ReadFile(inputPath)
if err != nil {
log.Fatalf("Failed to read input file: %v", err)
}
fmt.Printf("Encrypting buffer (%d bytes)...\n", len(data))
encrypted, err := client.EncryptBuffer(corpID, staffID, data)
if err != nil {
log.Fatalf("EncryptBuffer failed: %v", err)
}
outPath := inputPath + ".enc"
if err := os.WriteFile(outPath, encrypted, 0644); err != nil {
log.Fatalf("Failed to write output: %v", err)
}
fmt.Printf("Buffer encrypted successfully (%d bytes) -> %s\n", len(encrypted), outPath)
}
func decryptBuffer(client *safechat.Client, corpID, staffID, inputPath string) {
if inputPath == "" {
log.Fatal("decrypt-buf requires -input with file path containing data to decrypt")
}
data, err := os.ReadFile(inputPath)
if err != nil {
log.Fatalf("Failed to read input file: %v", err)
}
fmt.Printf("Decrypting buffer (%d bytes)...\n", len(data))
decrypted, err := client.DecryptBuffer(corpID, staffID, data)
if err != nil {
log.Fatalf("DecryptBuffer failed: %v", err)
}
outPath := inputPath + ".dec"
if err := os.WriteFile(outPath, decrypted, 0644); err != nil {
log.Fatalf("Failed to write output: %v", err)
}
fmt.Printf("Buffer decrypted successfully (%d bytes) -> %s\n", len(decrypted), outPath)
}
// ---------- loggers ----------
type stdLogger struct{}
func (l *stdLogger) Debug(format string, args ...interface{}) {
log.Printf("[DEBUG] "+format, args...)
}
func (l *stdLogger) Info(format string, args ...interface{}) {
log.Printf("[INFO] "+format, args...)
}
func (l *stdLogger) Error(format string, args ...interface{}) {
log.Printf("[ERROR] "+format, args...)
}
// quietLogger swallows SDK logs; only errors are surfaced via returned error values.
type quietLogger struct{}
func (l *quietLogger) Debug(format string, args ...interface{}) {}
func (l *quietLogger) Info(format string, args ...interface{}) {}
func (l *quietLogger) Error(format string, args ...interface{}) {}
+5
View File
@@ -0,0 +1,5 @@
module safechat-go-sdk
go 1.21
require github.com/google/uuid v1.6.0
+2
View File
@@ -0,0 +1,2 @@
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
+80
View File
@@ -0,0 +1,80 @@
#include "goproxy_bridge.h"
#include "csrc/safechat.h"
#include <stdlib.h>
/*
* goproxy_bridge.c - CGO callback bridge implementation
*
* This file MUST reside in the Go package root directory so that
* CGO compiles it together with the Go code. It bridges C library
* callback invocations to Go exported functions.
*
* Go exported functions are declared as extern here.
* The bridge functions simply forward calls from C library
* to the Go runtime via CGO mechanism.
*/
/* Go exported function declarations (defined in callback.go) */
extern int goProxy(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id);
extern int goBlock(char *corpid);
extern int goCancelBlock(char *corpid);
/*
* init - Wrapper for safechatInit
* Initializes the SafeChat library.
*/
int init(char *path, char *my_id, void *reserved) {
(void)reserved; /* unused */
return safechatInit(path, my_id);
}
/*
* clearCache - Clears the key cache for a specific enterprise
* Returns 0 on success, non-zero on error.
* Note: This is a placeholder - actual implementation may vary.
*/
int clearCache(char *corpid) {
(void)corpid; /* unused for now */
/* TODO: Implement actual cache clearing if needed */
return 0;
}
/*
* freeCryptoBuf - Frees a buffer allocated by the C library
* The C library allocates buffers with malloc, so we free with free.
*/
void freeCryptoBuf(void *buf) {
if (buf != NULL) {
free(buf);
}
}
/*
* goProxyBridge - Bridge for call_proxy_func typedef
* Called by C library when a key request needs to be sent.
* Forwards to Go's goProxy which performs HTTP request
* and calls setResponse to feed back the key data.
*/
int goProxyBridge(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id) {
return goProxy(corpid, uid, domain, url, param, seq_id);
}
/*
* goBlockBridge - Bridge for block_crypto_func typedef
* Called by C library when enterprise key is restricted.
* Notifies Go layer of the restriction.
*/
int goBlockBridge(char *corpid) {
return goBlock(corpid);
}
/*
* goCancelBlockBridge - Bridge for cancel_block_crypto_func typedef
* Called by C library when enterprise key restriction is lifted.
* Notifies Go layer to remove the restriction.
*/
int goCancelBlockBridge(char *corpid) {
return goCancelBlock(corpid);
}
+28
View File
@@ -0,0 +1,28 @@
#ifndef GOPROXY_BRIDGE_H
#define GOPROXY_BRIDGE_H
/*
* goproxy_bridge.h - CGO callback bridge declarations
*
* These bridge functions are called by the C library (safechat.c)
* and forward to Go exported functions via CGO.
* This indirection is required because CGO cannot directly pass
* Go function pointers to C code.
*/
/* Proxy callback bridge - forwards key requests to Go HTTP client */
int goProxyBridge(char *corpid, char *uid, char *domain,
char *url, char *param, char *seq_id);
/* Block crypto callback bridge - notifies Go of key restriction */
int goBlockBridge(char *corpid);
/* Cancel block crypto callback bridge - notifies Go of restriction lift */
int goCancelBlockBridge(char *corpid);
/* Wrapper functions for CGO */
int init(char *path, char *my_id, void *reserved);
int clearCache(char *corpid);
void freeCryptoBuf(void *buf);
#endif /* GOPROXY_BRIDGE_H */
+155
View File
@@ -0,0 +1,155 @@
package safechat
import (
"crypto/tls"
"fmt"
"io"
"net/http"
"net/url"
"strings"
"time"
)
// httpLogBodyLimit is the maximum body length (bytes) included in debug logs.
// Beyond this, the body is truncated with a "..." marker so the log line stays
// readable. Full bodies are still passed to the C library unchanged.
const httpLogBodyLimit = 4096
// keyClient handles HTTP communication with the key server.
// It is responsible for sending key requests (triggered by C library's goProxy callback)
// and returning the JSON response to be fed back via setResponse.
//
// The code field holds a DingTalk authCode used for server authentication.
// It is required when:
// - Fetching keys for the first time (empty keystore)
// - Server-side key version rotation (C library detects version mismatch)
type keyClient struct {
httpClient *http.Client
code string // DingTalk authCode
keyServer string // Optional override for key server URL
logger Logger
}
// newKeyClient creates a new key client with the given configuration.
func newKeyClient(cfg Config) *keyClient {
transport := &http.Transport{
TLSClientConfig: &tls.Config{
// Allow connecting to enterprise private servers with self-signed certs
InsecureSkipVerify: true,
},
MaxIdleConns: 10,
IdleConnTimeout: 30 * time.Second,
DisableCompression: true,
}
return &keyClient{
httpClient: &http.Client{
Timeout: cfg.HTTPTimeout,
Transport: transport,
},
code: cfg.Code,
keyServer: cfg.KeyServer,
logger: cfg.Logger,
}
}
func (kc *keyClient) doKeyRequest(fullURL, param, code string) (string, error) {
// Use override key server if configured
targetURL := fullURL
if kc.keyServer != "" {
targetURL = kc.keyServer
}
var body string
if strings.HasPrefix(param, "param=") {
// V1 format: C library output already has "param=" prefix and all fields.
body = fmt.Sprintf("%s&code=%s", param, url.QueryEscape(code))
} else {
body = fmt.Sprintf("%s&code=%s&appAlgVersion=1", param, url.QueryEscape(code))
}
// === Request logging (debug) ===
// Log full request line, headers and body (truncated) so we can verify
// the URL, the URL-encoded payload and the auth code are shaped as expected.
if kc.logger != nil {
kc.logger.Debug("=== HTTP key request ===")
kc.logger.Debug("URL: POST %s", targetURL)
kc.logger.Debug("C-URL: %s (from C library)", fullURL)
kc.logger.Debug("Server: %s%s",
func() string {
if kc.keyServer != "" {
return kc.keyServer + " (override)"
}
return "(C-provided)"
}(),
"")
kc.logger.Debug("Headers: Content-Type=application/x-www-form-urlencoded, User-Agent=SafeChat-Go-SDK/1.0")
kc.logger.Debug("Code (auth_token, length=%d): %s", len(code), previewString(code, 64))
kc.logger.Debug("Body (length=%d): %s", len(body), previewString(body, httpLogBodyLimit))
}
req, err := http.NewRequest("POST", targetURL, strings.NewReader(body))
if err != nil {
if kc.logger != nil {
kc.logger.Error("create request failed: %v", err)
}
return "", fmt.Errorf("create request failed: %w", err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
req.Header.Set("User-Agent", "SafeChat-Go-SDK/1.0")
// Capture timing so we can see if the server is slow / timing out.
start := time.Now()
resp, err := kc.httpClient.Do(req)
if err != nil {
if kc.logger != nil {
kc.logger.Error("HTTP request failed after %s: %v", time.Since(start), err)
}
return "", fmt.Errorf("HTTP request failed: %w", err)
}
defer resp.Body.Close()
respBody, err := io.ReadAll(resp.Body)
if err != nil {
if kc.logger != nil {
kc.logger.Error("read response body failed after %s: %v", time.Since(start), err)
}
return "", fmt.Errorf("read response body failed: %w", err)
}
duration := time.Since(start)
// === Response logging (debug on success, info on non-2xx) ===
if kc.logger != nil {
kc.logger.Debug("=== HTTP key response ===")
kc.logger.Debug("Status: %d %s (took %s)", resp.StatusCode, http.StatusText(resp.StatusCode), duration)
kc.logger.Debug("Headers: Content-Type=%s, Content-Length=%d", resp.Header.Get("Content-Type"), len(respBody))
kc.logger.Debug("Body (length=%d): %s", len(respBody), previewString(string(respBody), httpLogBodyLimit))
}
if resp.StatusCode != http.StatusOK {
if kc.logger != nil {
kc.logger.Error("key server returned HTTP %d %s (took %s, body length=%d): %s",
resp.StatusCode, http.StatusText(resp.StatusCode), duration, len(respBody),
previewString(string(respBody), httpLogBodyLimit))
}
return "", fmt.Errorf("key server returned HTTP %d", resp.StatusCode)
}
return string(respBody), nil
}
// previewString returns s unchanged when it fits in max bytes; otherwise it
// returns the first max bytes followed by "...[truncated, total=N]". This is
// used to keep HTTP request/response log lines readable for very large bodies
// while still preserving the head of the payload for debugging.
func previewString(s string, max int) string {
if max <= 0 || len(s) <= max {
return s
}
return s[:max] + fmt.Sprintf("...[truncated, total=%d]", len(s))
}
// updateCode updates the authentication code (may change during runtime).
func (kc *keyClient) updateCode(code string) {
kc.code = code
}
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
+615
View File
@@ -0,0 +1,615 @@
// Package safechat provides encryption/decryption capabilities for the
// Basic usage:
//
// client, err := safechat.New(safechat.Config{
// DataPath: "/path/to/keystore",
// UserID: "user123",
// Code: "authCode ", // authCode
// })
// if err != nil {
// log.Fatal(err)
// }
// defer client.Close()
//
// ciphertext, err := client.EncryptMsg("corp_id", "staff_id", []byte("hello"))
//
// The Code field must be a valid DingTalk authCode. It is required
// when fetching keys from the server (first-time use or key version rotation).
package safechat
/*
#include "csrc/safechat.h"
#include "goproxy_bridge.h"
#include <stdlib.h>
#include <string.h>
*/
import "C"
import (
"fmt"
"sync"
"unsafe"
"github.com/google/uuid"
)
// Client is the main SafeChat encryption client.
// It wraps the C library and provides a thread-safe Go API.
//
// IMPORTANT: Only one Client instance can exist per process because
// the underlying C library uses global state. Creating a second Client
// will return ErrAlreadyInitialized.
//
// All public methods are safe for concurrent use from multiple goroutines.
// Internal synchronization uses a two-lock design:
// - mu: serializes all C library calls (prevents C global state corruption)
// - kcMu: protects keyClient state during HTTP requests (used in goProxy callback)
type Client struct {
cfg Config
mu sync.Mutex // Serializes all C library calls
kcMu sync.Mutex // Protects keyClient during HTTP key requests
kc *keyClient // HTTP client for key server communication
inited bool // Whether C library init() has been called
blockedCorps sync.Map // map[string]bool - enterprises with restricted keys
}
// New creates a new SafeChat client and initializes the underlying C library.
//
// The Config.DataPath directory will be used to store encryption keys and
// related metadata. It must exist and be writable.
//
// Returns ErrAlreadyInitialized if called more than once per process.
func New(cfg Config) (*Client, error) {
cfg = defaultConfig(cfg)
if err := cfg.validate(); err != nil {
return nil, err
}
// Check if already initialized (C library is singleton)
if getGlobalClient() != nil {
return nil, ErrAlreadyInitialized
}
c := &Client{
cfg: cfg,
kc: newKeyClient(cfg),
}
// Register globally for CGO callbacks before calling init
registerGlobalClient(c)
// Initialize C library
if err := c.cInit(); err != nil {
globalClient.Store((*Client)(nil))
return nil, fmt.Errorf("safechat init failed: %w", err)
}
c.inited = true
return c, nil
}
// Close releases resources held by the client.
// After Close is called, no other methods should be called.
func (c *Client) Close() {
c.mu.Lock()
defer c.mu.Unlock()
c.inited = false
// Clear the global singleton so a new Client can be created after Close.
// The C library has no explicit cleanup function; keys are persisted to
// disk and process memory is reclaimed on exit.
globalClient.Store((*Client)(nil))
}
// EncryptMsg encrypts a plaintext message for the given enterprise.
//
// Returns the ciphertext in the standard SafeChat format:
// base64(encrypted_data)||key_version||method_num||plain_length
//
// If the encryption key is not yet available, the SDK will automatically
// request it from the key server (via goProxy callback) and retry.
func (c *Client) EncryptMsg(corpID, staffID string, plaintext []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(plaintext) == 0 {
return nil, fmt.Errorf("safechat: plaintext cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cEncryptData(corpID, staffID, plaintext)
if err == nil {
return result, nil
}
// Check if it's a "key requested" status - retry
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
// goProxy was called and setResponse was invoked synchronously.
// Next iteration should find the key in local cache.
continue
}
// Check for key restriction
if cerr, ok := err.(*CError); ok && cerr.Code == cEncryptDataKeyNegtive {
return nil, ErrKeyRestricted
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// DecryptMsg decrypts a ciphertext message.
//
// The ciphertext must be in the standard SafeChat format:
// base64(encrypted_data)||key_version||method_num||plain_length
func (c *Client) DecryptMsg(corpID, staffID string, ciphertext []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(ciphertext) == 0 {
return nil, fmt.Errorf("safechat: ciphertext cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cDecryptData(corpID, staffID, ciphertext)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
if cerr, ok := err.(*CError); ok && cerr.Code == cDecryptDataKeyNegtive {
return nil, ErrKeyRestricted
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// EncryptFile encrypts a file from srcPath to dstPath.
//
// The encrypted file uses a 12-byte header (msg_HandInfo_t) followed by
// SM4-ECB encrypted data in 8KiB blocks.
func (c *Client) EncryptFile(corpID, staffID, srcPath, dstPath string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
err := c.cEncryptFile(corpID, staffID, srcPath, dstPath)
if err == nil {
return nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return err
}
return ErrMaxRetryExceeded
}
// DecryptFile decrypts a file from srcPath to dstPath.
func (c *Client) DecryptFile(corpID, staffID, srcPath, dstPath string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
err := c.cDecryptFile(corpID, staffID, srcPath, dstPath)
if err == nil {
return nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return err
}
return ErrMaxRetryExceeded
}
// EncryptBuffer encrypts binary data in memory.
//
// The result includes a 12-byte header followed by SM4-ECB encrypted content.
func (c *Client) EncryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(data) == 0 {
return nil, fmt.Errorf("safechat: data cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cEncryptBuffer(corpID, staffID, data)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// DecryptBuffer decrypts binary data in memory.
func (c *Client) DecryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
if !c.inited {
return nil, ErrNotInitialized
}
if len(data) == 0 {
return nil, fmt.Errorf("safechat: data cannot be empty")
}
c.mu.Lock()
defer c.mu.Unlock()
for i := 0; i <= c.cfg.MaxRetry; i++ {
result, err := c.cDecryptBuffer(corpID, staffID, data)
if err == nil {
return result, nil
}
if cerr, ok := err.(*CError); ok && cerr.Code == cSendRequestParamOK {
continue
}
return nil, err
}
return nil, ErrMaxRetryExceeded
}
// SetResponse manually injects a key server response into the C library.
// This is an advanced API for cases where the caller handles HTTP
// communication externally.
func (c *Client) SetResponse(corpID, jsonStr string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
return c.cSetResponse(corpID, jsonStr)
}
// SetPushData processes a server push notification (key update, etc).
func (c *Client) SetPushData(corpID, staffID, pushData string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
return c.cSetPushData(corpID, staffID, pushData)
}
// ClearCache clears the local key cache for a specific enterprise.
func (c *Client) ClearCache(corpID string) error {
if !c.inited {
return ErrNotInitialized
}
c.mu.Lock()
defer c.mu.Unlock()
cCorpID := C.CString(corpID)
defer C.free(unsafe.Pointer(cCorpID))
ret := C.clearCache(cCorpID)
return mapCError(int(ret))
}
// IsBlocked returns true if the given enterprise's key is restricted.
func (c *Client) IsBlocked(corpID string) bool {
_, ok := c.blockedCorps.Load(corpID)
return ok
}
// UpdateCode updates the DingTalk authentication code at runtime.
func (c *Client) UpdateCode(code string) {
c.kcMu.Lock()
defer c.kcMu.Unlock()
c.cfg.Code = code
c.kc.updateCode(code)
}
// ============================================================
// CGO wrapper methods (called with c.mu held)
// ============================================================
func (c *Client) cInit() error {
cPath := C.CString(c.cfg.DataPath)
cUserID := C.CString(c.cfg.UserID)
defer C.free(unsafe.Pointer(cPath))
defer C.free(unsafe.Pointer(cUserID))
ret := C.init(cPath, cUserID, nil)
return mapCError(int(ret))
}
func (c *Client) cEncryptData(corpID, staffID string, plaintext []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.encryptData(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&plaintext[0])),
C.uint(len(plaintext)),
nil, // id - NULL for group messages
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
// Copy result to Go slice and free C buffer
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cDecryptData(corpID, staffID string, ciphertext []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
// The C layer treats the ciphertext as a NUL-terminated C string
// (sscanf/strtok). Go byte slices are not NUL-terminated and strtok would
// mutate the caller's buffer, so pass a NUL-terminated private copy.
cbuf := make([]byte, len(ciphertext)+1)
copy(cbuf, ciphertext)
ret := C.decryptData(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&cbuf[0])),
C.uint(len(ciphertext)),
nil, // id
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cEncryptFile(corpID, staffID, srcPath, dstPath string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
cSrcPath := C.CString(srcPath)
cDstPath := C.CString(dstPath)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSrcPath))
defer C.free(unsafe.Pointer(cDstPath))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.encryptFile(
cCorpID,
cStaffID,
cID,
cSrcPath,
cDstPath,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cDecryptFile(corpID, staffID, srcPath, dstPath string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
cSrcPath := C.CString(srcPath)
cDstPath := C.CString(dstPath)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSrcPath))
defer C.free(unsafe.Pointer(cDstPath))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.decryptFile(
cCorpID,
cStaffID,
cID,
cSrcPath,
cDstPath,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cEncryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.encryptBuffer(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&data[0])),
C.uint(len(data)),
cID,
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cDecryptBuffer(corpID, staffID string, data []byte) ([]byte, error) {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cID := C.CString("")
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cID))
defer C.free(unsafe.Pointer(cSeqID))
var outBuf *C.uchar
var outLen C.uint
ret := C.decryptBuffer(
cCorpID,
cStaffID,
(*C.uchar)(unsafe.Pointer(&data[0])),
C.uint(len(data)),
cID,
&outBuf,
&outLen,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return nil, mapCError(retCode)
}
result := C.GoBytes(unsafe.Pointer(outBuf), C.int(outLen))
C.freeCryptoBuf(unsafe.Pointer(outBuf))
return result, nil
}
func (c *Client) cSetResponse(corpID, jsonStr string) error {
cCorpID := C.CString(corpID)
cJSON := C.CString(jsonStr)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cJSON))
ret := C.setResponse(
cCorpID,
cJSON,
(C.block_crypto_func)(C.goBlockBridge),
)
retCode := int(ret)
if retCode != cFunctionOK {
return mapCError(retCode)
}
return nil
}
func (c *Client) cSetPushData(corpID, staffID, pushData string) error {
cCorpID := C.CString(corpID)
cStaffID := C.CString(staffID)
cPushData := C.CString(pushData)
seqID := uuid.New().String()
cSeqID := C.CString(seqID)
defer C.free(unsafe.Pointer(cCorpID))
defer C.free(unsafe.Pointer(cStaffID))
defer C.free(unsafe.Pointer(cPushData))
defer C.free(unsafe.Pointer(cSeqID))
ret := C.setPushData(
cCorpID,
cStaffID,
cPushData,
cSeqID,
(C.call_proxy_func)(C.goProxyBridge),
(C.cancel_block_crypto_func)(C.goCancelBlockBridge),
)
retCode := int(ret)
if retCode != cFunctionOK && retCode != cSendRequestParamOK {
return mapCError(retCode)
}
return nil
}
+5
View File
@@ -0,0 +1,5 @@
package safechat
// Version is the current version of the SafeChat Go SDK.
// This version follows semantic versioning (semver).
const Version = "1.0.0"