Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d6781d2ffb |
@@ -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.
|
||||
@@ -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 ./...
|
||||
|
||||
|
||||
@@ -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
@@ -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 一致 |
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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 不提供的高级多事件控制",
|
||||
},
|
||||
|
||||
@@ -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] {
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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 := ¶mAliasCaptureCaller{}
|
||||
_, 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 := ¶mAliasCaptureCaller{}
|
||||
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 := ¶mAliasCaptureCaller{}
|
||||
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 := ¶mAliasCaptureCaller{}
|
||||
_, 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 := ¶mAliasCaptureCaller{}
|
||||
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 {
|
||||
|
||||
@@ -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).
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
@@ -81,7 +81,6 @@ var schemaCatalogToolOptionalKeys = []string{
|
||||
"pagination",
|
||||
"positionals",
|
||||
"result",
|
||||
"wait",
|
||||
}
|
||||
|
||||
var schemaCatalogToolEnums = map[string][]string{
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -30,7 +30,6 @@ type ContractFinalPayload struct {
|
||||
Parameters []ParamDecl
|
||||
Safety *SafetySpec
|
||||
DryRun *DryRunSpec
|
||||
Wait *WaitSpec
|
||||
Result *ResultSpec
|
||||
Pagination *PaginationSpec
|
||||
Interface *InterfaceSpec
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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) })
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"}},
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
} {
|
||||
|
||||
@@ -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: "审批单终止",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
})
|
||||
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 © }
|
||||
|
||||
type cloneNode struct {
|
||||
|
||||
@@ -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 ©
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,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"
|
||||
---
|
||||
|
||||
|
||||
@@ -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 的订阅。
|
||||
|
||||
@@ -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`,否则会跳过自动退订。
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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{}) {}
|
||||
Vendored
+5
@@ -0,0 +1,5 @@
|
||||
module safechat-go-sdk
|
||||
|
||||
go 1.21
|
||||
|
||||
require github.com/google/uuid v1.6.0
|
||||
Vendored
+2
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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"
|
||||
Reference in New Issue
Block a user