Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
af000a8dfa | ||
|
|
f19a3ccfa5 | ||
|
|
89c5038446 | ||
|
|
4f915e4e2c | ||
|
|
ec03b7cca3 | ||
|
|
e0544579d2 | ||
|
|
ce43280c11 | ||
|
|
74ca40c197 | ||
|
|
cfaa673863 | ||
|
|
3bc6c31a2d | ||
|
|
383aeefaf6 | ||
|
|
a5bede3a19 | ||
|
|
bbf66e23d6 | ||
|
|
5e168c92cf | ||
|
|
725577103d | ||
|
|
f762117d4e | ||
|
|
750b6c04d6 | ||
|
|
59e51c348a | ||
|
|
a056a9abfb | ||
|
|
33ae780103 | ||
|
|
daf56514f7 | ||
|
|
8bcbceb971 | ||
|
|
df01f36442 | ||
|
|
b0024aa669 | ||
|
|
7b7aeadbbe | ||
|
|
94ad422a9f | ||
|
|
4f1ee37508 | ||
|
|
fec0347cd6 | ||
|
|
93318f4a83 | ||
|
|
a14fd0250c | ||
|
|
c99e228669 | ||
|
|
95d495f290 | ||
|
|
4bf300d862 | ||
|
|
1a1fc531f5 | ||
|
|
9fc570607f | ||
|
|
4851d19141 | ||
|
|
42fb25d150 | ||
|
|
416ad6571d | ||
|
|
9a119fbd64 | ||
|
|
c5decb2f90 | ||
|
|
fae2a4f5f0 | ||
|
|
d25b106e4f | ||
|
|
9f78e51ae7 | ||
|
|
d2752d8b5b | ||
|
|
8ecbff391c | ||
|
|
d259864a2b | ||
|
|
408098bdc1 | ||
|
|
658ec1676c | ||
|
|
e36d3b3474 | ||
|
|
0b9952c58d | ||
|
|
56af1ea091 | ||
|
|
ea5859b92b | ||
|
|
19f2ed5c69 | ||
|
|
efbaf7a49d | ||
|
|
374a9e9b13 | ||
|
|
d7d85c9e67 | ||
|
|
0fa982fe91 | ||
|
|
c4fb1bbd3e | ||
|
|
26d7d8f946 | ||
|
|
05ac342c4b | ||
|
|
5e491aef8f | ||
|
|
202187d5e2 | ||
|
|
13877b1c3a | ||
|
|
0e72e89ba3 | ||
|
|
f1b68271cc | ||
|
|
83efff21cd | ||
|
|
e6a4b35921 | ||
|
|
cc2d97ddba | ||
|
|
b78dd19cf9 | ||
|
|
1f0a75f836 | ||
|
|
16202c83a3 | ||
|
|
f4cc76c77d | ||
|
|
9fef6a9c43 | ||
|
|
810985b03a | ||
|
|
02633c6bd3 | ||
|
|
eb9416aa16 | ||
|
|
65b64af213 | ||
|
|
f1d160a481 | ||
|
|
f8c7f012a1 | ||
|
|
45618a55e6 | ||
|
|
c49583836b | ||
|
|
9dc8dc7065 | ||
|
|
f978e306cc | ||
|
|
aec852f971 | ||
|
|
143f781064 | ||
|
|
953b422295 | ||
|
|
da1a0f1299 | ||
|
|
bc7d19cfd8 | ||
|
|
df3122090f | ||
|
|
713fdf6188 | ||
|
|
70e21b58b4 | ||
|
|
18ebba1bb2 | ||
|
|
937404e6df | ||
|
|
88e155dd23 | ||
|
|
2e2cea0973 | ||
|
|
9b8c13a8b6 | ||
|
|
8238cc9f41 | ||
|
|
e59c4f30b8 | ||
|
|
fd7ef5edc2 | ||
|
|
a8e1acec09 | ||
|
|
5e003a41b1 | ||
|
|
4eaeb1dd4a | ||
|
|
84471bd6f0 | ||
|
|
cc4dd1e87b |
@@ -1 +1 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 52.5%"><title>coverage: 52.5%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">52.5%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">52.5%</text></g></svg>
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 49.8%"><title>coverage: 49.8%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">49.8%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">49.8%</text></g></svg>
|
||||
|
Before Width: | Height: | Size: 1.1 KiB After Width: | Height: | Size: 1.1 KiB |
@@ -28,3 +28,5 @@ credentials*
|
||||
plans
|
||||
_docs
|
||||
dws.zip
|
||||
*.code-workspace
|
||||
/dingtalk-workspace.zip
|
||||
|
||||
+293
@@ -4,6 +4,299 @@ All notable changes to this project will be documented in this file.
|
||||
|
||||
The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and this project follows [Semantic Versioning](https://semver.org/).
|
||||
|
||||
## [1.0.13] - 2026-04-22
|
||||
|
||||
IM / Messaging capability expansion: the `chat` (aka `im`) product surface grows from "group + bot messaging" into a full conversational layer — user-identity messaging, message reading & search, personal messages, topic replies, mentions, focused contacts, unread/top/common conversations, org-wide group creation, and first-class bot lifecycle.
|
||||
|
||||
### Added
|
||||
|
||||
- **`dws im` alias** — `dws im` is now registered as an alias of `dws chat` for intent clarity
|
||||
- **User-identity messaging** (`chat message send`) — send group or 1-on-1 messages as the current user
|
||||
- Recipient selection is mutually exclusive: `--group <openConversationId>` / `--user <userId>` / `--open-dingtalk-id <openDingTalkId>`
|
||||
- Markdown text via `--text` (or positional arg), optional `--title`
|
||||
- Group-only: `--at-all` to @everyone, `--at-users` for per-member @mentions
|
||||
- Image messages via `--media-id` (obtained from `dt_media_upload`)
|
||||
- **Personal messages** (`chat message send-personal`) — sensitive personal-channel send (⚠️ destructive/dangerous op, requires confirmation)
|
||||
- **Conversation read paths**:
|
||||
- `chat message list` — pull group / 1-on-1 conversation messages
|
||||
- `chat message list-all` — pull all conversations for the current user in a time range
|
||||
- `chat message list-topic-replies` — pull group topic reply threads
|
||||
- `chat message list-by-sender` — messages by a specific sender
|
||||
- `chat message list-mentions` — messages where the current user was @-mentioned
|
||||
- `chat message list-focused` — messages from focused / starred contacts
|
||||
- `chat message list-unread-conversations` — unread conversation list
|
||||
- `chat message search` — keyword search across conversations
|
||||
- `chat message info` — conversation metadata
|
||||
- `chat list-top-conversations` — pinned conversation list
|
||||
- **Group creation & discovery**:
|
||||
- `chat group create-org` — create an organization-wide group
|
||||
- `chat search-common` — search groups shared with a nickname list (`--nicks`, `--match-mode AND|OR`, cursor-based pagination)
|
||||
- **Bot lifecycle**:
|
||||
- `chat bot create` — create an enterprise bot
|
||||
- `chat bot search-groups` — search the groups a bot is present in
|
||||
|
||||
### Changed
|
||||
|
||||
- **`chat` skill reference** (`skills/references/products/chat.md`, #148) restructured into three sub-groups — `group` (9) / `message` (15) / `bot` (3) — with refreshed intent-routing table, workflow examples, and context-passing rules aligned with `dws-service-endpoints.json` (16 new group-chat tool overrides + 2 new bot tool overrides)
|
||||
- **README Key Services** sync:
|
||||
- `Chat` row: 10 → 20 commands; subcommand tags expanded to `message` `group` `search` `list-top-conversations`
|
||||
- `Bot` row: 6 → 7 commands; subcommand tags expanded with `create` `search-groups`
|
||||
- Total raised to **152 commands across 14 products**
|
||||
|
||||
## [1.0.12] - 2026-04-21
|
||||
|
||||
Product-surface expansion: first-class `doc` (DingTalk Docs) and `minutes` (AI Minutes) skill references, refreshed `aitable` guide aligned with the shipped binary (including dashboard / chart / export), and a README sync that brings the full command catalog to **141 commands across 14 products**.
|
||||
|
||||
### Added
|
||||
|
||||
- **`doc` skill reference** (`skills/references/products/doc.md`) — 16-command coverage of DingTalk Docs:
|
||||
- Discovery: `search`, `list`, `info`, `read`
|
||||
- Authoring: `create`, `update`, `folder create`
|
||||
- Files: `upload`, `download`
|
||||
- Block-level editing: block `query`, `insert`, `update`, `delete`
|
||||
- Comments: `comment list`, `create`, `reply`
|
||||
- URL → `doc_id` extraction rules and nodeId dual-format notes
|
||||
- **`minutes` skill reference** (`skills/references/products/minutes.md`) — coverage of AI Minutes:
|
||||
- Lists: personal / shared-with-me / all-accessible
|
||||
- Content: basic info, AI summary, keywords, transcription, extracted todos, batch detail
|
||||
- Editing: title update
|
||||
- Recording control: start, pause, resume, stop
|
||||
- **SKILL.md routing**:
|
||||
- Product overview table rows for `doc` and `minutes`
|
||||
- Intent decision tree routes — `钉钉文档/云文档/知识库/块级编辑/文档评论` → `doc`; `听记/AI听记/会议纪要/转写/摘要/思维导图/发言人/热词` → `minutes`
|
||||
- Danger-op table entries: `doc delete`, `doc block delete`
|
||||
- `aitable` description completed with the `附件` (attachment) group
|
||||
- **`aitable` skill enhancements**:
|
||||
- `field create` single-field mode (`--name` / `--type` / `--config`) with examples
|
||||
- `base get` URL → `baseId` quick-tip
|
||||
- Dedicated "URL → baseId 提取" chapter
|
||||
- "`--filters` 筛选语法排错与使用规范" chapter
|
||||
- "相关产品" cross-link section pointing to `doc`
|
||||
- **"复杂操作" chapter** (#141) — dashboard / chart workflow (with two-call sequencing and `chart share get` vs `dashboard share get` error semantics) and two-stage `export data` polling (`scope=all/table/view` parameter constraints)
|
||||
- **README Key Services sync** (#140):
|
||||
- New rows: `doc` (16 commands), `minutes` (22 commands — adds `hot-word`, `mind-graph`, `replace-text`, `speaker`, `upload` subgroups)
|
||||
- `aitable` expanded from 20 → 37 commands; surfaces `chart`, `dashboard`, `export`, `import`, `view` subgroups
|
||||
- Total command count updated from **86 → 141 across 14 products**
|
||||
- "Coming soon" list drops `doc` and `minutes`
|
||||
|
||||
### Changed
|
||||
|
||||
- `aitable record query` docs rename `--keyword` → `--query` to match the shipped binary
|
||||
- `aitable record query` docs clarify `--sort` direction semantics (avoids misuse of `order`)
|
||||
- `aitable base list` guidance strengthened — "only for recent browsing; use `base search` for lookups"; intent decision prioritizes `base search` for base discovery
|
||||
|
||||
## [1.0.11] - 2026-04-20
|
||||
|
||||
Plugin subsystem hardening: faster cold startup, cleaner lifecycle, stricter isolation, and polished UX for PAT / i18n / error routing.
|
||||
|
||||
### Added
|
||||
|
||||
- `feat: supports claw-like products` — overlay path for Claw-style embedded editions
|
||||
- `feat(plugin): inject user identity (UserID, CorpID) into stdio plugin subprocesses`
|
||||
- `feat(auth): improve login UX for terminal auth denial cases` — clearer messaging + retry affordance
|
||||
- `feat: PAT scope error visualization and auto-retry with authorization polling` (#113)
|
||||
- Human-readable error output (lark-cli style) with type/message/hint/authorization command
|
||||
- JSON payload also available via `--format json`
|
||||
- Auto-retry once the user completes scope authorization
|
||||
|
||||
### Changed
|
||||
|
||||
- `perf(plugin): serve plugin MCP tool list from disk cache on startup` — hot path skips Initialize+ListTools when snapshot exists
|
||||
- `perf(plugin): parallelize all plugin discovery and tighten cold timeouts` — HTTP cold budget 4s → 700ms (auth) / 500ms (plain); stdio and HTTP fan out concurrently
|
||||
- `perf(plugin): share cache.Store across discovery` — single `*cache.Store` above the fan-out instead of per-goroutine instances
|
||||
- `refactor(plugin): remove default/managed plugin privileged mechanism` (#124) — third-party plugins install on an equal footing via `dws plugin install`
|
||||
- `refactor(plugin): purge removed plugin settings instead of merely disabling` — `RemovePlugin` now deletes `EnabledPlugins` and `PluginConfigs` entries
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(transport): cap plugin MCP startup at ~4s when endpoints are unreachable` (#119) — eliminates the 10s `dws --help` stall caused by compounding transport timeouts
|
||||
- `fix(plugin): stop stdio child processes on exit and before removal` — no more orphaned plugin subprocesses
|
||||
- `fix(pat): avoid shared PAT command state in root registration` (#129)
|
||||
- `fix: -f json 模式下错误 JSON 从 stdout 改为输出到 stderr` (#133) — restores CI stderr-based failure assertions
|
||||
- `fix(cli): localize plugin/help command strings via i18n` (#118, #134) — zh locale now shows consistent Chinese `--help`; wraps plugin module, help command, and OAuth client-id/secret flag descriptions
|
||||
- `chore: remove workspace and bundled artifacts` (#127) — clean local-only repository leftovers
|
||||
|
||||
## [1.0.9] - 2026-04-16
|
||||
|
||||
Plugin system launch + execution-pipeline overhaul. This is the largest release since 1.0.0: third-party MCP servers become first-class commands, the command pipeline grows to five stages, and the edition overlay gains the hooks needed for embedded hosts.
|
||||
|
||||
### Added
|
||||
|
||||
#### Plugin system (new)
|
||||
|
||||
- `plugin` command family: `install`, `list`, `info`, `enable`, `disable`, `remove`, `create`, `dev`, `config set/get/list/unset`
|
||||
- Plugin manifest parsing/validation, managed/user directory-based identity
|
||||
- MCP server conversion and injection into the dynamic routing registry
|
||||
- Pipeline hook adapter for shell-based hooks
|
||||
- Stdio transport: subprocess lifecycle, `DWS_PLUGIN_ROOT` / `DWS_PLUGIN_DATA` variable expansion
|
||||
- Stdio server tools automatically registered as CLI subcommands (e.g. `dws hello greet --name Peter`)
|
||||
- Streamable-HTTP MCP tool discovery via `registerHTTPServer`
|
||||
- Updater: managed plugin update check on CLI startup (10 s timeout, best-effort)
|
||||
- `dws plugin create` scaffold (plugin.json, SKILL.md, hooks.json); `dws plugin dev` source-dir registration without copy
|
||||
- `SyncSkills` — copies plugin skills to agent directories on startup
|
||||
- **Auth Token Registry**: per-server HTTP headers declared in `plugin.json` for third-party MCP servers (e.g. Alibaba Cloud Bailian) independent from DingTalk OAuth
|
||||
- **Persistent plugin config** (`dws plugin config ...`): values persisted to `~/.dws/settings.json`, auto-injected as env vars; `${KEY}` in `plugin.json` resolves without manual `export`
|
||||
- **Build lifecycle**: `build` field compiles stdio servers to native binaries at install time
|
||||
- **Command-name conflict protection**: reserved built-in names (`auth`, `plugin`, `cache`, …) and plugin-vs-plugin duplicate detection
|
||||
- Parallel service discovery (`sync.WaitGroup`) — startup reduced from sequential `N*10s` to parallel `max(10s)`
|
||||
|
||||
#### Core commands & diagnostics
|
||||
|
||||
- `dws doctor` — one-stop environment/auth/network diagnostics
|
||||
- `dws config list` — centralized view of scattered configuration
|
||||
- Structured perf tracing (upgraded from debug tool to diagnostics output)
|
||||
- `feat(skill): restore find/get for legacy skill market API` — `skill find`, `skill get`; `skill add` still uses aihub download
|
||||
|
||||
#### Edition / overlay hooks
|
||||
|
||||
- `edition.Hooks.SaveToken` / `LoadToken` / `DeleteToken` — delegate token persistence with keychain fallback
|
||||
- `edition.Hooks.AuthClientID` / `AuthClientFromMCP` — overlay can override the OAuth client ID and route auth through MCP endpoints
|
||||
- `edition.Hooks.AfterPersistentPreRun` — wire non-MCP clients (e.g. A2A gateway) after root setup
|
||||
- `edition.Hooks.ClassifyToolResult` — custom MCP result classification before the default business-error detection
|
||||
- Token marker file (`token.json`) for embedded hosts to detect auth state without keychain access
|
||||
- `pkg/runtimetoken.ResolveAccessToken` mirroring MCP auth resolution; MCP identity headers exported via `pkg/cli` for auxiliary HTTP transports
|
||||
- `ExitCoder` interface — edition-specific errors carry custom exit codes
|
||||
- `RawStderrError` interface — errors that bypass CLI formatting and emit raw stderr (for desktop runtimes)
|
||||
|
||||
### Changed
|
||||
|
||||
- **Command execution pipeline: 3 → 5 stages** (`Register → PreParse → PostParse → PreRequest → PostResponse`)
|
||||
- `feat(schema): return structured degraded errors instead of silent empty catalog` — new `CatalogDegraded` error with reasons `unauthenticated` / `market_unreachable` / `runtime_all_failed`; auth pre-check short-circuits doomed MCP connections
|
||||
- `refactor(auth): unify auxiliary token resolution with MCP cached path` — shared `resolveAccessTokenFromDir`; overlays reuse the process-level token cache
|
||||
- `feat(plugin): improve CLI overlay resolution and plugin install robustness`
|
||||
- `plugin.json` `cli` field now accepts a file path (e.g. `"cli": "overlay.json"`) in addition to inline JSON
|
||||
- `description` field on `CLIToolOverride` for static fallback when MCP `tools/list` is unavailable
|
||||
- Windows install uses `cmd /C` instead of `sh -c` for build commands
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(plugin): harden plugin system security boundaries`
|
||||
- Reject `file://` / local paths in git URLs; allow only `https` / `ssh`
|
||||
- Reject symlink entries during ZIP extraction (path-traversal defense)
|
||||
- `build.output` must be a relative path within the plugin directory
|
||||
- Reject absolute paths in stdio command declarations
|
||||
- Block dangerous env var names (`PATH`, `LD_PRELOAD`, …) from plugin config injection
|
||||
- `fix(plugin): schema flag params, HTTP tool discovery, and integration tests`
|
||||
- `fix(plugin): skip min version check in dev mode`
|
||||
|
||||
## [1.0.8] - 2026-04-07
|
||||
|
||||
AITable command surface expansion, installer alignment with npm conventions, and execution-timeout hardening.
|
||||
|
||||
### Added
|
||||
|
||||
- **AITable static helper commands** (20 commands in total) replacing dynamic routing:
|
||||
- `base`: `list`, `search`, `get`, `create`, `update`
|
||||
- `table`: `get`, `create`, `update`
|
||||
- `field`: `get`, `create`, `update`
|
||||
- `record`: `query`, `create`, `update`
|
||||
- `template`: `search`
|
||||
- `attachment`: `upload`
|
||||
- `feat(install): align skill dirs with npm and add OpenClaw` — skill install paths follow npm conventions; OpenClaw added to supported agents
|
||||
- Label rendering optimization for AITable records (`to #73551688`)
|
||||
- README: npm install method documented
|
||||
- README: note that `dws upgrade` requires v1.0.7+
|
||||
|
||||
### Changed
|
||||
|
||||
- `perf: optimize command timeout handling, instrumentation, and diagnostics`
|
||||
|
||||
## [1.0.7] - 2026-04-02
|
||||
|
||||
Self-upgrade, edition overlay foundation, and fail-closed auth enforcement.
|
||||
|
||||
### Added
|
||||
|
||||
- **`dws upgrade`** — self-upgrade via GitHub Releases; atomic replace; cross-platform (macOS/Linux/Windows)
|
||||
- `feat: edition layer for Wukong overlay` — build-time edition hook lets downstream overlays customize auth UX, config dir, static server list, visible products, and extra root commands
|
||||
- `pkg/edition` defaults + `pkg/editiontest` contract tests
|
||||
- `Makefile` target `edition-test`; CI job `edition-tests`
|
||||
- Static server injection skips market discovery when configured
|
||||
- Deduplicates top-level commands so overlay wins
|
||||
- `hideNonDirectRuntimeCommands` respects edition `VisibleProducts`
|
||||
- Gated `auth login` subcommand + hints for embedded editions
|
||||
- Optional token auto-purge; edition `ConfigDir` override
|
||||
- `dws version` — human-readable multi-line output plus JSON with edition, architecture, build, commit
|
||||
- Tag reporting for case suites (`to #73551688`)
|
||||
- `feat(auth): unify MCP retry constant and add retry to remaining endpoints`
|
||||
|
||||
### Changed
|
||||
|
||||
- `style(auth): redesign OAuth authorization pages UI`
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(auth): switch CLI auth check from fail-open to fail-closed`
|
||||
- When `/cli/cliAuthEnabled` is unreachable (network error/timeout/5xx), OAuth callback now routes to the permission request page instead of silently marking "enabled"
|
||||
- Device Flow blocks login and asks the user to verify network connectivity
|
||||
- `CheckCLIAuthEnabled` retries with backoff (3 attempts, 0s/1s/2s) to tolerate transient issues
|
||||
|
||||
## [1.0.6] - 2026-04-01
|
||||
|
||||
Error diagnostics overhaul, destructive-command confirmation, and credential auto-persistence.
|
||||
|
||||
### Added
|
||||
|
||||
- **Interactive confirmation for destructive dynamic commands** — prompts before delete/remove operations unless `--yes` is set
|
||||
- **Enhanced error diagnostics**
|
||||
- `ServerDiagnostics` struct extracts `trace_id`, `server_error_code`, `technical_detail`, `server_retryable` from MCP responses
|
||||
- Pulls diagnostics from JSON-RPC `error.data`, tool call result content, and HTTP headers (`X-Trace-Id`, `X-Request-Id`, `x-dingtalk-trace-id`)
|
||||
- Three verbosity levels for `PrintHuman`: Normal (trace ID + server code), Verbose (+ technical detail), Debug (+ RPC code / operation / reason)
|
||||
- Local logging now includes sanitized request body, response body on error, retry attempts, and classification events
|
||||
- `TruncateBody` / `SanitizeArguments` / `RedactHeaders` helpers with sensitive-key substring detection
|
||||
- **Auth credential persistence**
|
||||
- `feat(auth): enhance device flow with CLI auth check and admin guidance`
|
||||
- `feat(auth): persist OAuth credentials for reliable token refresh`
|
||||
- `feat(auth): persist client credentials and optimize keychain access` — auto-persist `--client-id` / `--client-secret`; keychain credential cache to avoid repeated reads; enhanced logout cleans `app.json` + keychain secrets + `token.json`
|
||||
- `add report helper with flexible date parsing and defaults`
|
||||
- `feat: to #73551688 支持消息通知`
|
||||
- README: Official App mode (recommended, direct login without creating an app) + Custom App mode; admin guide for enabling CLI access
|
||||
|
||||
### Changed
|
||||
|
||||
- Getting Started simplified with inline login commands; whitelist references removed from the IMPORTANT banner
|
||||
- Version bump documentation updated to v1.0.5 internal; co-creation group QR code refreshed
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix: resolve verbosity flag lookup, FileLogger lazy binding, and business error logging`
|
||||
- `resolveVerbosity` uses `cmd.Flags()` instead of `PersistentFlags()` so subcommands inherit `--verbose` / `--debug`
|
||||
- `FileLogger` lazy-binds in `executeInvocation` (after `configureLogLevel` init)
|
||||
- Business errors (HTTP 200 + `success=false`) now written to the file logger for offline diagnosis
|
||||
- OAuth callback race condition (write response before sending code)
|
||||
- `import path for errors package in skill_command.go`
|
||||
|
||||
## [1.0.4] - 2026-03-30
|
||||
|
||||
Token-refresh reliability and onboarding clarity.
|
||||
|
||||
### Added
|
||||
|
||||
- `feat(auth): persist client credentials for token refresh` — `--client-id` / `--client-secret` are stored for automatic refresh after expiration; client secret lives in the system Keychain with a file reference
|
||||
- README onboarding flow rewrite with step-by-step first-time setup and more realistic examples
|
||||
- Agent skill reference polish: clearer examples, updated intent routing patterns, expanded `simple.md` onboarding, cross-skill reference fixes
|
||||
|
||||
## [1.0.3] - 2026-03-29
|
||||
|
||||
Filtering power, schema rendering, and a native `todo` command family.
|
||||
|
||||
### Added
|
||||
|
||||
- **Nested / array-indexed output filtering**
|
||||
- `--fields` now accepts dot-notation (e.g. `--fields response.content`) and array index access (e.g. `response.items[0]`)
|
||||
- New field-path parser with recursive extraction logic
|
||||
- **`schema` command enhancements**
|
||||
- Table format output for human consumption
|
||||
- Product-level endpoint loading in the CLI loader
|
||||
- Schema-text rendering wired into the runner output pipeline
|
||||
- **`todo` task helper family** — static `create` / `update` / `done` / `get` / `delete` with `preferLegacyLeaf` replacing dynamic commands
|
||||
- MCP tool alignment: `create_personal_todo`, `update_todo_task`, `update_todo_done_status`, `query_todo_detail`, `delete_todo`
|
||||
- ISO-8601 due-time parsing
|
||||
- Hidden title aliases and delete confirmation
|
||||
- Priority field on `todo` helper
|
||||
- Expanded zh / en i18n coverage (fixes `en.json` spacing/wording issues)
|
||||
- README restructured with collapsible feature sections
|
||||
|
||||
## [1.0.2] - 2026-03-29
|
||||
|
||||
Deep workspace tooling upgrade: pipeline-based input correction, output filtering, enhanced stdin handling, and multi-endpoint routing.
|
||||
|
||||
@@ -350,24 +350,26 @@ dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
| Service | Command | Commands | Subcommands | Description |
|
||||
|---------|---------|:--------:|-------------|-------------|
|
||||
| Contact | `contact` | 6 | `user` `dept` | Search users by name/mobile, batch query, departments, current user profile |
|
||||
| Chat | `chat` | 10 | `message` `group` `search` | Group CRUD, member management, bot messaging, webhook |
|
||||
| Bot | `chat bot` | 6 | `bot` `group` `message` `search` | Robot creation/search, group/single messaging, webhook, message recall |
|
||||
| Chat / IM | `chat` (alias `im`) | 20 | `message` `group` `search` `list-top-conversations` | User-identity send (group / 1-on-1 / open-dingtalk-id), Markdown + image, @mentions; read & search conversations (list, list-all, topic replies, by-sender, mentions, focused, unread, search, info, top / common groups); group CRUD + member management |
|
||||
| Bot | `chat bot` | 7 | `bot` `group` `message` `search` `create` `search-groups` | Bot create / search, search bot groups; bot-identity group & batch-1:1 messaging, Webhook, message recall; add bot to group |
|
||||
| Calendar | `calendar` | 13 | `event` `room` `participant` `busy` | Events CRUD, meeting room booking, free-busy query, participant management |
|
||||
| Todo | `todo` | 6 | `task` | Create, list, update, done, get detail, delete |
|
||||
| Approval | `oa` | 9 | `approval` | Approve/reject/revoke, pending tasks, initiated instances, process list |
|
||||
| Attendance | `attendance` | 4 | `record` `shift` `summary` `rules` | Clock-in records, shift schedules, attendance summary, group rules |
|
||||
| Ding | `ding` | 2 | `message` | Send/recall DING messages |
|
||||
| Report | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | Create reports, sent/received list, templates, statistics |
|
||||
| AITable | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | Full CRUD for bases/tables/records/fields, templates |
|
||||
| AITable | `aitable` | 37 | `base` `table` `record` `field` `attachment` `template` `chart` `dashboard` `export` `import` `view` | Full CRUD for bases/tables/records/fields; charts/dashboards; data import/export; views; templates |
|
||||
| Doc | `doc` | 16 | `search` `list` `info` `read` `create` `update` `upload` `download` `folder` `block` `comment` | Search, read, create/update documents; block-level editing; file upload/download; comments |
|
||||
| Minutes | `minutes` | 22 | `list` `get` `update` `record` `hot-word` `mind-graph` `replace-text` `speaker` `upload` | List/search AI meeting transcripts; summaries, transcriptions, todos, mind-maps; recording control; speaker management, hot-words, file upload |
|
||||
| Workbench | `workbench` | 2 | `app` | Batch query app details |
|
||||
| DevDoc | `devdoc` | 1 | `article` | Search platform docs and error codes |
|
||||
|
||||
> 86 commands across 12 products. Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
|
||||
> 152 commands across 14 products. Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
|
||||
|
||||
<details>
|
||||
<summary>Coming soon</summary>
|
||||
|
||||
`doc` (documents) · `mail` (email) · `minutes` (AI transcription) · `drive` (cloud drive) · `conference` (video) · `tb` (Teambition) · `aiapp` (AI apps) · `live` (streaming) · `skill` (marketplace)
|
||||
`mail` (email) · `drive` (cloud drive) · `conference` (video) · `tb` (Teambition) · `aiapp` (AI apps) · `live` (streaming) · `skill` (marketplace)
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
@@ -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 app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
|
||||
// the same host compatibility hooks as MCP. It mirrors the former body of
|
||||
// getCachedRuntimeToken (excluding process-level cache and timing).
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, disc)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token), nil
|
||||
}
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
return "", tokenErr
|
||||
}
|
||||
manager := authpkg.NewManager(configDir, nil)
|
||||
configureLegacyAuthManagerCompatibility(manager)
|
||||
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
|
||||
return strings.TrimSpace(leg), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
|
||||
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
|
||||
// the active edition config directory, the same process-cached path as MCP is used.
|
||||
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
if t := strings.TrimSpace(explicitToken); t != "" {
|
||||
return t, nil
|
||||
}
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
|
||||
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
|
||||
return tok, nil
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
}
|
||||
tok, err := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if tok != "" {
|
||||
return tok, nil
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
}
|
||||
|
||||
func noCredentialsError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return fmt.Errorf("认证信息已失效,请重新认证")
|
||||
}
|
||||
return fmt.Errorf("no credentials found, run: dws auth login")
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResolveAuxiliaryAccessToken_explicitToken(t *testing.T) {
|
||||
tok, err := ResolveAuxiliaryAccessToken(context.Background(), "/any/dir", " bearer-xyz ")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if tok != "bearer-xyz" {
|
||||
t.Fatalf("got %q, want bearer-xyz", tok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveAuxiliaryAccessToken_emptyConfigDir(t *testing.T) {
|
||||
_, err := ResolveAuxiliaryAccessToken(context.Background(), " ", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty config directory")
|
||||
}
|
||||
}
|
||||
@@ -121,6 +121,7 @@ func newAuthLoginCommand() *cobra.Command {
|
||||
}
|
||||
}
|
||||
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
@@ -206,6 +207,7 @@ func newAuthLogoutCommand() *cobra.Command {
|
||||
_ = os.Remove(filepath.Join(configDir, "mcp_url"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token.json"))
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintln(w, "[OK] 已清除所有认证信息")
|
||||
@@ -308,6 +310,7 @@ func newAuthExchangeCommand() *cobra.Command {
|
||||
if err != nil {
|
||||
return apperrors.NewAuth(fmt.Sprintf("failed to exchange authorization code: %v", err))
|
||||
}
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
@@ -347,6 +350,7 @@ func newAuthResetCommand() *cobra.Command {
|
||||
}
|
||||
_ = os.Remove(filepath.Join(configDir, "mcp_url"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token"))
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintln(w, "[OK] 认证信息已重置")
|
||||
|
||||
@@ -44,7 +44,7 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
CorpID: "dingcorp",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
}
|
||||
|
||||
originalTransport := http.DefaultTransport
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
// 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 app
|
||||
|
||||
import "sync"
|
||||
|
||||
// PluginAuth holds authentication credentials for a plugin-owned
|
||||
// streamable-http MCP server. Each server is keyed by its canonical
|
||||
// product ID (CLI.ID) so that different servers can use independent
|
||||
// tokens without interfering with each other or with the default
|
||||
// DingTalk OAuth token.
|
||||
type PluginAuth struct {
|
||||
// Token is the Bearer token extracted from the plugin's
|
||||
// "Authorization" header (e.g. a third-party API key).
|
||||
Token string
|
||||
|
||||
// ExtraHeaders contains any additional custom HTTP headers
|
||||
// declared by the plugin (excluding Authorization).
|
||||
ExtraHeaders map[string]string
|
||||
|
||||
// TrustedDomains lists the hostnames that the token is allowed
|
||||
// to be sent to. Typically derived from the server endpoint.
|
||||
TrustedDomains []string
|
||||
}
|
||||
|
||||
var (
|
||||
pluginAuthMu sync.RWMutex
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
)
|
||||
|
||||
// RegisterPluginAuth stores authentication credentials for a plugin
|
||||
// server keyed by its canonical product ID. The runner looks up these
|
||||
// credentials at execution time to inject the correct Bearer token
|
||||
// instead of the default DingTalk OAuth token.
|
||||
func RegisterPluginAuth(productID string, auth *PluginAuth) {
|
||||
pluginAuthMu.Lock()
|
||||
defer pluginAuthMu.Unlock()
|
||||
pluginAuthRegistry[productID] = auth
|
||||
}
|
||||
|
||||
// LookupPluginAuth returns the authentication credentials registered
|
||||
// for the given product ID, or nil if none exists.
|
||||
func LookupPluginAuth(productID string) (*PluginAuth, bool) {
|
||||
pluginAuthMu.RLock()
|
||||
defer pluginAuthMu.RUnlock()
|
||||
auth, ok := pluginAuthRegistry[productID]
|
||||
return auth, ok
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestPluginAuthRegistry(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "test-product")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
// Initially not found
|
||||
if _, ok := LookupPluginAuth("test-product"); ok {
|
||||
t.Error("expected LookupPluginAuth to return false for unregistered product")
|
||||
}
|
||||
|
||||
// Register auth credentials
|
||||
auth := &PluginAuth{
|
||||
Token: "sk-test-token-12345",
|
||||
ExtraHeaders: map[string]string{"X-Custom": "value"},
|
||||
TrustedDomains: []string{"api.example.com", "*.example.com"},
|
||||
}
|
||||
RegisterPluginAuth("test-product", auth)
|
||||
|
||||
// Now should be found
|
||||
got, ok := LookupPluginAuth("test-product")
|
||||
if !ok {
|
||||
t.Fatal("expected LookupPluginAuth to return true after registration")
|
||||
}
|
||||
if got != auth {
|
||||
t.Error("LookupPluginAuth returned different auth instance")
|
||||
}
|
||||
if got.Token != "sk-test-token-12345" {
|
||||
t.Errorf("Token = %q, want sk-test-token-12345", got.Token)
|
||||
}
|
||||
if got.ExtraHeaders["X-Custom"] != "value" {
|
||||
t.Errorf("ExtraHeaders[X-Custom] = %q, want value", got.ExtraHeaders["X-Custom"])
|
||||
}
|
||||
if len(got.TrustedDomains) != 2 {
|
||||
t.Errorf("TrustedDomains len = %d, want 2", len(got.TrustedDomains))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginAuthRegistryIsolation(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "product-a")
|
||||
delete(pluginAuthRegistry, "product-b")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
authA := &PluginAuth{Token: "token-a"}
|
||||
authB := &PluginAuth{Token: "token-b"}
|
||||
|
||||
RegisterPluginAuth("product-a", authA)
|
||||
RegisterPluginAuth("product-b", authB)
|
||||
|
||||
gotA, okA := LookupPluginAuth("product-a")
|
||||
gotB, okB := LookupPluginAuth("product-b")
|
||||
|
||||
if !okA || !okB {
|
||||
t.Fatal("expected both products to be registered")
|
||||
}
|
||||
if gotA.Token != "token-a" {
|
||||
t.Errorf("product-a Token = %q, want token-a", gotA.Token)
|
||||
}
|
||||
if gotB.Token != "token-b" {
|
||||
t.Errorf("product-b Token = %q, want token-b", gotB.Token)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeriveToolCLIName(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"web_search", "web-search"},
|
||||
{"maps.search_poi", "search-poi"},
|
||||
{"maps.geo", "geo"},
|
||||
{"simple", "simple"},
|
||||
{"a.b.deep_nested_name", "deep-nested-name"},
|
||||
{"already-kebab", "already-kebab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
got := deriveToolCLIName(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("deriveToolCLIName(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPluginAuthFromHeaders(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "test-srv")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Key: "test-srv",
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
CLI: market.CLIOverlay{ID: "test-srv", Command: "test-srv"},
|
||||
AuthHeaders: map[string]string{
|
||||
"Authorization": "Bearer sk-my-secret-key",
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
|
||||
auth, ok := LookupPluginAuth("test-srv")
|
||||
if !ok {
|
||||
t.Fatal("expected auth to be registered after registerPluginAuthFromHeaders")
|
||||
}
|
||||
if auth.Token != "sk-my-secret-key" {
|
||||
t.Errorf("Token = %q, want sk-my-secret-key", auth.Token)
|
||||
}
|
||||
if auth.ExtraHeaders["X-Custom"] != "custom-value" {
|
||||
t.Errorf("ExtraHeaders[X-Custom] = %q, want custom-value", auth.ExtraHeaders["X-Custom"])
|
||||
}
|
||||
if len(auth.TrustedDomains) != 2 {
|
||||
t.Fatalf("TrustedDomains len = %d, want 2", len(auth.TrustedDomains))
|
||||
}
|
||||
if auth.TrustedDomains[0] != "api.example.com" {
|
||||
t.Errorf("TrustedDomains[0] = %q, want api.example.com", auth.TrustedDomains[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPluginAuthFromHeadersNoAuth(t *testing.T) {
|
||||
srv := market.ServerDescriptor{
|
||||
Key: "no-auth-srv",
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
CLI: market.CLIOverlay{ID: "no-auth-srv"},
|
||||
AuthHeaders: map[string]string{
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
|
||||
// Should not register because there's no Authorization header
|
||||
if _, ok := LookupPluginAuth("no-auth-srv"); ok {
|
||||
t.Error("expected no auth registration when Authorization header is missing")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildPluginAuthClient(t *testing.T) {
|
||||
base := transport.NewClient(nil)
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1/mcp",
|
||||
AuthHeaders: map[string]string{
|
||||
"Authorization": "Bearer sk-test-api-key",
|
||||
"X-Extra": "extra-value",
|
||||
},
|
||||
}
|
||||
|
||||
client := buildPluginAuthClient(base, srv)
|
||||
|
||||
// Should return a different client instance
|
||||
if client == base {
|
||||
t.Error("expected buildPluginAuthClient to return a new client, not the base")
|
||||
}
|
||||
|
||||
// Verify trusted domains
|
||||
if len(client.TrustedDomains) != 2 {
|
||||
t.Fatalf("TrustedDomains len = %d, want 2", len(client.TrustedDomains))
|
||||
}
|
||||
if client.TrustedDomains[0] != "dashscope.aliyuncs.com" {
|
||||
t.Errorf("TrustedDomains[0] = %q, want dashscope.aliyuncs.com", client.TrustedDomains[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildPluginAuthClientNoAuth(t *testing.T) {
|
||||
base := transport.NewClient(nil)
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
AuthHeaders: map[string]string{
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
client := buildPluginAuthClient(base, srv)
|
||||
|
||||
// Should return the base client when no Authorization header
|
||||
if client != base {
|
||||
t.Error("expected buildPluginAuthClient to return base client when no Authorization header")
|
||||
}
|
||||
}
|
||||
@@ -17,9 +17,20 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CONFIG_DIR",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "覆盖默认配置目录 (~/.dws)",
|
||||
DefaultValue: "~/.dws",
|
||||
Example: "/opt/dws/config",
|
||||
})
|
||||
}
|
||||
|
||||
// Build-time variables injected via ldflags when available.
|
||||
var (
|
||||
buildTime = "unknown"
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newConfigCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "config",
|
||||
Short: "配置管理",
|
||||
Long: "管理 DWS CLI 的配置项。查看所有支持的环境变量及其当前值。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
cmd.AddCommand(newConfigListCommand())
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "列出所有可用配置项",
|
||||
Long: "显示 DWS CLI 支持的全部环境变量配置项,包括名称、分类、描述和默认值。",
|
||||
RunE: runConfigList,
|
||||
}
|
||||
cmd.Flags().String("category", "", "按分类过滤 (core|auth|network|security|runtime|debug|external)")
|
||||
cmd.Flags().Bool("show-values", false, "显示配置项的当前实际值 (敏感信息会脱敏)")
|
||||
cmd.Flags().Bool("show-hidden", false, "包含隐藏的内部调试配置项")
|
||||
cmd.Flags().Bool("json", false, "以 JSON 格式输出")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runConfigList(cmd *cobra.Command, _ []string) error {
|
||||
category, _ := cmd.Flags().GetString("category")
|
||||
showValues, _ := cmd.Flags().GetBool("show-values")
|
||||
showHidden, _ := cmd.Flags().GetBool("show-hidden")
|
||||
jsonOut, _ := cmd.Flags().GetBool("json")
|
||||
|
||||
var items []configmeta.ConfigItem
|
||||
if category != "" {
|
||||
items = configmeta.ByCategory(configmeta.Category(category))
|
||||
} else {
|
||||
items = configmeta.All()
|
||||
}
|
||||
|
||||
if !showHidden {
|
||||
items = filterVisible(items)
|
||||
}
|
||||
|
||||
if jsonOut {
|
||||
return writeConfigJSON(cmd, items, showValues)
|
||||
}
|
||||
return writeConfigTable(cmd, items, showValues)
|
||||
}
|
||||
|
||||
func filterVisible(items []configmeta.ConfigItem) []configmeta.ConfigItem {
|
||||
out := make([]configmeta.ConfigItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
if !item.Hidden {
|
||||
out = append(out, item)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func writeConfigJSON(cmd *cobra.Command, items []configmeta.ConfigItem, showValues bool) error {
|
||||
type jsonItem struct {
|
||||
Name string `json:"name"`
|
||||
Category string `json:"category"`
|
||||
Description string `json:"description"`
|
||||
DefaultValue string `json:"default_value,omitempty"`
|
||||
Example string `json:"example,omitempty"`
|
||||
Sensitive bool `json:"sensitive,omitempty"`
|
||||
CurrentValue string `json:"current_value,omitempty"`
|
||||
IsSet bool `json:"is_set"`
|
||||
}
|
||||
|
||||
result := make([]jsonItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
ji := jsonItem{
|
||||
Name: item.Name,
|
||||
Category: string(item.Category),
|
||||
Description: item.Description,
|
||||
DefaultValue: item.DefaultValue,
|
||||
Example: item.Example,
|
||||
Sensitive: item.Sensitive,
|
||||
}
|
||||
val, ok := configmeta.Resolve(item.Name)
|
||||
ji.IsSet = ok
|
||||
if showValues && ok {
|
||||
ji.CurrentValue = val
|
||||
}
|
||||
result = append(result, ji)
|
||||
}
|
||||
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "config_list",
|
||||
"count": len(result),
|
||||
"configs": result,
|
||||
})
|
||||
}
|
||||
|
||||
func writeConfigTable(cmd *cobra.Command, items []configmeta.ConfigItem, showValues bool) error {
|
||||
w := cmd.OutOrStdout()
|
||||
|
||||
if len(items) == 0 {
|
||||
_, _ = fmt.Fprintln(w, "没有找到匹配的配置项。")
|
||||
return nil
|
||||
}
|
||||
|
||||
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
|
||||
|
||||
if showValues {
|
||||
_, _ = fmt.Fprintln(tw, "分类\t配置项\t描述\t默认值\t当前值")
|
||||
_, _ = fmt.Fprintln(tw, "────\t──────\t────\t──────\t──────")
|
||||
} else {
|
||||
_, _ = fmt.Fprintln(tw, "分类\t配置项\t描述\t默认值")
|
||||
_, _ = fmt.Fprintln(tw, "────\t──────\t────\t──────")
|
||||
}
|
||||
|
||||
for _, item := range items {
|
||||
def := item.DefaultValue
|
||||
if def == "" {
|
||||
def = "(空)"
|
||||
}
|
||||
if showValues {
|
||||
val, ok := configmeta.Resolve(item.Name)
|
||||
display := "(未设置)"
|
||||
if ok {
|
||||
display = val
|
||||
}
|
||||
_, _ = fmt.Fprintf(tw, "%s\t%s\t%s\t%s\t%s\n",
|
||||
item.Category, item.Name, item.Description, def, display)
|
||||
} else {
|
||||
_, _ = fmt.Fprintf(tw, "%s\t%s\t%s\t%s\n",
|
||||
item.Category, item.Name, item.Description, def)
|
||||
}
|
||||
}
|
||||
|
||||
_ = tw.Flush()
|
||||
_, _ = fmt.Fprintf(w, "\n共 %d 个配置项。使用 --show-values 查看当前值,--show-hidden 显示隐藏项。\n", len(items))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func seedTestConfig(t *testing.T) {
|
||||
t.Helper()
|
||||
configmeta.Reset()
|
||||
t.Cleanup(configmeta.Reset)
|
||||
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CONFIG_DIR", Category: configmeta.CategoryCore,
|
||||
Description: "覆盖默认配置目录", DefaultValue: "~/.dws",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET", Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret", Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CATALOG_FIXTURE", Category: configmeta.CategoryDebug,
|
||||
Description: "目录 Fixture 路径", Hidden: true,
|
||||
})
|
||||
}
|
||||
|
||||
func TestConfigListTable(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CONFIG_DIR") {
|
||||
t.Error("expected DWS_CONFIG_DIR in output")
|
||||
}
|
||||
if !strings.Contains(out, "DWS_CLIENT_SECRET") {
|
||||
t.Error("expected DWS_CLIENT_SECRET in output")
|
||||
}
|
||||
// Hidden items should be excluded by default
|
||||
if strings.Contains(out, "DWS_CATALOG_FIXTURE") {
|
||||
t.Error("expected DWS_CATALOG_FIXTURE to be hidden")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListShowHidden(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--show-hidden"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CATALOG_FIXTURE") {
|
||||
t.Error("expected DWS_CATALOG_FIXTURE with --show-hidden")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListCategory(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--category", "auth"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CLIENT_SECRET") {
|
||||
t.Error("expected DWS_CLIENT_SECRET for auth category")
|
||||
}
|
||||
if strings.Contains(out, "DWS_CONFIG_DIR") {
|
||||
t.Error("DWS_CONFIG_DIR should not appear for auth category")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListJSON(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--json", "--show-hidden"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var result map[string]any
|
||||
if err := json.Unmarshal(buf.Bytes(), &result); err != nil {
|
||||
t.Fatalf("invalid JSON output: %v", err)
|
||||
}
|
||||
if result["kind"] != "config_list" {
|
||||
t.Errorf("expected kind=config_list, got %v", result["kind"])
|
||||
}
|
||||
count, ok := result["count"].(float64)
|
||||
if !ok || count != 3 {
|
||||
t.Errorf("expected count=3, got %v", result["count"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListShowValues(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
t.Setenv("DWS_CONFIG_DIR", "/custom/dir")
|
||||
t.Setenv("DWS_CLIENT_SECRET", "supersecret123")
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--show-values"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "/custom/dir") {
|
||||
t.Error("expected actual value for DWS_CONFIG_DIR")
|
||||
}
|
||||
if strings.Contains(out, "supersecret123") {
|
||||
t.Error("sensitive value should be masked")
|
||||
}
|
||||
if !strings.Contains(out, "当前值") {
|
||||
t.Error("expected '当前值' column header")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListEmpty(t *testing.T) {
|
||||
configmeta.Reset()
|
||||
defer configmeta.Reset()
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "没有找到") {
|
||||
t.Error("expected empty message")
|
||||
}
|
||||
}
|
||||
@@ -157,6 +157,66 @@ func DirectRuntimeProductIDs() map[string]bool {
|
||||
return ids
|
||||
}
|
||||
|
||||
// AppendDynamicServer adds a single server descriptor to the existing
|
||||
// dynamic server registry without replacing the current entries. This
|
||||
// is used by the plugin loader to inject plugin servers alongside
|
||||
// Market-discovered servers.
|
||||
func AppendDynamicServer(server market.ServerDescriptor) {
|
||||
dynamicMu.Lock()
|
||||
defer dynamicMu.Unlock()
|
||||
|
||||
if dynamicEndpoints == nil {
|
||||
dynamicEndpoints = make(map[string]string)
|
||||
}
|
||||
if dynamicProducts == nil {
|
||||
dynamicProducts = make(map[string]bool)
|
||||
}
|
||||
if dynamicAliases == nil {
|
||||
dynamicAliases = make(map[string]string)
|
||||
}
|
||||
if dynamicToolEndpoints == nil {
|
||||
dynamicToolEndpoints = make(map[string]string)
|
||||
}
|
||||
|
||||
if server.CLI.Skip {
|
||||
return
|
||||
}
|
||||
|
||||
id := strings.TrimSpace(server.CLI.ID)
|
||||
endpoint := strings.TrimSpace(server.Endpoint)
|
||||
if id != "" && endpoint != "" {
|
||||
dynamicEndpoints[id] = endpoint
|
||||
dynamicProducts[id] = true
|
||||
}
|
||||
cmd := strings.TrimSpace(server.CLI.Command)
|
||||
if cmd != "" && cmd != id && endpoint != "" {
|
||||
dynamicEndpoints[cmd] = endpoint
|
||||
dynamicProducts[cmd] = true
|
||||
}
|
||||
for _, alias := range server.CLI.Aliases {
|
||||
alias = strings.TrimSpace(alias)
|
||||
if alias != "" && endpoint != "" {
|
||||
dynamicEndpoints[alias] = endpoint
|
||||
dynamicProducts[alias] = true
|
||||
dynamicAliases[alias] = id
|
||||
}
|
||||
}
|
||||
if endpoint != "" {
|
||||
for _, tool := range server.CLI.Tools {
|
||||
toolName := strings.TrimSpace(tool.Name)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
for toolName := range server.CLI.ToolOverrides {
|
||||
toolName = strings.TrimSpace(toolName)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeDirectRuntimeProductID(productID string) string {
|
||||
dynamicMu.RLock()
|
||||
da := dynamicAliases
|
||||
|
||||
@@ -0,0 +1,438 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/upgrade"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// checkStatus represents the outcome of a single doctor check.
|
||||
type checkStatus string
|
||||
|
||||
const (
|
||||
statusPass checkStatus = "pass"
|
||||
statusWarn checkStatus = "warn"
|
||||
statusFail checkStatus = "fail"
|
||||
)
|
||||
|
||||
// checkResult holds the outcome of a single doctor check.
|
||||
type checkResult struct {
|
||||
Name string `json:"name"`
|
||||
Status checkStatus `json:"status"`
|
||||
Message string `json:"message"`
|
||||
Hint string `json:"hint,omitempty"`
|
||||
Detail any `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
func newDoctorCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "doctor",
|
||||
Short: "环境健康检查",
|
||||
Long: "一键检查登录态、网络连通性、缓存状态和版本更新,快速定位常见问题。",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runDoctor,
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "以 JSON 格式输出")
|
||||
cmd.Flags().Int("timeout", 10, "网络检查超时时间 (秒)")
|
||||
cmd.Flags().Bool("perf", false, "额外展示最近一次性能报告")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runDoctor(cmd *cobra.Command, _ []string) error {
|
||||
jsonOut, _ := cmd.Flags().GetBool("json")
|
||||
timeout, _ := cmd.Flags().GetInt("timeout")
|
||||
if timeout <= 0 {
|
||||
timeout = 10
|
||||
}
|
||||
networkTimeout := time.Duration(timeout) * time.Second
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
checks := make([]checkResult, 0, 4)
|
||||
|
||||
authResult := doctorCheckAuth(cmd.Context(), w, jsonOut)
|
||||
checks = append(checks, authResult)
|
||||
|
||||
networkResult := doctorCheckNetwork(cmd.Context(), w, jsonOut, networkTimeout)
|
||||
checks = append(checks, networkResult)
|
||||
|
||||
cacheResult := doctorCheckCache(w, jsonOut)
|
||||
checks = append(checks, cacheResult)
|
||||
|
||||
versionResult := doctorCheckVersion(w, jsonOut, networkTimeout)
|
||||
checks = append(checks, versionResult)
|
||||
|
||||
showPerf, _ := cmd.Flags().GetBool("perf")
|
||||
if showPerf {
|
||||
perfResult := doctorCheckPerf(w, jsonOut)
|
||||
checks = append(checks, perfResult)
|
||||
}
|
||||
|
||||
pass, warn, fail := countResults(checks)
|
||||
|
||||
if jsonOut {
|
||||
result := map[string]any{
|
||||
"kind": "doctor",
|
||||
"checks": checks,
|
||||
"summary": map[string]int{
|
||||
"pass": pass,
|
||||
"warn": warn,
|
||||
"fail": fail,
|
||||
},
|
||||
}
|
||||
if showPerf {
|
||||
if report, err := LoadLatestReport(); err == nil {
|
||||
result["perf_report"] = report
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(w, result)
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "\n诊断完成: %d 项通过, %d 项警告, %d 项失败\n", pass, warn, fail)
|
||||
if fail > 0 {
|
||||
return fmt.Errorf("诊断发现 %d 项失败", fail)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── Auth check ──────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckAuth(ctx context.Context, w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查登录状态... ")
|
||||
}
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
|
||||
data, err := provider.Status()
|
||||
if err != nil || data == nil {
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "未登录"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 进行登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if data.IsAccessTokenValid() || data.IsRefreshTokenValid() {
|
||||
if !data.IsAccessTokenValid() {
|
||||
refreshCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
|
||||
_, refreshErr := provider.GetAccessToken(refreshCtx)
|
||||
cancel()
|
||||
if refreshErr != nil {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusWarn,
|
||||
Message: "Refresh Token 有效, 但自动刷新 Access Token 失败",
|
||||
Hint: "运行 dws auth login 重新登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "登录已过期"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 重新登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Network check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckNetwork(ctx context.Context, w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查网络连通性... ")
|
||||
}
|
||||
|
||||
baseURL := cli.DefaultMarketBaseURL
|
||||
httpClient := &http.Client{Timeout: timeout}
|
||||
client := market.NewClient(baseURL, httpClient)
|
||||
|
||||
start := time.Now()
|
||||
reqCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
_, err := client.FetchServers(reqCtx, 1)
|
||||
latency := time.Since(start)
|
||||
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 不可达: %v", err),
|
||||
Hint: "请检查网络连接或代理设置",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 可达 (延迟 %dms)", latency.Milliseconds()),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Cache check ─────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckCache(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查缓存状态... ")
|
||||
}
|
||||
|
||||
store := cacheStoreFromEnv()
|
||||
files, _, err := cacheDirectoryStats(store.Root)
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("缓存目录不可读: %v", err),
|
||||
Hint: "运行 dws cache clean 清理后重试",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
entries, _ := store.ListToolsCacheEntries(config.DefaultPartition)
|
||||
|
||||
if files == 0 && len(entries) == 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: "缓存为空 (首次使用)",
|
||||
Hint: "运行任意 dws 命令后将自动建立缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
staleCount := 0
|
||||
for _, e := range entries {
|
||||
if e.Freshness == cache.FreshnessStale {
|
||||
staleCount++
|
||||
}
|
||||
}
|
||||
|
||||
if staleCount > 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("%d 个文件, %d 个工具缓存, %d 个已过期", files, len(entries), staleCount),
|
||||
Hint: "运行 dws cache refresh 刷新缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("%d 个文件, %d 个工具缓存", files, len(entries))
|
||||
if len(entries) > 0 {
|
||||
msg += ", 全部新鲜"
|
||||
}
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusPass,
|
||||
Message: msg,
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Version check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckVersion(w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查版本更新... ")
|
||||
}
|
||||
|
||||
currentVer := version
|
||||
|
||||
client := upgrade.NewClient()
|
||||
latest, err := client.FetchLatestRelease()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("无法获取最新版本: %v", err),
|
||||
Hint: "请检查网络连接",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if upgrade.NeedsUpgrade(currentVer, latest.Version) {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("有新版本 (当前 %s, 最新 v%s)", ensureV(currentVer), latest.Version),
|
||||
Hint: "运行 dws upgrade 升级到最新版本",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("已是最新版本 %s", ensureV(currentVer)),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Output helpers ──────────────────────────────────────────────────────
|
||||
|
||||
func printCheckResult(w io.Writer, r checkResult) {
|
||||
icon := statusIcon(r.Status)
|
||||
fmt.Fprintf(w, "%s %s\n", icon, r.Message)
|
||||
if r.Hint != "" {
|
||||
fmt.Fprintf(w, " %s\n", r.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func statusIcon(s checkStatus) string {
|
||||
switch s {
|
||||
case statusPass:
|
||||
return "✅"
|
||||
case statusWarn:
|
||||
return "⚠️"
|
||||
case statusFail:
|
||||
return "❌"
|
||||
default:
|
||||
return "?"
|
||||
}
|
||||
}
|
||||
|
||||
func countResults(checks []checkResult) (pass, warn, fail int) {
|
||||
for _, c := range checks {
|
||||
switch c.Status {
|
||||
case statusPass:
|
||||
pass++
|
||||
case statusWarn:
|
||||
warn++
|
||||
case statusFail:
|
||||
fail++
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// ── Perf report check ──────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckPerf(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查性能报告... ")
|
||||
}
|
||||
|
||||
report, err := LoadLatestReport()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusWarn,
|
||||
Message: "未找到性能报告",
|
||||
Hint: "设置 DWS_PERF_REPORT=auto 后运行任意命令生成报告",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("报告可用 (%s, %s)", report.Command, report.Timestamp.Local().Format("2006-01-02 15:04")),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
printPerfReportSummary(w, report)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func printPerfReportSummary(w io.Writer, report *PerfReport) {
|
||||
fmt.Fprintf(w, "\n最近一次性能报告 (%s, %s):\n",
|
||||
report.Command, report.Timestamp.Local().Format("2006-01-02 15:04"))
|
||||
|
||||
for _, p := range report.Phases {
|
||||
marker := ""
|
||||
if p.Name == report.Slowest {
|
||||
marker = " ← 最慢"
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s %dms%s\n", p.Name, p.DurationMs, marker)
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s ─────────\n", "─────────────────────────")
|
||||
fmt.Fprintf(w, " %-25s %dms (框架开销 %dms)\n", "总耗时", report.TotalMs, report.OverheadMs)
|
||||
}
|
||||
|
||||
func formatLocalTime(t time.Time) string {
|
||||
if t.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return t.Local().Format("2006-01-02 15:04")
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCountResults(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
{Status: statusWarn},
|
||||
{Status: statusFail},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 1 || fail != 1 {
|
||||
t.Errorf("expected (2,1,1), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountResultsAllPass(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 0 || fail != 0 {
|
||||
t.Errorf("expected (2,0,0), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStatusIcon(t *testing.T) {
|
||||
tests := []struct {
|
||||
status checkStatus
|
||||
want string
|
||||
}{
|
||||
{statusPass, "✅"},
|
||||
{statusWarn, "⚠️"},
|
||||
{statusFail, "❌"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
got := statusIcon(tc.status)
|
||||
if got != tc.want {
|
||||
t.Errorf("statusIcon(%q) = %q, want %q", tc.status, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResult(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusFail,
|
||||
Message: "something broke",
|
||||
Hint: "try fixing it",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "❌") {
|
||||
t.Error("expected fail icon")
|
||||
}
|
||||
if !strings.Contains(out, "something broke") {
|
||||
t.Error("expected message")
|
||||
}
|
||||
if !strings.Contains(out, "try fixing it") {
|
||||
t.Error("expected hint")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResultNoHint(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusPass,
|
||||
Message: "all good",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "✅") {
|
||||
t.Error("expected pass icon")
|
||||
}
|
||||
lines := strings.Split(strings.TrimSpace(out), "\n")
|
||||
if len(lines) != 1 {
|
||||
t.Errorf("expected 1 line (no hint), got %d", len(lines))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmpty(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, false)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if !strings.Contains(r.Message, "缓存为空") {
|
||||
t.Errorf("expected empty cache message, got %q", r.Message)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmptyJSON(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, true)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if buf.Len() != 0 {
|
||||
t.Error("expected no output in JSON mode")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCommandStructure(t *testing.T) {
|
||||
cmd := newDoctorCommand()
|
||||
if cmd.Use != "doctor" {
|
||||
t.Errorf("Use = %q, want doctor", cmd.Use)
|
||||
}
|
||||
|
||||
jsonFlag := cmd.Flags().Lookup("json")
|
||||
if jsonFlag == nil {
|
||||
t.Error("expected --json flag")
|
||||
}
|
||||
timeoutFlag := cmd.Flags().Lookup("timeout")
|
||||
if timeoutFlag == nil {
|
||||
t.Error("expected --timeout flag")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckResultJSONMarshal(t *testing.T) {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
data, err := json.Marshal(r)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parsed["name"] != "auth" {
|
||||
t.Errorf("expected name=auth, got %v", parsed["name"])
|
||||
}
|
||||
if parsed["status"] != "pass" {
|
||||
t.Errorf("expected status=pass, got %v", parsed["status"])
|
||||
}
|
||||
if _, hasHint := parsed["hint"]; hasHint {
|
||||
t.Error("empty hint should be omitted")
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,7 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -35,8 +36,8 @@ type GlobalFlags struct {
|
||||
}
|
||||
|
||||
func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", "Override OAuth client ID (DingTalk AppKey)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", "Override OAuth client secret (DingTalk AppSecret)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", i18n.T("覆盖 OAuth 客户端 ID (钉钉 AppKey)"))
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", i18n.T("覆盖 OAuth 客户端密钥 (钉钉 AppSecret)"))
|
||||
cmd.PersistentFlags().BoolVar(&flags.Debug, "debug", false, "显示调试日志")
|
||||
cmd.PersistentFlags().BoolVar(&flags.DryRun, "dry-run", false, "预览操作内容,不实际执行")
|
||||
cmd.PersistentFlags().StringVar(&flags.Fields, "fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
|
||||
@@ -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.
|
||||
|
||||
package app
|
||||
|
||||
// MCPIdentityHeaders returns the same header map used for MCP HTTP requests
|
||||
// (agent identity, env trace headers, edition MergeHeaders). Intended for
|
||||
// non-MCP transports such as the A2A gateway client.
|
||||
func MCPIdentityHeaders() map[string]string {
|
||||
return resolveIdentityHeaders()
|
||||
}
|
||||
@@ -0,0 +1,577 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/fatih/color"
|
||||
)
|
||||
|
||||
const (
|
||||
// PatAuthRetryTimeout is the maximum time to wait for user authorization
|
||||
// when a PAT scope error is detected.
|
||||
PatAuthRetryTimeout = 10 * time.Minute
|
||||
|
||||
// PatAuthPollInterval is how often we poll to check if the user has
|
||||
// completed authorization.
|
||||
PatAuthPollInterval = 5 * time.Second
|
||||
)
|
||||
|
||||
// PatScopeError holds information about a missing PAT scope.
|
||||
type PatScopeError struct {
|
||||
OriginalError string
|
||||
Identity string
|
||||
ErrorType string
|
||||
Message string
|
||||
Hint string
|
||||
MissingScope string
|
||||
}
|
||||
|
||||
func (e *PatScopeError) Error() string {
|
||||
return e.OriginalError
|
||||
}
|
||||
|
||||
// patScopeRegex matches PAT-protocol scope error patterns from the API.
|
||||
// Only matches explicit scope-related keywords; generic "permission denied" or
|
||||
// "forbidden" are intentionally excluded to avoid false positives on business
|
||||
// authorization errors (e.g. mailbox access denied, 403 Forbidden).
|
||||
var patScopeRegex = regexp.MustCompile(`(?i)(missing_scope|insufficient_scope|scope.*required)`)
|
||||
|
||||
// scopeValueRegex extracts a scope identifier (e.g. "calendar:read",
|
||||
// "mail:user_mailbox.message:send") from an error message.
|
||||
// Supports multi-segment scopes with multiple colons (resource:sub:action).
|
||||
var scopeValueRegex = regexp.MustCompile(`([a-zA-Z][a-zA-Z0-9_.]*(?::[a-zA-Z][a-zA-Z0-9_.]*)+)`)
|
||||
|
||||
// identityValueRegex extracts an identity label from an error message.
|
||||
var identityValueRegex = regexp.MustCompile(`(?i)identity["\s:]+([a-zA-Z_]+)`)
|
||||
|
||||
// isPatScopeError checks if an error looks like a PAT scope/permission error
|
||||
// that can be resolved by re-authorizing with additional scopes.
|
||||
func isPatScopeError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
|
||||
// Check for missing_scope pattern in error message or hint
|
||||
if patScopeRegex.MatchString(msg) {
|
||||
return true
|
||||
}
|
||||
|
||||
var typed *apperrors.Error
|
||||
if stderrors.As(err, &typed) {
|
||||
// Check message, reason, and hint for scope-related patterns
|
||||
fullText := strings.ToLower(typed.Message + " " + typed.Reason + " " + typed.Hint)
|
||||
if typed.Category == apperrors.CategoryAuth {
|
||||
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") ||
|
||||
(strings.Contains(fullText, "scope") && strings.Contains(fullText, "required")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
// Any category with scope/permission hints
|
||||
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// extractPatScopeError parses an error to extract PAT scope details.
|
||||
func extractPatScopeError(err error) *PatScopeError {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
msg := err.Error()
|
||||
scope := ""
|
||||
|
||||
var typed *apperrors.Error
|
||||
if stderrors.As(err, &typed) {
|
||||
msg = typed.Message
|
||||
if typed.Reason != "" {
|
||||
msg += " (" + typed.Reason + ")"
|
||||
}
|
||||
}
|
||||
|
||||
// Try to extract scope value (e.g. "calendar:read") from error message.
|
||||
scopeMatch := scopeValueRegex.FindStringSubmatch(msg)
|
||||
if len(scopeMatch) > 1 {
|
||||
scope = scopeMatch[1]
|
||||
}
|
||||
|
||||
// Try to extract identity from error message.
|
||||
identity := "user"
|
||||
identityMatch := identityValueRegex.FindStringSubmatch(msg)
|
||||
if len(identityMatch) > 1 {
|
||||
identity = identityMatch[1]
|
||||
}
|
||||
|
||||
return &PatScopeError{
|
||||
OriginalError: err.Error(),
|
||||
Identity: identity,
|
||||
ErrorType: "missing_scope",
|
||||
Message: msg,
|
||||
Hint: fmt.Sprintf("run `dws auth login --scope %q` to authorize the missing scope", scope),
|
||||
MissingScope: scope,
|
||||
}
|
||||
}
|
||||
|
||||
// PrintPatAuthError prints a human-readable PAT authorization error.
|
||||
func PrintPatAuthError(w io.Writer, scopeErr *PatScopeError) {
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
cyan := color.New(color.FgCyan).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
green := color.New(color.FgGreen).SprintFunc()
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "{\n")
|
||||
fmt.Fprintf(w, " %s: %s,\n", bold("\"ok\""), "false")
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"identity\""), scopeErr.Identity)
|
||||
fmt.Fprintf(w, " %s: {\n", bold("\"error\""))
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"type\""), scopeErr.ErrorType)
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"message\""), scopeErr.Message)
|
||||
fmt.Fprintf(w, " %s: %q\n", bold("\"hint\""), scopeErr.Hint)
|
||||
fmt.Fprintf(w, " }\n")
|
||||
fmt.Fprintf(w, "}\n")
|
||||
fmt.Fprintln(w)
|
||||
|
||||
// Print authorization instructions
|
||||
fmt.Fprintf(w, "%s %s\n", green("▶"), bold("需要额外授权"))
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, " %s %s\n", dim("#"), dim("运行以下命令完成授权"))
|
||||
|
||||
if scopeErr.MissingScope != "" {
|
||||
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan(fmt.Sprintf("dws auth login --scope %q", scopeErr.MissingScope)))
|
||||
} else {
|
||||
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan("dws auth login"))
|
||||
}
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, " %s 在浏览器中打开授权链接,完成授权后重新执行命令\n", dim("ℹ"))
|
||||
fmt.Fprintln(w)
|
||||
}
|
||||
|
||||
// PrintPatAuthJSON prints a machine-readable PAT authorization error.
|
||||
func PrintPatAuthJSON(w io.Writer, scopeErr *PatScopeError) {
|
||||
payload := map[string]any{
|
||||
"ok": false,
|
||||
"identity": scopeErr.Identity,
|
||||
"error": map[string]any{
|
||||
"type": scopeErr.ErrorType,
|
||||
"message": scopeErr.Message,
|
||||
"hint": scopeErr.Hint,
|
||||
},
|
||||
}
|
||||
if scopeErr.MissingScope != "" {
|
||||
payload["missing_scope"] = scopeErr.MissingScope
|
||||
}
|
||||
|
||||
data, _ := json.MarshalIndent(payload, "", " ")
|
||||
fmt.Fprintln(w, string(data))
|
||||
}
|
||||
|
||||
// WaitForPatAuthorization polls until the user completes authorization or timeout.
|
||||
// It returns true if authorization was completed, false if timed out or cancelled.
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
yellow := color.New(color.FgYellow).SprintFunc()
|
||||
green := color.New(color.FgGreen).SprintFunc()
|
||||
red := color.New(color.FgRed).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
|
||||
timeout := PatAuthRetryTimeout
|
||||
deadline := time.Now().Add(timeout)
|
||||
pollTicker := time.NewTicker(PatAuthPollInterval)
|
||||
defer pollTicker.Stop()
|
||||
start := time.Now()
|
||||
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", yellow("⏳"), bold("等待用户授权..."))
|
||||
fmt.Fprintf(output, " %s 请在另一个终端完成 dws auth login 授权\n", dim("ℹ"))
|
||||
fmt.Fprintf(output, " %s 超时时间: %s\n", dim("⏱"), timeout)
|
||||
fmt.Fprintln(output)
|
||||
|
||||
pollCount := 0
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fmt.Fprintf(output, "%s 操作已取消\n", red("✗"))
|
||||
return false
|
||||
|
||||
case <-time.After(time.Until(deadline)):
|
||||
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", red("✗"), timeout)
|
||||
fmt.Fprintf(output, " %s 请重新执行命令\n", dim("ℹ"))
|
||||
return false
|
||||
|
||||
case <-pollTicker.C:
|
||||
pollCount++
|
||||
elapsed := time.Since(start).Truncate(time.Second)
|
||||
remaining := time.Until(deadline).Truncate(time.Second)
|
||||
|
||||
// Check if token is now valid
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err == nil && tokenData != nil {
|
||||
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
green("✓"), bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// Show polling status
|
||||
fmt.Fprintf(output, "\r%s [%d] 等待授权中... (%s 已用, %s 剩余) ",
|
||||
dim("⟳"), pollCount, elapsed, remaining)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// retryWithPatAuthRetry wraps an invocation that failed with a PAT scope error.
|
||||
// It waits for the user to complete authorization and then retries the invocation.
|
||||
func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocation executor.Invocation, scopeErr *PatScopeError, configDir string, output io.Writer) (executor.Result, error) {
|
||||
// Print the PAT error in human-readable format
|
||||
PrintPatAuthError(output, scopeErr)
|
||||
|
||||
// Wait for user to complete authorization
|
||||
authorized := WaitForPatAuthorization(ctx, configDir, output)
|
||||
if !authorized {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"等待用户授权超时",
|
||||
apperrors.WithReason("pat_auth_timeout"),
|
||||
apperrors.WithHint(fmt.Sprintf("授权超时 (%s),请重新执行命令", PatAuthRetryTimeout)),
|
||||
apperrors.WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
// Clear the token cache so the new token is loaded
|
||||
ResetRuntimeTokenCache()
|
||||
|
||||
// Retry the invocation
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", color.New(color.FgGreen).SprintFunc()("▶"),
|
||||
color.New(color.Bold).SprintFunc()("授权完成,正在重试..."))
|
||||
fmt.Fprintln(output)
|
||||
|
||||
return runner.Run(ctx, invocation)
|
||||
}
|
||||
|
||||
// ---- handlePatAuthCheck (runner.go entry point) -----------------------------
|
||||
|
||||
const (
|
||||
// patPollInterval is how often we poll the device flow status endpoint.
|
||||
patPollInterval = 2 * time.Second
|
||||
// patPollTimeout is the maximum time to wait for user authorization via device flow.
|
||||
patPollTimeout = 10 * time.Minute
|
||||
)
|
||||
|
||||
// patRetryingKey is a context key to prevent recursive PAT auth checks.
|
||||
// After APPROVED, the retry should not trigger another PAT flow.
|
||||
type patRetryingKeyType struct{}
|
||||
|
||||
var patRetryingKey = patRetryingKeyType{}
|
||||
|
||||
// IsPatRetrying returns true if the current context is already in a PAT retry.
|
||||
func IsPatRetrying(ctx context.Context) bool {
|
||||
v, _ := ctx.Value(patRetryingKey).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
// handlePatAuthCheck is called by runner.executeInvocation when a PAT
|
||||
// authorization error is detected. It injects the server-assigned clientId
|
||||
// as x-robot-uid header, prints authorization details, opens the browser,
|
||||
// polls the device flow endpoint until the user authorizes, and retries the
|
||||
// original invocation on success.
|
||||
func handlePatAuthCheck(
|
||||
ctx context.Context,
|
||||
r *runtimeRunner,
|
||||
invocation executor.Invocation,
|
||||
patErr *apperrors.PATError,
|
||||
configDir string,
|
||||
output io.Writer,
|
||||
) (executor.Result, error) {
|
||||
// Parse authorization details from PATError.RawJSON.
|
||||
var patData struct {
|
||||
Code string `json:"code"`
|
||||
Data struct {
|
||||
Desc string `json:"desc"`
|
||||
FlowID string `json:"flowId"`
|
||||
URI string `json:"uri"`
|
||||
ClientID string `json:"clientId"`
|
||||
ClientSecret string `json:"clientSecret"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(patErr.RawJSON), &patData); err != nil {
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
slog.Debug("PAT auth check",
|
||||
"clientId", patData.Data.ClientID,
|
||||
"flowId", patData.Data.FlowID,
|
||||
"hasSecret", patData.Data.ClientSecret != "",
|
||||
)
|
||||
|
||||
// Inject clientId/clientSecret from PAT response as runtime credentials
|
||||
// so that subsequent device flow auth uses the server-assigned app identity.
|
||||
if patData.Data.ClientID != "" {
|
||||
if patData.Data.ClientSecret != "" {
|
||||
// When both clientId and clientSecret are provided, use direct mode
|
||||
// (DingTalk API) rather than MCP proxy — the MCP proxy does not hold
|
||||
// the secret for this particular app.
|
||||
authpkg.SetClientID(patData.Data.ClientID)
|
||||
authpkg.SetClientSecret(patData.Data.ClientSecret)
|
||||
} else {
|
||||
// No clientSecret — rely on MCP proxy to manage the secret server-side.
|
||||
authpkg.SetClientIDFromMCP(patData.Data.ClientID)
|
||||
}
|
||||
|
||||
// Persist clientId (and optionally secret) to ~/.dws/app.json so that
|
||||
// future process invocations can load it at startup and populate
|
||||
// DWS_CLIENT_ID env before the first MCP request.
|
||||
appCfg := &authpkg.AppConfig{
|
||||
ClientID: patData.Data.ClientID,
|
||||
}
|
||||
if patData.Data.ClientSecret != "" {
|
||||
appCfg.ClientSecret = authpkg.PlainSecret(patData.Data.ClientSecret)
|
||||
}
|
||||
if err := authpkg.SaveAppConfig(configDir, appCfg); err != nil {
|
||||
slog.Warn("failed to persist app config from PAT", "error", err)
|
||||
fmt.Fprintf(output, " \u26a0 保存应用配置失败: %v (下次启动可能需要重新授权)\n", err)
|
||||
}
|
||||
}
|
||||
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
cyan := color.New(color.FgCyan).SprintFunc()
|
||||
greenFn := color.New(color.FgGreen).SprintFunc()
|
||||
yellowFn := color.New(color.FgYellow).SprintFunc()
|
||||
redFn := color.New(color.FgRed).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("需要 PAT 授权"))
|
||||
if patData.Data.Desc != "" {
|
||||
fmt.Fprintf(output, " %s %s\n", dim("ℹ"), patData.Data.Desc)
|
||||
}
|
||||
if patData.Data.URI != "" {
|
||||
fmt.Fprintf(output, " %s %s\n\n", dim("🔗"), cyan(patData.Data.URI))
|
||||
// Best-effort browser open.
|
||||
_ = tryOpenBrowser(patData.Data.URI)
|
||||
}
|
||||
|
||||
// If no flowId, we can't poll — fall back to returning PATError for host-app.
|
||||
if patData.Data.FlowID == "" {
|
||||
fmt.Fprintln(output)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
// Poll the device flow status until user authorizes, rejects, or timeout.
|
||||
fmt.Fprintf(output, "%s %s\n", yellowFn("⏳"), bold("等待用户授权..."))
|
||||
fmt.Fprintf(output, " %s 请在浏览器中完成授权,超时时间: %s\n", dim("ℹ"), patPollTimeout)
|
||||
fmt.Fprintln(output)
|
||||
|
||||
pollCtx, cancel := context.WithTimeout(ctx, patPollTimeout)
|
||||
defer cancel()
|
||||
|
||||
status, authCode, err := pollPatDeviceFlow(pollCtx, patData.Data.FlowID, configDir, output)
|
||||
if err != nil {
|
||||
fmt.Fprintf(output, "%s 轮询授权状态失败: %v\n", redFn("✗"), err)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
switch status {
|
||||
case authpkg.StatusApproved:
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("✓"), bold("授权成功!"))
|
||||
fmt.Fprintln(output)
|
||||
|
||||
// Exchange authCode for a fresh access token (mirrors device_flow loginOnce).
|
||||
if authCode != "" {
|
||||
slog.Debug("PAT retry: exchanging authCode for token", "hasCode", true)
|
||||
tokenData, exchErr := authpkg.ExchangeCodeForToken(ctx, configDir, authCode)
|
||||
if exchErr != nil {
|
||||
slog.Warn("PAT retry: exchangeCode failed, retrying with existing token", "error", exchErr)
|
||||
fmt.Fprintf(output, " %s 换取新 token 失败: %v (将使用现有凭证重试)\n", yellowFn("⚠"), exchErr)
|
||||
} else {
|
||||
if err := authpkg.SaveTokenData(configDir, tokenData); err != nil {
|
||||
slog.Warn("PAT retry: failed to save new token", "error", err)
|
||||
fmt.Fprintf(output, " %s 保存新 token 失败: %v\n", yellowFn("⚠"), err)
|
||||
} else {
|
||||
slog.Debug("PAT retry: token refreshed and saved")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Clear token cache so the new credentials take effect.
|
||||
ResetRuntimeTokenCache()
|
||||
|
||||
// Workaround: brief delay to let server-side authorization state propagate
|
||||
// before retrying. Without this the retry may use stale credentials.
|
||||
slog.Debug("PAT retry: waiting for server-side state propagation", "delay", "1s")
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// Retry the original invocation with pat-retrying flag to prevent recursion.
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("授权完成,正在重试..."))
|
||||
fmt.Fprintln(output)
|
||||
slog.Debug("PAT retry: identity env check",
|
||||
"DWS_CLIENT_ID", os.Getenv("DWS_CLIENT_ID"),
|
||||
)
|
||||
retryCtx := context.WithValue(ctx, patRetryingKey, true)
|
||||
return r.Run(retryCtx, invocation)
|
||||
|
||||
case authpkg.StatusRejected:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("用户已拒绝授权"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"用户已拒绝授权",
|
||||
apperrors.WithReason("pat_auth_rejected"),
|
||||
apperrors.WithHint("用户在浏览器中拒绝了授权请求,请重新执行命令。"),
|
||||
)
|
||||
|
||||
case authpkg.StatusExpired:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("授权超时"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"授权超时",
|
||||
apperrors.WithReason("pat_auth_expired"),
|
||||
apperrors.WithHint("授权链接已过期,请重新执行命令。"),
|
||||
)
|
||||
|
||||
case authpkg.StatusCancelled:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("操作已取消"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"操作已取消",
|
||||
apperrors.WithReason("pat_auth_cancelled"),
|
||||
apperrors.WithHint("用户取消了授权操作。"),
|
||||
)
|
||||
|
||||
default:
|
||||
fmt.Fprintf(output, "%s 未知授权状态: %s\n", redFn("✗"), status)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
}
|
||||
|
||||
// pollPatDeviceFlow polls the PAT device flow status endpoint until a terminal
|
||||
// state (APPROVED/REJECTED/EXPIRED) is reached or the context is cancelled.
|
||||
// Returns the final status string and the authCode (non-empty only on APPROVED).
|
||||
func pollPatDeviceFlow(ctx context.Context, flowID string, configDir string, output io.Writer) (string, string, error) {
|
||||
pollURL := fmt.Sprintf("%s%s?flowId=%s",
|
||||
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
|
||||
|
||||
// Load user access token for the poll request header.
|
||||
var accessToken string
|
||||
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
|
||||
accessToken = tokenData.AccessToken
|
||||
}
|
||||
|
||||
// Use a client that does NOT follow redirects, so we can detect SSO 302.
|
||||
noRedirectClient := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(patPollInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
pollCount := 0
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if ctx.Err() == context.Canceled {
|
||||
return authpkg.StatusCancelled, "", nil
|
||||
}
|
||||
return authpkg.StatusExpired, "", nil
|
||||
case <-ticker.C:
|
||||
pollCount++
|
||||
fmt.Fprintf(output, "\r%s [%d] 等待授权中... ", dim("⟳"), pollCount)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, pollURL, nil)
|
||||
if err != nil {
|
||||
slog.Debug("PAT poll: failed to create request", "error", err)
|
||||
continue
|
||||
}
|
||||
if accessToken != "" {
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
}
|
||||
resp, err := noRedirectClient.Do(req)
|
||||
if err != nil {
|
||||
slog.Debug("PAT poll: request failed", "error", err)
|
||||
continue // transient network error, keep polling
|
||||
}
|
||||
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
|
||||
// If we got a redirect (302/301), SSO gateway intercepted — skip JSON parse.
|
||||
if resp.StatusCode == http.StatusFound || resp.StatusCode == http.StatusMovedPermanently {
|
||||
continue
|
||||
}
|
||||
|
||||
var pollResp authpkg.DevicePollResponse
|
||||
if err := json.Unmarshal(bodyBytes, &pollResp); err != nil {
|
||||
slog.Debug("PAT poll: failed to parse response", "error", err, "body", string(bodyBytes))
|
||||
continue
|
||||
}
|
||||
|
||||
status := authpkg.ParseDeviceFlowStatus(pollResp.Data.Status, pollResp.Success)
|
||||
switch status {
|
||||
case authpkg.StatusApproved:
|
||||
fmt.Fprintln(output) // clear the polling line
|
||||
return status, pollResp.Data.AuthCode, nil
|
||||
case authpkg.StatusRejected, authpkg.StatusExpired:
|
||||
fmt.Fprintln(output) // clear the polling line
|
||||
return status, "", nil
|
||||
case authpkg.StatusPending:
|
||||
// keep polling
|
||||
default:
|
||||
// ParseDeviceFlowStatus normalizes empty+!success to EXPIRED,
|
||||
// so this branch handles truly unknown statuses.
|
||||
fmt.Fprintln(output)
|
||||
return status, "", nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tryOpenBrowser opens url in the default browser; errors are silently ignored.
|
||||
func tryOpenBrowser(url string) error {
|
||||
var cmd *exec.Cmd
|
||||
switch runtime.GOOS {
|
||||
case "darwin":
|
||||
cmd = exec.Command("open", url)
|
||||
case "linux":
|
||||
cmd = exec.Command("xdg-open", url)
|
||||
case "windows":
|
||||
cmd = exec.Command("cmd", "/c", "start", url)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
return cmd.Start()
|
||||
}
|
||||
@@ -0,0 +1,606 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
|
||||
func TestIsPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected missing_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_PlainString(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user lacks required scope",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user lacks required scope",
|
||||
}
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected plain string with missing_scope to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_NotScopeError(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewValidation("invalid parameter")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("expected validation error NOT to be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if isPatScopeError(nil) {
|
||||
t.Fatal("nil error should not be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_WithReason(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("API error",
|
||||
apperrors.WithReason("missing_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected error with missing_scope reason to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_InsufficientScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("insufficient_scope for resource",
|
||||
apperrors.WithReason("insufficient_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected insufficient_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.ErrorType != "missing_scope" {
|
||||
t.Errorf("expected error type 'missing_scope', got %q", scopeErr.ErrorType)
|
||||
}
|
||||
if !strings.Contains(scopeErr.Hint, "dws auth login") {
|
||||
t.Errorf("expected hint to contain 'dws auth login', got %q", scopeErr.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ExtractsScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user needs calendar:read",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user needs calendar:read",
|
||||
}
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:read" {
|
||||
t.Errorf("expected MissingScope 'calendar:read', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthError_HumanReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:user_mailbox.message:send",
|
||||
Hint: "run `dws auth login --scope \"mail:user_mailbox.message:send\"` to authorize",
|
||||
MissingScope: "mail:user_mailbox.message:send",
|
||||
}
|
||||
PrintPatAuthError(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "missing_scope") {
|
||||
t.Errorf("expected output to contain 'missing_scope', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "dws auth login") {
|
||||
t.Errorf("expected output to contain 'dws auth login', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "需要额外授权") {
|
||||
t.Errorf("expected output to contain Chinese auth prompt, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthJSON_MachineReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:send",
|
||||
Hint: "run dws auth login --scope mail:send",
|
||||
MissingScope: "mail:send",
|
||||
}
|
||||
PrintPatAuthJSON(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, `"ok": false`) {
|
||||
t.Errorf("expected JSON to contain ok: false, got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, `"missing_scope": "mail:send"`) {
|
||||
t.Errorf("expected JSON to contain missing_scope, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_BusinessPermissionDenied(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Generic business "permission denied" should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("User has no permission to access this mailbox, permission denied")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("generic 'permission denied' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_GenericForbidden(t *testing.T) {
|
||||
t.Parallel()
|
||||
// HTTP 403 Forbidden should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("403 Forbidden")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("'403 Forbidden' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ComplexScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "mail:user_mailbox.message:send" {
|
||||
t.Errorf("expected MissingScope 'mail:user_mailbox.message:send', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatScopeError_Error(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "test error message",
|
||||
}
|
||||
if err.Error() != "test error message" {
|
||||
t.Errorf("expected Error() to return OriginalError, got %q", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow integration tests — httptest mock covering four terminal
|
||||
// states: APPROVED, REJECTED, EXPIRED, CANCELLED (ctx cancel).
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// setupPollServer creates an httptest server that responds to
|
||||
// /cli/oauth/device/poll?flowId=<fid> with the given status sequence.
|
||||
// It also writes the server URL into a temp DWS_CONFIG_DIR/mcp_url so that
|
||||
// GetMCPBaseURL() returns the test server address.
|
||||
func setupPollServer(t *testing.T, statuses []authpkg.DevicePollResponse) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var callCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
idx := int(callCount.Add(1)) - 1
|
||||
if idx >= len(statuses) {
|
||||
idx = len(statuses) - 1
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(statuses[idx])
|
||||
}))
|
||||
|
||||
// Write mcp_url so GetMCPBaseURL picks up the test server.
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Approved(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "APPROVED", AuthCode: "code123"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-1", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED, got %q", status)
|
||||
}
|
||||
if authCode != "code123" {
|
||||
t.Errorf("expected authCode 'code123', got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Rejected(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "REJECTED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-2", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "REJECTED" {
|
||||
t.Errorf("expected REJECTED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for REJECTED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Expired(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "EXPIRED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-3", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for EXPIRED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Cancelled(t *testing.T) {
|
||||
// Server always returns PENDING so context cancellation is the only exit.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Cancel immediately after first poll tick.
|
||||
go func() {
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
cancel()
|
||||
}()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-4", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "CANCELLED" {
|
||||
t.Errorf("expected CANCELLED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for CANCELLED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// IsPatRetrying tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsPatRetrying_Default(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
if IsPatRetrying(ctx) {
|
||||
t.Fatal("expected false for plain context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatRetrying_WithValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.WithValue(context.Background(), patRetryingKey, true)
|
||||
if !IsPatRetrying(ctx) {
|
||||
t.Fatal("expected true when pat retry key is set")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPollPatDeviceFlow_ServerErrorFallback(t *testing.T) {
|
||||
// When server returns success=false with empty status, should treat as EXPIRED.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: ""}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-err", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED for server error fallback, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for server error, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_RedirectSkipped(t *testing.T) {
|
||||
// When server returns 302 (SSO redirect), poll should continue until real response.
|
||||
var callCount int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
callCount++
|
||||
if callCount <= 1 {
|
||||
// First call: simulate SSO redirect
|
||||
w.Header().Set("Location", "https://sso.example.com")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
return
|
||||
}
|
||||
// Second call: return APPROVED
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
resp := authpkg.DevicePollResponse{
|
||||
Success: true,
|
||||
Data: authpkg.DevicePollData{Status: "APPROVED"},
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, _, err := pollPatDeviceFlow(ctx, "flow-redirect", tmpDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED after redirect, got %q", status)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// extractPatScopeError edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestExtractPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := extractPatScopeError(nil); got != nil {
|
||||
t.Fatalf("expected nil for nil error, got %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_WithIdentity(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth(`insufficient_scope: identity "app_user" needs calendar:write`)
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.Identity != "app_user" {
|
||||
t.Errorf("expected Identity 'app_user', got %q", scopeErr.Identity)
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:write" {
|
||||
t.Errorf("expected MissingScope 'calendar:write', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// handlePatAuthCheck integration tests — cover the main orchestrator with
|
||||
// mock runner + httptest poll server for APPROVED, REJECTED, EmptyFlowID.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// mockRunner is a simple executor.Runner for testing handlePatAuthCheck.
|
||||
type mockRunner struct {
|
||||
runFunc func(ctx context.Context, inv executor.Invocation) (executor.Result, error)
|
||||
}
|
||||
|
||||
func (m *mockRunner) Run(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
return m.runFunc(ctx, inv)
|
||||
}
|
||||
|
||||
// setupHandlePATServer creates an httptest server for handlePatAuthCheck tests.
|
||||
// It responds to device poll requests with the given status after the first poll.
|
||||
func setupHandlePATServer(t *testing.T, terminalStatus string, authCode string) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var pollCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.Contains(r.URL.Path, "/cli/oauth/device/poll") {
|
||||
idx := int(pollCount.Add(1)) - 1
|
||||
var resp authpkg.DevicePollResponse
|
||||
if idx == 0 {
|
||||
resp = authpkg.DevicePollResponse{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}}
|
||||
} else {
|
||||
resp = authpkg.DevicePollResponse{
|
||||
Success: terminalStatus == "APPROVED",
|
||||
Data: authpkg.DevicePollData{Status: terminalStatus, AuthCode: authCode},
|
||||
}
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func makePATErrorJSON(flowID, clientID string) string {
|
||||
type patData struct {
|
||||
Desc string `json:"desc"`
|
||||
FlowID string `json:"flowId"`
|
||||
URI string `json:"uri"`
|
||||
ClientID string `json:"clientId"`
|
||||
}
|
||||
payload := struct {
|
||||
Code string `json:"code"`
|
||||
Data patData `json:"data"`
|
||||
}{
|
||||
Code: "AGENT_CODE_NOT_EXISTS",
|
||||
Data: patData{
|
||||
Desc: "test auth",
|
||||
FlowID: flowID,
|
||||
URI: "", // empty to avoid opening browser in test
|
||||
ClientID: clientID,
|
||||
},
|
||||
}
|
||||
data, _ := json.Marshal(payload)
|
||||
return string(data)
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Approved(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "APPROVED", "test-auth-code")
|
||||
defer server.Close()
|
||||
|
||||
var retryCalled bool
|
||||
var retryHasKey bool
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
retryCalled = true
|
||||
retryHasKey = IsPatRetrying(ctx)
|
||||
return executor.Result{Response: map[string]any{"ok": true}}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-approved", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !retryCalled {
|
||||
t.Fatal("expected mock runner to be called for retry")
|
||||
}
|
||||
if !retryHasKey {
|
||||
t.Fatal("expected retry context to have patRetryingKey")
|
||||
}
|
||||
// Verify SetClientIDFromMCP was called with the PAT response clientId.
|
||||
if cid := authpkg.ClientID(); cid != "test-client-id" {
|
||||
t.Errorf("expected ClientID 'test-client-id', got %q", cid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Rejected(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "REJECTED", "")
|
||||
defer server.Close()
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called on REJECTED")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-rejected", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for REJECTED")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "用户已拒绝授权") {
|
||||
t.Errorf("expected rejection error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_EmptyFlowID_FallsBackToPATError(t *testing.T) {
|
||||
// No poll server needed — empty flowId means no polling, return PATError directly.
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called when flowId is empty")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, tmpDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected PATError when flowId is empty")
|
||||
}
|
||||
// Should return the original PATError.
|
||||
if _, ok := err.(*apperrors.PATError); !ok {
|
||||
t.Errorf("expected *PATError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,664 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newPluginCommand() *cobra.Command {
|
||||
pluginCmd := newPlaceholderParent("plugin", i18n.T("插件管理"))
|
||||
|
||||
pluginCmd.AddCommand(
|
||||
newPluginListCommand(),
|
||||
newPluginInstallCommand(),
|
||||
newPluginInfoCommand(),
|
||||
newPluginEnableCommand(),
|
||||
newPluginDisableCommand(),
|
||||
newPluginRemoveCommand(),
|
||||
newPluginValidateCommand(),
|
||||
newPluginCreateCommand(),
|
||||
newPluginDevCommand(),
|
||||
newPluginConfigCommand(),
|
||||
newPluginBuildCommand(),
|
||||
)
|
||||
|
||||
return pluginCmd
|
||||
}
|
||||
|
||||
func newPluginListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("列出已安装的插件"),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
if wantJSON {
|
||||
return output.WriteJSON(cmd.OutOrStdout(), plugins)
|
||||
}
|
||||
|
||||
if len(plugins) == 0 {
|
||||
fmt.Fprintln(cmd.OutOrStdout(), "No plugins installed.")
|
||||
return nil
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
"NAME", "VERSION", "TYPE", "STATUS", "DESCRIPTION")
|
||||
fmt.Fprintln(w, strings.Repeat("-", 85))
|
||||
for _, p := range plugins {
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
p.Name, p.Version, p.Type, statusStr(p.Enabled), p.Description)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install",
|
||||
Short: i18n.T("安装插件"),
|
||||
Example: ` dws plugin install --dir ./conference
|
||||
dws plugin install --git https://github.com/DingTalk-Real-AI/conference.git`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dirPath, _ := cmd.Flags().GetString("dir")
|
||||
gitURL, _ := cmd.Flags().GetString("git")
|
||||
|
||||
if dirPath == "" && gitURL == "" {
|
||||
return apperrors.NewValidation("specify install source: --dir <path> or --git <url>")
|
||||
}
|
||||
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if gitURL != "" {
|
||||
p, err := loader.InstallFromGit(gitURL)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
}
|
||||
|
||||
p, err := loader.InstallFromDir(dirPath)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("dir", "", "Install from a local directory")
|
||||
cmd.Flags().String("git", "", "Install from a Git repository")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInfoCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "info <name>",
|
||||
Short: i18n.T("查看插件详情"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
for _, p := range plugins {
|
||||
if p.Name == name {
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Name: %s\n", p.Name)
|
||||
fmt.Fprintf(w, "Version: %s\n", p.Version)
|
||||
fmt.Fprintf(w, "Type: %s\n", p.Type)
|
||||
fmt.Fprintf(w, "Status: %s\n", statusStr(p.Enabled))
|
||||
fmt.Fprintf(w, "Path: %s\n", p.Path)
|
||||
if p.Description != "" {
|
||||
fmt.Fprintf(w, "Description: %s\n", p.Description)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found", name))
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginEnableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "enable <name>",
|
||||
Short: i18n.T("启用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], true); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s enabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginDisableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "disable <name>",
|
||||
Short: i18n.T("禁用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], false); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s disabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginRemoveCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "remove <name>",
|
||||
Short: i18n.T("卸载已安装的插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
// Stop stdio clients before removing to release file locks
|
||||
StopStdioClientsByPlugin(args[0])
|
||||
keepData, _ := cmd.Flags().GetBool("keep-data")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.RemovePlugin(args[0], keepData); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s removed.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("keep-data", false, "Keep plugin data directory")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginValidateCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "validate <dir>",
|
||||
Short: i18n.T("校验 plugin.json"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
m, err := plugin.ParseManifest(dir + "/plugin.json")
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("parse failed: %v", err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Valid: %s (%s)\n", m.Name, m.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginCreateCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create <name>",
|
||||
Short: i18n.T("脚手架生成新插件目录"),
|
||||
Example: ` dws plugin create my-tool
|
||||
dws plugin create my-tool --description "My awesome tool"`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
desc, _ := cmd.Flags().GetString("description")
|
||||
pluginType := "user"
|
||||
|
||||
// Validate name format
|
||||
m := &plugin.Manifest{Name: name, Version: "0.1.0", Type: pluginType}
|
||||
if err := m.Validate(""); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin name: %v", err))
|
||||
}
|
||||
|
||||
dir := filepath.Join(".", name)
|
||||
if _, err := os.Stat(dir); err == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("directory %q already exists", dir))
|
||||
}
|
||||
|
||||
// Create directory structure
|
||||
dirs := []string{
|
||||
dir,
|
||||
filepath.Join(dir, "skills", name),
|
||||
filepath.Join(dir, "hooks"),
|
||||
}
|
||||
for _, d := range dirs {
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
|
||||
}
|
||||
}
|
||||
|
||||
// Write plugin.json
|
||||
pluginJSON := fmt.Sprintf(`{
|
||||
"name": %q,
|
||||
"version": "0.1.0",
|
||||
"description": %q,
|
||||
"type": %q,
|
||||
"minCLIVersion": %q,
|
||||
"mcpServers": {
|
||||
%q: {
|
||||
"type": "stdio",
|
||||
"command": "${DWS_PLUGIN_ROOT}/bin/server",
|
||||
"args": []
|
||||
}
|
||||
},
|
||||
"build": {
|
||||
"command": "echo 'TODO: replace with your build command, e.g.: bun build --compile src/server.ts --outfile bin/server'",
|
||||
"output": "bin/server"
|
||||
},
|
||||
"skills": "./skills/",
|
||||
"hooks": "./hooks/hooks.json"
|
||||
}
|
||||
`, name, desc, pluginType, RawVersion(), name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(pluginJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write plugin.json: %v", err))
|
||||
}
|
||||
|
||||
// Write SKILL.md template
|
||||
skillMD := fmt.Sprintf(`---
|
||||
name: %s
|
||||
description: %s
|
||||
cli_version: ">=%s"
|
||||
---
|
||||
|
||||
# %s
|
||||
|
||||
## Intent Recognition
|
||||
|
||||
Use this skill when the user mentions:
|
||||
- TODO: add your intent keywords here
|
||||
|
||||
## Command Decision Tree
|
||||
|
||||
| User Intent | Command | Required Parameters |
|
||||
|-------------|---------|---------------------|
|
||||
| TODO | `+"`dws %s <sub-command>`"+` | `+"`--param`"+` |
|
||||
|
||||
## Parameter Rules
|
||||
|
||||
### TODO: parameter type
|
||||
- Format description
|
||||
- Conversion rules
|
||||
`, name, desc, RawVersion(), name, name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "skills", name, "SKILL.md"), []byte(skillMD), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write SKILL.md: %v", err))
|
||||
}
|
||||
|
||||
// Write hooks.json template
|
||||
hooksJSON := `{
|
||||
"hooks": []
|
||||
}
|
||||
`
|
||||
if err := os.WriteFile(filepath.Join(dir, "hooks", "hooks.json"), []byte(hooksJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write hooks.json: %v", err))
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Created plugin scaffold at ./%s/\n", name)
|
||||
fmt.Fprintf(w, " %s/\n", name)
|
||||
fmt.Fprintf(w, " ├── plugin.json\n")
|
||||
fmt.Fprintf(w, " ├── skills/%s/SKILL.md\n", name)
|
||||
fmt.Fprintf(w, " └── hooks/hooks.json\n")
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "Next steps:\n")
|
||||
fmt.Fprintf(w, " 1. Edit plugin.json to configure your MCP servers\n")
|
||||
fmt.Fprintf(w, " 2. Edit skills/%s/SKILL.md to describe your commands\n", name)
|
||||
fmt.Fprintf(w, " 3. Run: dws plugin validate ./%s\n", name)
|
||||
fmt.Fprintf(w, " 4. Run: dws plugin dev ./%s\n", name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("description", "", "Plugin description")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginDevCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "dev <dir>",
|
||||
Short: i18n.T("将本地目录注册为开发态插件"),
|
||||
Long: `Registers a plugin from a local source directory for development.
|
||||
The plugin is loaded directly from the source directory on next CLI invocation,
|
||||
without copying files to ~/.dws/plugins/. Use 'dws plugin dev --off <name>'
|
||||
to unregister.`,
|
||||
Example: ` dws plugin dev ./my-tool
|
||||
dws plugin dev --off my-tool`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
off, _ := cmd.Flags().GetBool("off")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if off {
|
||||
// Unregister dev plugin
|
||||
name := args[0]
|
||||
if err := loader.UnregisterDevPlugin(name); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q unregistered.\n", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Register dev plugin
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
// Validate the plugin first
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
|
||||
if err := loader.RegisterDevPlugin(m.Name, absDir); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to register: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q registered from %s\n", m.Name, absDir)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "It will be loaded on next dws invocation.\n")
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "To unregister: dws plugin dev --off %s\n", m.Name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("off", false, "Unregister a dev plugin")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigCommand() *cobra.Command {
|
||||
configCmd := newPlaceholderParent("config", i18n.T("管理插件配置"))
|
||||
configCmd.AddCommand(
|
||||
newPluginConfigSetCommand(),
|
||||
newPluginConfigGetCommand(),
|
||||
newPluginConfigListCommand(),
|
||||
newPluginConfigUnsetCommand(),
|
||||
)
|
||||
return configCmd
|
||||
}
|
||||
|
||||
func newPluginConfigSetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "set <plugin-name> <key> <value>",
|
||||
Short: i18n.T("设置插件配置项"),
|
||||
Long: `Persistently set a configuration value for a plugin.
|
||||
The value is stored in ~/.dws/settings.json and automatically injected
|
||||
as an environment variable when the plugin is loaded.
|
||||
|
||||
Environment variables set by the user (e.g. via export) take precedence
|
||||
over values stored in settings.json.`,
|
||||
Example: ` dws plugin config set demo-devtool DASHSCOPE_API_KEY sk-xxx
|
||||
dws plugin config set my-plugin API_ENDPOINT https://api.example.com`,
|
||||
Args: cobra.ExactArgs(3),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key, value := args[0], args[1], args[2]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// Validate that the plugin exists.
|
||||
plugins := loader.ListInstalled()
|
||||
found := false
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found; use 'dws plugin list' to see installed plugins", pluginName))
|
||||
}
|
||||
|
||||
loader.SetPluginConfig(pluginName, key, value)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config saved: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigGetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "get <plugin-name> <key>",
|
||||
Short: i18n.T("读取插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
val, ok := loader.GetPluginConfig(pluginName, key)
|
||||
if !ok {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintln(cmd.OutOrStdout(), val)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list <plugin-name>",
|
||||
Short: i18n.T("列出插件所有配置项"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
configs := loader.ListPluginConfig(pluginName)
|
||||
|
||||
// Also load the plugin manifest to show declared userConfig keys.
|
||||
declaredKeys := loadDeclaredUserConfig(loader, pluginName)
|
||||
|
||||
if wantJSON {
|
||||
result := make(map[string]any)
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
if sensitive {
|
||||
result[k] = maskSensitiveValue(v)
|
||||
} else {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
// Include declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
entry := map[string]any{
|
||||
"value": nil,
|
||||
"description": ci.Description,
|
||||
"required": ci.Default == "",
|
||||
}
|
||||
result[k] = entry
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "plugin_config",
|
||||
"plugin": pluginName,
|
||||
"config": result,
|
||||
})
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
if len(configs) == 0 && len(declaredKeys) == 0 {
|
||||
fmt.Fprintf(w, "No configuration for plugin %q.\n", pluginName)
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "Configuration for %s:\n\n", pluginName)
|
||||
|
||||
// Show set values.
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
displayVal := v
|
||||
if sensitive {
|
||||
displayVal = maskSensitiveValue(v)
|
||||
}
|
||||
fmt.Fprintf(w, " %s = %s\n", k, displayVal)
|
||||
}
|
||||
|
||||
// Show declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
desc := ""
|
||||
if ci.Description != "" {
|
||||
desc = " # " + ci.Description
|
||||
}
|
||||
fmt.Fprintf(w, " %s = (not set)%s\n", k, desc)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigUnsetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "unset <plugin-name> <key>",
|
||||
Short: i18n.T("删除插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if !loader.UnsetPluginConfig(pluginName, key) {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config removed: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// loadDeclaredUserConfig loads the userConfig section from a plugin's manifest.
|
||||
func loadDeclaredUserConfig(loader *plugin.Loader, pluginName string) map[string]plugin.ConfigItem {
|
||||
plugins := loader.ListInstalled()
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
m, err := plugin.ParseManifest(filepath.Join(p.Path, "plugin.json"))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return m.UserConfig
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// maskSensitiveValue masks a sensitive value, showing only the first 4
|
||||
// and last 2 characters for values longer than 8 characters.
|
||||
func maskSensitiveValue(value string) string {
|
||||
if len(value) <= 8 {
|
||||
return strings.Repeat("*", len(value))
|
||||
}
|
||||
return value[:4] + strings.Repeat("*", len(value)-6) + value[len(value)-2:]
|
||||
}
|
||||
|
||||
func newPluginBuildCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "build <dir>",
|
||||
Short: i18n.T("将插件 stdio server 编译为原生二进制"),
|
||||
Long: `Runs the build command declared in plugin.json to compile the
|
||||
plugin's server into a single executable. This ensures plugin users
|
||||
don't need any language runtime (Node.js, Python, etc.) installed.
|
||||
|
||||
The build configuration is read from the "build" field in plugin.json:
|
||||
|
||||
{
|
||||
"build": {
|
||||
"command": "bun build --compile src/server.ts --outfile bin/server",
|
||||
"output": "bin/server"
|
||||
}
|
||||
}`,
|
||||
Example: ` dws plugin build ./my-plugin
|
||||
dws plugin build .`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
|
||||
if m.Build == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf(
|
||||
"plugin %q has no \"build\" field in plugin.json.\n"+
|
||||
"Add a build config, e.g.:\n\n"+
|
||||
" \"build\": {\n"+
|
||||
" \"command\": \"bun build --compile src/server.js --outfile bin/server\",\n"+
|
||||
" \"output\": \"bin/server\"\n"+
|
||||
" }", m.Name))
|
||||
}
|
||||
|
||||
if err := plugin.BuildPlugin(absDir); err != nil {
|
||||
return apperrors.NewInternal(err.Error())
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Build succeeded: %s\n", m.Build.Output)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func statusStr(enabled bool) string {
|
||||
if enabled {
|
||||
return "enabled"
|
||||
}
|
||||
return "disabled"
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
// TestSharedCacheStoreConcurrentSaveTools verifies that a single *cache.Store
|
||||
// instance is safe for goroutines saving tool snapshots concurrently, as long
|
||||
// as each goroutine targets a distinct (partition, serverKey). This mirrors
|
||||
// the real plugin discovery path where each goroutine owns one plugin/server.
|
||||
//
|
||||
// Each call serializes to its own "<key>.json.tmp" file followed by a
|
||||
// rename(2) to the final path, so concurrent writers targeting distinct keys
|
||||
// never collide. The invariant asserted here: after N parallel writes, the
|
||||
// Store returns each written snapshot intact under LoadTools.
|
||||
func TestSharedCacheStoreConcurrentSaveTools(t *testing.T) {
|
||||
const (
|
||||
partition = "default/default"
|
||||
writers = 16
|
||||
)
|
||||
|
||||
store := cache.NewStore(t.TempDir())
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < writers; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", idx)
|
||||
if err := store.SaveTools(partition, key, cache.ToolsSnapshot{
|
||||
ServerKey: key,
|
||||
}); err != nil {
|
||||
t.Errorf("SaveTools(%s): %v", key, err)
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < writers; i++ {
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", i)
|
||||
snapshot, _, err := store.LoadTools(partition, key)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTools(%s): %v", key, err)
|
||||
}
|
||||
if snapshot.ServerKey != key {
|
||||
t.Errorf("LoadTools(%s) returned ServerKey %q", key, snapshot.ServerKey)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppendDynamicServerConcurrent exercises the dynamicMu mutex on the
|
||||
// write path by spraying distinct server descriptors in parallel. Afterwards
|
||||
// every injected product ID must be resolvable — a missing entry would
|
||||
// indicate a lost write through an un-synchronized map update.
|
||||
func TestAppendDynamicServerConcurrent(t *testing.T) {
|
||||
dynamicMu.Lock()
|
||||
prev := struct {
|
||||
endpoints map[string]string
|
||||
products map[string]bool
|
||||
aliases map[string]string
|
||||
toolEndpoints map[string]string
|
||||
}{dynamicEndpoints, dynamicProducts, dynamicAliases, dynamicToolEndpoints}
|
||||
dynamicEndpoints = nil
|
||||
dynamicProducts = nil
|
||||
dynamicAliases = nil
|
||||
dynamicToolEndpoints = nil
|
||||
dynamicMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
dynamicMu.Lock()
|
||||
dynamicEndpoints = prev.endpoints
|
||||
dynamicProducts = prev.products
|
||||
dynamicAliases = prev.aliases
|
||||
dynamicToolEndpoints = prev.toolEndpoints
|
||||
dynamicMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
id := fmt.Sprintf("plugin-id-%d", idx)
|
||||
endpoint := fmt.Sprintf("https://example.test/%d", idx)
|
||||
AppendDynamicServer(market.ServerDescriptor{
|
||||
Endpoint: endpoint,
|
||||
CLI: market.CLIOverlay{
|
||||
ID: id,
|
||||
Command: id,
|
||||
},
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
id := fmt.Sprintf("plugin-id-%d", i)
|
||||
if endpoint, ok := directRuntimeEndpoint(id, ""); !ok || endpoint == "" {
|
||||
t.Errorf("directRuntimeEndpoint(%q) = (%q, %v), want non-empty", id, endpoint, ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegisterStdioClientConcurrent verifies the stdioMu-protected registry
|
||||
// survives concurrent writers — every registered client must be looked up
|
||||
// afterwards. Uses nil client pointers since LookupStdioClient only compares
|
||||
// keys, not values.
|
||||
func TestRegisterStdioClientConcurrent(t *testing.T) {
|
||||
stdioMu.Lock()
|
||||
prev := stdioClients
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
stdioMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
stdioMu.Lock()
|
||||
stdioClients = prev
|
||||
stdioMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
RegisterStdioClient(fmt.Sprintf("plugin/%d", idx), nil)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
key := fmt.Sprintf("plugin/%d", i)
|
||||
if _, ok := LookupStdioClient(key); !ok {
|
||||
t.Errorf("LookupStdioClient(%q) missing after concurrent registration", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolvePluginColdTimeouts covers the three code paths of the env
|
||||
// parser: unset (defaults), valid duration (applied to all three slots),
|
||||
// and invalid duration (logged and ignored, defaults returned).
|
||||
func TestResolvePluginColdTimeouts(t *testing.T) {
|
||||
t.Run("defaults when env unset", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("httpNoAuth = %v, want 1s", got.httpNoAuth)
|
||||
}
|
||||
if got.httpAuth != 1500*time.Millisecond {
|
||||
t.Errorf("httpAuth = %v, want 1.5s", got.httpAuth)
|
||||
}
|
||||
if got.stdio != 2*time.Second {
|
||||
t.Errorf("stdio = %v, want 2s", got.stdio)
|
||||
}
|
||||
})
|
||||
t.Run("env override applies to all slots", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "3500ms")
|
||||
got := resolvePluginColdTimeouts()
|
||||
want := 3500 * time.Millisecond
|
||||
if got.httpNoAuth != want || got.httpAuth != want || got.stdio != want {
|
||||
t.Errorf("override not propagated: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("invalid env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "not-a-duration")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second || got.stdio != 2*time.Second {
|
||||
t.Errorf("invalid env should not override defaults: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("non-positive env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "0")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("zero duration should not override defaults, got %v", got.httpNoAuth)
|
||||
}
|
||||
})
|
||||
}
|
||||
+671
-5
@@ -15,20 +15,25 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/compat"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
@@ -36,8 +41,10 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pat"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline/handlers"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
@@ -51,9 +58,20 @@ type outputFileContextKey struct{}
|
||||
const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
|
||||
|
||||
// Execute runs the root command and returns the process exit code.
|
||||
func Execute() int {
|
||||
func Execute() (exitCode int) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
fmt.Fprintf(os.Stderr, "Error: internal panic: %v\n", r)
|
||||
exitCode = 5
|
||||
}
|
||||
}()
|
||||
|
||||
timing := NewTimingCollector()
|
||||
defer func() { timing.PrintIfEnabled() }()
|
||||
defer func() {
|
||||
StopAllStdioClients() // Ensure child processes are terminated on exit
|
||||
timing.PrintIfEnabled()
|
||||
timing.WriteReportIfEnabled(RawVersion(), SanitizeCommand(os.Args))
|
||||
}()
|
||||
|
||||
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||
defer cancel()
|
||||
@@ -123,8 +141,13 @@ func flagErrorWithSuggestions(cmd *cobra.Command, err error) error {
|
||||
}
|
||||
|
||||
func printExecutionError(root *cobra.Command, stdout, stderr io.Writer, err error) error {
|
||||
var raw apperrors.RawStderrError
|
||||
if stderrors.As(err, &raw) {
|
||||
_, writeErr := fmt.Fprintln(stderr, raw.RawStderr())
|
||||
return writeErr
|
||||
}
|
||||
if wantsJSONErrors(root) {
|
||||
return apperrors.PrintJSON(stdout, err)
|
||||
return apperrors.PrintJSON(stderr, err)
|
||||
}
|
||||
return apperrors.PrintHumanAt(stderr, err, resolveVerbosity(root))
|
||||
}
|
||||
@@ -242,9 +265,16 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
// Configure global slog level based on --debug / --verbose flags.
|
||||
configureLogLevel(flags)
|
||||
|
||||
return configureOutputSink(cmd)
|
||||
if err := configureOutputSink(cmd); err != nil {
|
||||
return err
|
||||
}
|
||||
if fn := edition.Get().AfterPersistentPreRun; fn != nil {
|
||||
return fn(cmd, args)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
PersistentPostRunE: func(cmd *cobra.Command, args []string) error {
|
||||
StopAllStdioClients()
|
||||
CloseFileLogger()
|
||||
return closeOutputSink(cmd)
|
||||
},
|
||||
@@ -262,18 +292,34 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
newAuthCommand(),
|
||||
newSkillCommand(),
|
||||
newCacheCommand(),
|
||||
newConfigCommand(),
|
||||
newDoctorCommand(),
|
||||
newCompletionCommand(root),
|
||||
newRecoveryCommand(rootCtx, loader, flags),
|
||||
newUpgradeCommand(),
|
||||
newVersionCommand(),
|
||||
newPluginCommand(),
|
||||
schemaCmd,
|
||||
genSkillsCmd,
|
||||
mcpCmd,
|
||||
}
|
||||
root.AddCommand(utilityCommands...)
|
||||
|
||||
root.AddCommand(newLegacyPublicCommands(rootCtx, runner)...)
|
||||
root.AddCommand(newLegacyHiddenCommands(runner)...)
|
||||
|
||||
// --- Plugin loading: runs AFTER legacy commands so that
|
||||
// AppendDynamicServer adds plugin endpoints on top of Market
|
||||
// endpoints (SetDynamicServers is called inside loadDynamicCommands).
|
||||
pluginCmds := loadPlugins(engine, runner)
|
||||
if len(pluginCmds) > 0 {
|
||||
addPluginCommandsSafe(root, pluginCmds)
|
||||
}
|
||||
|
||||
// PAT authorization commands (open-source core)
|
||||
patCaller := newToolCallerAdapter(runner, flags)
|
||||
pat.RegisterCommands(root, patCaller)
|
||||
|
||||
if fn := edition.Get().RegisterExtraCommands; fn != nil {
|
||||
caller := newToolCallerAdapter(runner, flags)
|
||||
fn(root, caller)
|
||||
@@ -617,7 +663,11 @@ func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
staticCommands := map[string]bool{
|
||||
"auth": true,
|
||||
"cache": true,
|
||||
"config": true,
|
||||
"doctor": true,
|
||||
"completion": true,
|
||||
"skill": true,
|
||||
"plugin": true,
|
||||
"version": true,
|
||||
"help": true,
|
||||
"recovery": true,
|
||||
@@ -639,6 +689,66 @@ func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
}
|
||||
}
|
||||
|
||||
// reservedCommands is the set of built-in command names that plugins must
|
||||
// not override. This protects core CLI functionality from being hijacked
|
||||
// by a malicious or misconfigured plugin.
|
||||
var reservedCommands = map[string]bool{
|
||||
"auth": true, "login": true, "logout": true,
|
||||
"plugin": true, "skill": true, "cache": true,
|
||||
"config": true, "doctor": true, "completion": true,
|
||||
"recovery": true, "upgrade": true, "version": true,
|
||||
"schema": true, "mcp": true, "help": true,
|
||||
}
|
||||
|
||||
// addPluginCommandsSafe registers plugin commands with conflict detection.
|
||||
//
|
||||
// Rules:
|
||||
// - Plugin vs reserved (auth/plugin/cache/...) → reject, warn
|
||||
// - Plugin vs plugin (same name) → reject later one, warn
|
||||
// - Plugin vs Market dynamic command → allow, plugin wins
|
||||
func addPluginCommandsSafe(root *cobra.Command, pluginCmds []*cobra.Command) {
|
||||
// Build index of existing commands before plugin registration.
|
||||
existing := make(map[string]bool)
|
||||
for _, cmd := range root.Commands() {
|
||||
existing[cmd.Name()] = true
|
||||
}
|
||||
|
||||
pluginSeen := make(map[string]bool)
|
||||
|
||||
for _, cmd := range pluginCmds {
|
||||
name := cmd.Name()
|
||||
|
||||
// Rule 1: never override reserved built-in commands.
|
||||
if reservedCommands[name] {
|
||||
slog.Warn("plugin: command name conflicts with built-in command, skipping",
|
||||
"command", name)
|
||||
continue
|
||||
}
|
||||
|
||||
// Rule 2: plugin vs plugin — first plugin wins.
|
||||
if pluginSeen[name] {
|
||||
slog.Warn("plugin: duplicate command from another plugin, skipping",
|
||||
"command", name)
|
||||
continue
|
||||
}
|
||||
pluginSeen[name] = true
|
||||
|
||||
// Rule 3: plugin vs Market — plugin wins, remove the old one.
|
||||
if existing[name] {
|
||||
for _, old := range root.Commands() {
|
||||
if old.Name() == name {
|
||||
root.RemoveCommand(old)
|
||||
slog.Debug("plugin: overriding Market command",
|
||||
"command", name)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
root.AddCommand(cmd)
|
||||
}
|
||||
}
|
||||
|
||||
// deduplicateCommands removes duplicate top-level commands, keeping the last
|
||||
// registered one. This ensures overlay commands take precedence over
|
||||
// open-source defaults when both register the same product name.
|
||||
@@ -662,6 +772,43 @@ func cacheStoreFromEnv() *cache.Store {
|
||||
return cache.NewStore(cacheDir)
|
||||
}
|
||||
|
||||
// pluginColdTimeouts holds the cold-path discovery budget for plugin MCP
|
||||
// servers. Timeouts only apply to the *first* discovery for a given
|
||||
// plugin/server; subsequent startups take the warm cache path and bypass
|
||||
// the network entirely.
|
||||
type pluginColdTimeouts struct {
|
||||
httpNoAuth time.Duration
|
||||
httpAuth time.Duration
|
||||
stdio time.Duration
|
||||
}
|
||||
|
||||
// resolvePluginColdTimeouts returns the cold-discovery budget for plugin MCP
|
||||
// servers, applying the DWS_PLUGIN_COLD_TIMEOUT override when set. Defaults
|
||||
// are tuned so healthy cross-region HTTP endpoints succeed on a cold start
|
||||
// and Python/Node-based stdio plugins have headroom for interpreter load,
|
||||
// while an unreachable host still surrenders in bounded time.
|
||||
func resolvePluginColdTimeouts() pluginColdTimeouts {
|
||||
t := pluginColdTimeouts{
|
||||
httpNoAuth: 1 * time.Second,
|
||||
httpAuth: 1500 * time.Millisecond,
|
||||
stdio: 2 * time.Second,
|
||||
}
|
||||
raw := strings.TrimSpace(os.Getenv(cli.PluginColdTimeoutEnv))
|
||||
if raw == "" {
|
||||
return t
|
||||
}
|
||||
d, err := time.ParseDuration(raw)
|
||||
if err != nil || d <= 0 {
|
||||
slog.Warn("plugin: ignoring invalid DWS_PLUGIN_COLD_TIMEOUT",
|
||||
"value", raw, "error", err)
|
||||
return t
|
||||
}
|
||||
t.httpNoAuth = d
|
||||
t.httpAuth = d
|
||||
t.stdio = d
|
||||
return t
|
||||
}
|
||||
|
||||
func configureOutputSink(cmd *cobra.Command) error {
|
||||
if local := cmd.LocalFlags().Lookup("output"); local != nil {
|
||||
return nil
|
||||
@@ -918,11 +1065,524 @@ func CloseFileLogger() {
|
||||
}
|
||||
}
|
||||
|
||||
// loadPlugins scans plugin directories, injects their MCP servers into
|
||||
// the dynamic server registry, and registers their pipeline hooks.
|
||||
// This runs before legacy command construction so that plugin servers
|
||||
// are available for EnvironmentLoader.Load().
|
||||
func loadPlugins(engine *pipeline.Engine, runner executor.Runner) []*cobra.Command {
|
||||
pluginLoader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// 0a. Inject plugin config values from settings.json as environment
|
||||
// variables so that expandPluginVars can resolve ${KEY} references
|
||||
// in plugin.json headers, endpoints, etc. User-set env vars take
|
||||
// precedence (InjectPluginConfigEnv skips already-set keys).
|
||||
pluginLoader.InjectPluginConfigEnv()
|
||||
|
||||
// Load TokenData once; reused for stdio injection below.
|
||||
tokenData, _ := authpkg.LoadTokenData(defaultConfigDir())
|
||||
var userCtx *plugin.UserContext
|
||||
if tokenData != nil {
|
||||
// Inject user context if either UserID or CorpID is present.
|
||||
if tokenData.UserID != "" || tokenData.CorpID != "" {
|
||||
userCtx = &plugin.UserContext{
|
||||
UserID: tokenData.UserID,
|
||||
CorpID: tokenData.CorpID,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Load user plugins (per settings.json)
|
||||
userPlugins := pluginLoader.LoadUser()
|
||||
|
||||
// 2. Load dev plugins (registered via `dws plugin dev`)
|
||||
devPlugins := pluginLoader.LoadDev()
|
||||
|
||||
allPlugins := append(userPlugins, devPlugins...)
|
||||
|
||||
// 3. Discover tools from streamable-http servers and build CLI commands.
|
||||
// Third-party servers with auth headers are discovered in parallel
|
||||
// to avoid sequential 10s timeouts when multiple remote servers exist.
|
||||
var pluginCmds []*cobra.Command
|
||||
tc := transport.NewClient(nil)
|
||||
|
||||
// Collect all server descriptors and register auth first (fast, no I/O).
|
||||
type pluginServer struct {
|
||||
plugin *plugin.Plugin
|
||||
srv market.ServerDescriptor
|
||||
}
|
||||
var httpServers []pluginServer
|
||||
|
||||
for _, p := range allPlugins {
|
||||
for _, srv := range p.ToServerDescriptors() {
|
||||
AppendDynamicServer(srv)
|
||||
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
}
|
||||
|
||||
if srv.HasCLIMeta {
|
||||
httpServers = append(httpServers, pluginServer{plugin: p, srv: srv})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Collect all stdio clients up front so HTTP + stdio discovery can run
|
||||
// concurrently — the slowest plugin (typically an unreachable HTTP
|
||||
// endpoint hitting its dial timeout) dominates the parallel wall-clock,
|
||||
// not the sum of every plugin's cold timeout.
|
||||
type stdioEntry struct {
|
||||
plugin *plugin.Plugin
|
||||
sc plugin.StdioServerClient
|
||||
}
|
||||
var stdioEntries []stdioEntry
|
||||
for _, p := range allPlugins {
|
||||
for _, sc := range p.StdioClients(userCtx) {
|
||||
// Use background context so the subprocess lives for the CLI
|
||||
// process lifetime (not killed by a short timeout).
|
||||
if err := sc.Client.Start(context.Background()); err != nil {
|
||||
slog.Warn("plugin: failed to start stdio server",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
continue
|
||||
}
|
||||
stdioEntries = append(stdioEntries, stdioEntry{plugin: p, sc: sc})
|
||||
}
|
||||
}
|
||||
|
||||
// Share one cache.Store across all discovery goroutines. Each goroutine
|
||||
// writes to a distinct serverKey path ("tools/<plugin>_<server>.json") with
|
||||
// atomic tmp+rename, so concurrent writes to different keys never collide
|
||||
// on the filesystem. Global in-process registries (AppendDynamicServer,
|
||||
// RegisterStdioClient) carry their own sync.Mutex; see direct_runtime.go
|
||||
// and stdio_registry.go.
|
||||
sharedStore := cacheStoreFromEnv()
|
||||
coldTimeouts := resolvePluginColdTimeouts()
|
||||
|
||||
// Fan out HTTP and stdio discovery in parallel. Each goroutine resolves
|
||||
// its cache hit locally (no network) or runs a bounded cold-path probe.
|
||||
// Wall-clock cost ≈ max(individual plugin latencies), not the sum.
|
||||
httpResults := make([][]*cobra.Command, len(httpServers))
|
||||
stdioResults := make([][]*cobra.Command, len(stdioEntries))
|
||||
var wg sync.WaitGroup
|
||||
for i, ps := range httpServers {
|
||||
wg.Add(1)
|
||||
go func(idx int, ps pluginServer) {
|
||||
defer wg.Done()
|
||||
httpResults[idx] = registerHTTPServer(ps.plugin, ps.srv, tc, runner, sharedStore, coldTimeouts)
|
||||
}(i, ps)
|
||||
}
|
||||
for i, e := range stdioEntries {
|
||||
wg.Add(1)
|
||||
go func(idx int, e stdioEntry) {
|
||||
defer wg.Done()
|
||||
stdioResults[idx] = registerStdioServer(e.plugin, e.sc, runner, sharedStore, coldTimeouts)
|
||||
}(i, e)
|
||||
}
|
||||
wg.Wait()
|
||||
for _, cmds := range httpResults {
|
||||
pluginCmds = append(pluginCmds, cmds...)
|
||||
}
|
||||
for _, cmds := range stdioResults {
|
||||
pluginCmds = append(pluginCmds, cmds...)
|
||||
}
|
||||
|
||||
// 5. Register plugin hooks into pipeline engine
|
||||
if engine != nil {
|
||||
for _, p := range allPlugins {
|
||||
hooksCfg, err := p.LoadHooks()
|
||||
if err != nil {
|
||||
slog.Warn("plugin: failed to load hooks",
|
||||
"plugin", p.Manifest.Name, "error", err)
|
||||
continue
|
||||
}
|
||||
if hooksCfg == nil {
|
||||
continue
|
||||
}
|
||||
for _, entry := range hooksCfg.Hooks {
|
||||
engine.Register(plugin.NewHookAdapter(p.Manifest.Name, entry))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Sync plugin skills to agent directories
|
||||
plugin.SyncSkills(allPlugins)
|
||||
|
||||
if len(allPlugins) > 0 {
|
||||
slog.Debug("plugins loaded",
|
||||
"user", len(userPlugins),
|
||||
"dev", len(devPlugins),
|
||||
)
|
||||
}
|
||||
|
||||
return pluginCmds
|
||||
}
|
||||
|
||||
// pluginCacheKey derives the cache key used to persist a plugin MCP server's
|
||||
// tool list. Prefixed with "plugin:" so entries are namespaced apart from the
|
||||
// Market-derived cache, and visible distinctly via `dws cache status`.
|
||||
func pluginCacheKey(pluginName, serverKey string) string {
|
||||
return "plugin:" + pluginName + ":" + serverKey
|
||||
}
|
||||
|
||||
// registerHTTPServer discovers tools from a streamable-http MCP server and
|
||||
// builds CLI commands. Used for plugin-owned HTTP servers that provide CLI metadata.
|
||||
//
|
||||
// Startup-latency strategy (issue #119):
|
||||
// - Warm cache: build commands from the persisted tools snapshot
|
||||
// synchronously — no network I/O. `dws --help` returns in ms even when
|
||||
// the plugin endpoint is unreachable.
|
||||
// - Cold cache: synchronous discovery (Initialize + ListTools) with a tight
|
||||
// timeout. The outcome — success or failure — is persisted so the next
|
||||
// invocation hits the warm path. Refresh on demand via `dws cache clean`
|
||||
// / `dws cache refresh`; the cache TTL (7d) otherwise expires naturally.
|
||||
//
|
||||
// When the server descriptor carries AuthHeaders (from plugin.json "headers"),
|
||||
// a dedicated transport.Client is created with the plugin's Bearer token and
|
||||
// trusted domains so that third-party MCP servers requiring independent
|
||||
// authentication (e.g. Alibaba Cloud Bailian) can be discovered at startup.
|
||||
func registerHTTPServer(p *plugin.Plugin, srv market.ServerDescriptor, tc *transport.Client, runner executor.Runner, store *cache.Store, timeouts pluginColdTimeouts) []*cobra.Command {
|
||||
partition := config.DefaultPartition
|
||||
cacheKey := pluginCacheKey(p.Manifest.Name, srv.Key)
|
||||
|
||||
if snapshot, freshness, err := store.LoadTools(partition, cacheKey); err == nil {
|
||||
slog.Debug("plugin: http server served from cache",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key,
|
||||
"tools", len(snapshot.Tools), "freshness", string(freshness))
|
||||
return buildHTTPCommandsFromTools(srv, snapshot.Tools, runner)
|
||||
}
|
||||
|
||||
// Cold cache: synchronous discovery. Persist the outcome even on failure
|
||||
// (empty tools == negative cache) so the next invocation takes the fast
|
||||
// path regardless of endpoint health.
|
||||
tools := discoverHTTPTools(p, srv, tc, timeouts)
|
||||
_ = store.SaveTools(partition, cacheKey, cache.ToolsSnapshot{
|
||||
ServerKey: cacheKey,
|
||||
Tools: tools,
|
||||
})
|
||||
return buildHTTPCommandsFromTools(srv, tools, runner)
|
||||
}
|
||||
|
||||
// discoverHTTPTools performs the blocking Initialize + ListTools handshake
|
||||
// for an HTTP MCP server and returns the discovered tools. Returns nil on
|
||||
// any transport/protocol error; errors are logged at Debug level.
|
||||
func discoverHTTPTools(p *plugin.Plugin, srv market.ServerDescriptor, tc *transport.Client, timeouts pluginColdTimeouts) []transport.ToolDescriptor {
|
||||
// Cold-path budget. An unreachable endpoint will burn the full window
|
||||
// via the TCP dial timeout; a healthy localhost/third-party endpoint
|
||||
// typically responds in <200 ms. Third-party servers with auth get a
|
||||
// slightly larger window to accommodate TLS + auth RTT. Operators with
|
||||
// cross-region endpoints can relax the window via DWS_PLUGIN_COLD_TIMEOUT.
|
||||
// The outcome is persisted as a negative cache so subsequent startups
|
||||
// (80 ms warm) are unaffected. See issue #119.
|
||||
timeout := timeouts.httpNoAuth
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
timeout = timeouts.httpAuth
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
|
||||
discoveryClient := tc
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
discoveryClient = buildPluginAuthClient(tc, srv)
|
||||
}
|
||||
|
||||
if _, err := discoveryClient.Initialize(ctx, srv.Endpoint); err != nil {
|
||||
slog.Debug("plugin: http server offline, skipping tool discovery",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
toolsResult, err := discoveryClient.ListTools(ctx, srv.Endpoint)
|
||||
if err != nil {
|
||||
slog.Debug("plugin: http ListTools failed",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
return toolsResult.Tools
|
||||
}
|
||||
|
||||
// buildHTTPCommandsFromTools converts a tool list into Cobra commands via
|
||||
// the BuildDynamicCommands path. Returns nil for an empty tool list.
|
||||
func buildHTTPCommandsFromTools(srv market.ServerDescriptor, tools []transport.ToolDescriptor, runner executor.Runner) []*cobra.Command {
|
||||
if len(tools) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
detailsByID := make(map[string][]market.DetailTool)
|
||||
var detailTools []market.DetailTool
|
||||
for _, tool := range tools {
|
||||
schemaJSON := ""
|
||||
if tool.InputSchema != nil {
|
||||
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
|
||||
schemaJSON = string(data)
|
||||
}
|
||||
}
|
||||
detailTools = append(detailTools, market.DetailTool{
|
||||
ToolName: tool.Name,
|
||||
ToolTitle: tool.Title,
|
||||
ToolDesc: tool.Description,
|
||||
IsSensitive: tool.Sensitive,
|
||||
ToolRequest: schemaJSON,
|
||||
})
|
||||
}
|
||||
detailsByID[strings.TrimSpace(srv.CLI.ID)] = detailTools
|
||||
|
||||
// If the server has no ToolOverrides (e.g. third-party MCP servers that
|
||||
// only declare cli.id and cli.command), auto-generate one override per
|
||||
// discovered tool so BuildDynamicCommands can create leaf commands.
|
||||
if len(srv.CLI.ToolOverrides) == 0 {
|
||||
srv.CLI.ToolOverrides = make(map[string]market.CLIToolOverride, len(tools))
|
||||
for _, tool := range tools {
|
||||
srv.CLI.ToolOverrides[tool.Name] = market.CLIToolOverride{
|
||||
CLIName: deriveToolCLIName(tool.Name),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return compat.BuildDynamicCommands(
|
||||
[]market.ServerDescriptor{srv}, runner, detailsByID)
|
||||
}
|
||||
|
||||
// deriveToolCLIName converts an MCP tool name (e.g. "web_search" or
|
||||
// "maps.search_poi") into a kebab-case CLI command name ("search" or
|
||||
// "search-poi"). It strips common prefixes and replaces underscores/dots
|
||||
// with hyphens.
|
||||
func deriveToolCLIName(toolName string) string {
|
||||
// Use the last segment after "." as the base name.
|
||||
if idx := strings.LastIndex(toolName, "."); idx >= 0 {
|
||||
toolName = toolName[idx+1:]
|
||||
}
|
||||
// Replace underscores with hyphens for kebab-case.
|
||||
return strings.ReplaceAll(toolName, "_", "-")
|
||||
}
|
||||
|
||||
// buildPluginAuthClient creates a transport.Client copy with the plugin's
|
||||
// Bearer token and trusted domains injected. This allows third-party MCP
|
||||
// servers that require independent authentication to be discovered at startup.
|
||||
func buildPluginAuthClient(base *transport.Client, srv market.ServerDescriptor) *transport.Client {
|
||||
authToken := ""
|
||||
extraHeaders := make(map[string]string)
|
||||
for key, value := range srv.AuthHeaders {
|
||||
if strings.EqualFold(key, "Authorization") {
|
||||
authToken = strings.TrimPrefix(value, "Bearer ")
|
||||
authToken = strings.TrimSpace(authToken)
|
||||
} else {
|
||||
extraHeaders[key] = value
|
||||
}
|
||||
}
|
||||
if authToken == "" {
|
||||
return base
|
||||
}
|
||||
client := base.WithAuth(authToken, extraHeaders)
|
||||
// Trust the endpoint's hostname so the token is actually sent.
|
||||
if parsed, err := url.Parse(srv.Endpoint); err == nil {
|
||||
host := parsed.Hostname()
|
||||
client.TrustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
// registerPluginAuthFromHeaders extracts authentication credentials from
|
||||
// a server descriptor's AuthHeaders and registers them in the global
|
||||
// PluginAuth registry. The runner uses this registry at execution time
|
||||
// to inject the correct Bearer token for third-party MCP servers.
|
||||
func registerPluginAuthFromHeaders(srv market.ServerDescriptor) {
|
||||
authToken := ""
|
||||
extraHeaders := make(map[string]string)
|
||||
for key, value := range srv.AuthHeaders {
|
||||
if strings.EqualFold(key, "Authorization") {
|
||||
authToken = strings.TrimPrefix(value, "Bearer ")
|
||||
authToken = strings.TrimSpace(authToken)
|
||||
} else {
|
||||
extraHeaders[key] = value
|
||||
}
|
||||
}
|
||||
if authToken == "" {
|
||||
return
|
||||
}
|
||||
var trustedDomains []string
|
||||
if parsed, err := url.Parse(srv.Endpoint); err == nil {
|
||||
host := parsed.Hostname()
|
||||
trustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
productID := strings.TrimSpace(srv.CLI.ID)
|
||||
if productID == "" {
|
||||
productID = srv.Key
|
||||
}
|
||||
RegisterPluginAuth(productID, &PluginAuth{
|
||||
Token: authToken,
|
||||
ExtraHeaders: extraHeaders,
|
||||
TrustedDomains: trustedDomains,
|
||||
})
|
||||
}
|
||||
|
||||
// registerStdioServer initializes a stdio MCP server, discovers its tools
|
||||
// via ListTools, builds CLI commands, and registers the StdioClient for
|
||||
// runtime dispatch. Returns generated cobra commands.
|
||||
//
|
||||
// Warm-cache fast path (issue #119): when a tools snapshot is already cached
|
||||
// for this plugin/server, skip the Initialize + ListTools RPC round-trip and
|
||||
// rebuild commands directly from the snapshot. Cold cache falls back to
|
||||
// synchronous discovery with a 4s cap and persists the outcome.
|
||||
func registerStdioServer(p *plugin.Plugin, sc plugin.StdioServerClient, runner executor.Runner, store *cache.Store, timeouts pluginColdTimeouts) []*cobra.Command {
|
||||
partition := config.DefaultPartition
|
||||
cacheKey := pluginCacheKey(p.Manifest.Name, sc.Key)
|
||||
|
||||
if snapshot, freshness, err := store.LoadTools(partition, cacheKey); err == nil {
|
||||
slog.Debug("plugin: stdio server served from cache",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key,
|
||||
"tools", len(snapshot.Tools), "freshness", string(freshness))
|
||||
return buildStdioCommands(p, sc, snapshot.Tools, runner)
|
||||
}
|
||||
|
||||
tools := discoverStdioTools(p, sc, timeouts)
|
||||
_ = store.SaveTools(partition, cacheKey, cache.ToolsSnapshot{
|
||||
ServerKey: cacheKey,
|
||||
Tools: tools,
|
||||
})
|
||||
return buildStdioCommands(p, sc, tools, runner)
|
||||
}
|
||||
|
||||
// discoverStdioTools performs the blocking Initialize + ListTools handshake
|
||||
// on a stdio MCP subprocess. Returns nil on any error (logged at Warn level).
|
||||
// The default 2s budget comfortably accommodates Python/Node runtimes whose
|
||||
// interpreter + dependency load dominates the first response. Operators with
|
||||
// heavier startup chains can relax further via DWS_PLUGIN_COLD_TIMEOUT.
|
||||
func discoverStdioTools(p *plugin.Plugin, sc plugin.StdioServerClient, timeouts pluginColdTimeouts) []transport.ToolDescriptor {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeouts.stdio)
|
||||
defer cancel()
|
||||
|
||||
if _, err := sc.Client.Initialize(ctx); err != nil {
|
||||
slog.Warn("plugin: stdio initialize failed",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
toolsResult, err := sc.Client.ListTools(ctx)
|
||||
if err != nil {
|
||||
slog.Warn("plugin: stdio ListTools failed",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
return toolsResult.Tools
|
||||
}
|
||||
|
||||
// buildStdioCommands constructs Cobra commands from a tool list and
|
||||
// registers the runtime dispatch state (StdioClient + dynamic server).
|
||||
// Returns nil for an empty tool list.
|
||||
func buildStdioCommands(p *plugin.Plugin, sc plugin.StdioServerClient, tools []transport.ToolDescriptor, runner executor.Runner) []*cobra.Command {
|
||||
if len(tools) == 0 {
|
||||
slog.Debug("plugin: stdio server has no tools",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build CLIOverlay: use manifest CLI metadata if present, else auto-generate.
|
||||
serverID := sc.Key
|
||||
overlay := market.CLIOverlay{
|
||||
ID: serverID,
|
||||
Command: serverID,
|
||||
}
|
||||
if srv, ok := p.Manifest.MCPServers[sc.Key]; ok && len(srv.CLI) > 0 {
|
||||
cliData := srv.CLI
|
||||
// If cli is a JSON string, treat it as a relative file path to an overlay file.
|
||||
if len(cliData) > 0 && cliData[0] == '"' {
|
||||
var cliPath string
|
||||
if err := json.Unmarshal(cliData, &cliPath); err == nil && cliPath != "" {
|
||||
absPath := filepath.Join(p.Root, cliPath)
|
||||
if fileData, readErr := os.ReadFile(absPath); readErr == nil {
|
||||
cliData = fileData
|
||||
} else {
|
||||
slog.Warn("plugin: failed to read CLI overlay file",
|
||||
"plugin", p.Manifest.Name, "path", absPath, "error", readErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := json.Unmarshal(cliData, &overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLI overlay for stdio server",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
}
|
||||
if overlay.ID == "" {
|
||||
overlay.ID = serverID
|
||||
}
|
||||
if overlay.Command == "" {
|
||||
overlay.Command = serverID
|
||||
}
|
||||
}
|
||||
|
||||
// Auto-generate ToolOverrides from discovered tools when not provided.
|
||||
if len(overlay.ToolOverrides) == 0 {
|
||||
overlay.ToolOverrides = make(map[string]market.CLIToolOverride)
|
||||
if len(overlay.Prefixes) == 0 {
|
||||
overlay.Prefixes = []string{serverID}
|
||||
}
|
||||
for _, tool := range tools {
|
||||
overlay.ToolOverrides[tool.Name] = market.CLIToolOverride{
|
||||
IsSensitive: tool.Sensitive,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Construct virtual endpoint and server descriptor.
|
||||
endpoint := StdioEndpoint(p.Manifest.Name, sc.Key)
|
||||
|
||||
descriptor := market.ServerDescriptor{
|
||||
Key: sc.Key,
|
||||
DisplayName: p.Manifest.Name + "/" + sc.Key,
|
||||
Description: p.Manifest.Description,
|
||||
Endpoint: endpoint,
|
||||
Source: "plugin",
|
||||
CLI: overlay,
|
||||
HasCLIMeta: true,
|
||||
}
|
||||
|
||||
AppendDynamicServer(descriptor)
|
||||
// Register with pluginName/serverKey format for cleanup by plugin name
|
||||
RegisterStdioClient(p.Manifest.Name+"/"+serverID, sc.Client)
|
||||
|
||||
// Convert tool descriptors to DetailTool entries for flag generation.
|
||||
detailsByID := make(map[string][]market.DetailTool)
|
||||
var detailTools []market.DetailTool
|
||||
for _, tool := range tools {
|
||||
schemaJSON := ""
|
||||
if tool.InputSchema != nil {
|
||||
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
|
||||
schemaJSON = string(data)
|
||||
}
|
||||
}
|
||||
detailTools = append(detailTools, market.DetailTool{
|
||||
ToolName: tool.Name,
|
||||
ToolTitle: tool.Title,
|
||||
ToolDesc: tool.Description,
|
||||
IsSensitive: tool.Sensitive,
|
||||
ToolRequest: schemaJSON,
|
||||
})
|
||||
}
|
||||
detailsByID[serverID] = detailTools
|
||||
|
||||
cmds := compat.BuildDynamicCommands(
|
||||
[]market.ServerDescriptor{descriptor}, runner, detailsByID)
|
||||
|
||||
slog.Debug("plugin: stdio server registered",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key,
|
||||
"tools", len(tools), "commands", len(cmds))
|
||||
|
||||
return cmds
|
||||
}
|
||||
|
||||
// newPipelineEngine creates and configures the pipeline engine with
|
||||
// the standard set of handlers for model input correction.
|
||||
// handlers for all five pipeline phases. The phases execute in order:
|
||||
// Register → PreParse → PostParse → PreRequest → PostResponse.
|
||||
//
|
||||
// Phases are invoked at their respective integration points:
|
||||
// - Register: during command tree construction (newMCPCommand)
|
||||
// - PreParse: before Cobra parses raw argv (RunPreParse)
|
||||
// - PostParse: after Cobra parsing, before validation (canonical RunE)
|
||||
// - PreRequest: after validation, before JSON-RPC dispatch (canonical RunE)
|
||||
// - PostResponse: after transport returns, before stdout (canonical RunE)
|
||||
func newPipelineEngine() *pipeline.Engine {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
// Register handler runs during command tree building.
|
||||
handlers.RegisterHandler{},
|
||||
|
||||
// PreParse handlers run in order: alias → sticky → paramname.
|
||||
// Alias normalises case first (--userId → --user-id), then
|
||||
// sticky splits glued values (--limit100 → --limit 100), then
|
||||
@@ -933,6 +1593,12 @@ func newPipelineEngine() *pipeline.Engine {
|
||||
|
||||
// PostParse handlers normalise structured values.
|
||||
handlers.ParamValueHandler{},
|
||||
|
||||
// PreRequest handler inspects the validated payload before dispatch.
|
||||
handlers.PreRequestHandler{},
|
||||
|
||||
// PostResponse handler processes the response before output.
|
||||
handlers.PostResponseHandler{},
|
||||
)
|
||||
return engine
|
||||
}
|
||||
|
||||
@@ -29,6 +29,14 @@ import (
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
// patLikeError simulates an edition-specific PAT error that implements both
|
||||
// ExitCoder (exit code 4) and RawStderrError (raw JSON to stderr).
|
||||
type patLikeError struct{ raw string }
|
||||
|
||||
func (e *patLikeError) Error() string { return e.raw }
|
||||
func (e *patLikeError) ExitCode() int { return 4 }
|
||||
func (e *patLikeError) RawStderr() string { return e.raw }
|
||||
|
||||
func TestPrintExecutionErrorDefaultsToJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -43,11 +51,11 @@ func TestPrintExecutionErrorDefaultsToJSON(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -65,11 +73,11 @@ func TestPrintExecutionErrorUsesJSONWhenFormatIsJSON(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -97,11 +105,11 @@ func TestPrintExecutionErrorUsesJSONWhenCommandSetsJSONFlag(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -255,6 +263,11 @@ func TestRootHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
if !strings.Contains(out.String(), "Discovered MCP Services:") {
|
||||
t.Fatalf("root help output missing MCP summary:\n%s", out.String())
|
||||
}
|
||||
for _, want := range []string{"Utility Commands:", "skill", "auth", "version"} {
|
||||
if !strings.Contains(out.String(), want) {
|
||||
t.Fatalf("root help output missing %q:\n%s", want, out.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRootShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
@@ -342,3 +355,86 @@ func TestNestedShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
t.Fatalf("nested short help output missing command title:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_writes_raw_JSON_to_stderr(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_LOW_RISK_NO_PERMISSION","data":{}}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty for RawStderrError", stdout.String())
|
||||
}
|
||||
got := strings.TrimSpace(stderr.String())
|
||||
if got != rawJSON {
|
||||
t.Fatalf("stderr = %q, want raw JSON %q", got, rawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_exit_code_is_4(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := &patLikeError{raw: `{"code":"PAT_MEDIUM_RISK_NO_PERMISSION"}`}
|
||||
exitCode := apperrors.ExitCode(err)
|
||||
if exitCode != 4 {
|
||||
t.Fatalf("apperrors.ExitCode(patLikeError) = %d, want 4", exitCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_takes_precedence_over_JSON_mode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_HIGH_RISK_NO_PERMISSION"}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
_ = root.PersistentFlags().Set("format", "json")
|
||||
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty — RawStderrError should bypass JSON mode", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stderr.String(), "PAT_HIGH_RISK_NO_PERMISSION") {
|
||||
t.Fatalf("stderr = %q, want raw PAT JSON", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
// simulateExecuteWithPanic mirrors the recovery pattern in Execute():
|
||||
// named return + defer recover → exitCode = 5 on panic.
|
||||
func simulateExecuteWithPanic(doPanic bool) (exitCode int) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
exitCode = 5
|
||||
}
|
||||
}()
|
||||
if doPanic {
|
||||
panic("test panic")
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func TestExecute_panic_recovery_returns_exit_5(t *testing.T) {
|
||||
t.Parallel()
|
||||
code := simulateExecuteWithPanic(true)
|
||||
if code != 5 {
|
||||
t.Fatalf("panic recovery exitCode = %d, want 5", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecute_no_panic_returns_0(t *testing.T) {
|
||||
t.Parallel()
|
||||
code := simulateExecuteWithPanic(false)
|
||||
if code != 0 {
|
||||
t.Fatalf("no-panic exitCode = %d, want 0", code)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
@@ -14,6 +15,26 @@ func configureRootHelp(root *cobra.Command) {
|
||||
return
|
||||
}
|
||||
|
||||
// Replace the cobra-default English help command with a localized one so
|
||||
// that both its listing short (shown in `dws --help`) and its own
|
||||
// `dws help --help` long text follow the active locale.
|
||||
root.SetHelpCommand(&cobra.Command{
|
||||
Use: "help [command]",
|
||||
Short: i18n.T("查看任意命令的帮助信息"),
|
||||
Long: i18n.T("显示任意命令的帮助文案。\n" +
|
||||
"用法:dws help [命令路径] 查看完整说明。"),
|
||||
DisableAutoGenTag: true,
|
||||
Run: func(c *cobra.Command, args []string) {
|
||||
target, _, err := c.Root().Find(args)
|
||||
if target == nil || err != nil {
|
||||
c.Root().HelpFunc()(c.Root(), args)
|
||||
return
|
||||
}
|
||||
target.InitDefaultHelpFlag()
|
||||
_ = target.Help()
|
||||
},
|
||||
})
|
||||
|
||||
defaultHelpFunc := root.HelpFunc()
|
||||
root.SetHelpFunc(func(cmd *cobra.Command, args []string) {
|
||||
if cmd != root {
|
||||
@@ -26,6 +47,7 @@ func configureRootHelp(root *cobra.Command) {
|
||||
|
||||
func renderRootHelp(root *cobra.Command) {
|
||||
services := visibleMCPRootCommands(root)
|
||||
utilities := visibleUtilityRootCommands(root)
|
||||
w := root.OutOrStdout()
|
||||
|
||||
if len(services) == 0 {
|
||||
@@ -45,8 +67,21 @@ func renderRootHelp(root *cobra.Command) {
|
||||
|
||||
_, _ = fmt.Fprintln(w, "Usage:")
|
||||
_, _ = fmt.Fprintln(w, " dws <service> [command] [flags]")
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, " dws <command> [flags]")
|
||||
}
|
||||
_, _ = fmt.Fprintln(w)
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service.`)
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, "Utility Commands:")
|
||||
_, _ = fmt.Fprintln(w)
|
||||
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
|
||||
for _, utility := range utilities {
|
||||
_, _ = fmt.Fprintf(tw, " %s\t%s\n", utility.Name(), strings.TrimSpace(utility.Short))
|
||||
}
|
||||
_ = tw.Flush()
|
||||
_, _ = fmt.Fprintln(w)
|
||||
}
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service or "dws <command> --help" for utility commands.`)
|
||||
}
|
||||
|
||||
func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
@@ -80,3 +115,29 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
func visibleUtilityRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
if root == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
productCommands := DirectRuntimeProductIDs()
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
productCommands = make(map[string]bool, len(fn()))
|
||||
for _, product := range fn() {
|
||||
productCommands[product] = true
|
||||
}
|
||||
}
|
||||
|
||||
commands := make([]*cobra.Command, 0)
|
||||
for _, cmd := range root.Commands() {
|
||||
if cmd == nil || cmd.Hidden {
|
||||
continue
|
||||
}
|
||||
if productCommands[cmd.Name()] {
|
||||
continue
|
||||
}
|
||||
commands = append(commands, cmd)
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
+207
-21
@@ -19,7 +19,6 @@ import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
@@ -34,9 +33,51 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/safety"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "启用 MCP 响应内容安全扫描",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN_ENFORCE",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "内容安全扫描发现问题时阻断响应",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN_REPORT",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "在 JSON 输出中包含安全扫描报告",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_AGENT",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-agent 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_TRACE_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-trace-id 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_SESSION_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-session-id 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_MESSAGE_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-message-id 头",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
runtimeContentScanEnv = "DWS_RUNTIME_CONTENT_SCAN"
|
||||
runtimeContentScanEnforceEnv = "DWS_RUNTIME_CONTENT_SCAN_ENFORCE"
|
||||
@@ -47,9 +88,21 @@ const (
|
||||
envDingtalkTraceID = "DINGTALK_TRACE_ID"
|
||||
envDingtalkSessionID = "DINGTALK_SESSION_ID"
|
||||
envDingtalkMessageID = "DINGTALK_MESSAGE_ID"
|
||||
|
||||
// Environment variables for third-party channel integration
|
||||
envDWSChannel = "DWS_CHANNEL"
|
||||
)
|
||||
|
||||
func newCommandRunnerWithFlags(loader cli.CatalogLoader, flags *GlobalFlags) executor.Runner {
|
||||
// Ensure DWS_CLIENT_ID env is populated from persisted config before
|
||||
// resolveIdentityHeaders reads it. This covers fresh-process cold starts
|
||||
// where no env var has been inherited from a parent process.
|
||||
if os.Getenv("DWS_CLIENT_ID") == "" {
|
||||
if cid := authpkg.ClientID(); cid != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", cid)
|
||||
}
|
||||
}
|
||||
|
||||
var httpClient *http.Client
|
||||
if flags != nil && flags.Timeout > 0 {
|
||||
httpClient = &http.Client{Timeout: time.Duration(flags.Timeout) * time.Second}
|
||||
@@ -107,7 +160,10 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
catalog, err := r.loader.Load(ctx)
|
||||
RecordTiming(ctx, "catalog_load", time.Since(catalogStart))
|
||||
if err != nil {
|
||||
return executor.Result{}, err
|
||||
var degraded *cli.CatalogDegraded
|
||||
if !errors.As(err, °raded) {
|
||||
return executor.Result{}, err
|
||||
}
|
||||
}
|
||||
|
||||
product, ok := catalog.FindProduct(invocation.CanonicalProduct)
|
||||
@@ -129,6 +185,11 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (result executor.Result, retErr error) {
|
||||
// Route stdio:// endpoints to the local StdioClient — no HTTP, no auth.
|
||||
if IsStdioEndpoint(endpoint) {
|
||||
return r.executeStdioInvocation(ctx, invocation)
|
||||
}
|
||||
|
||||
invokeStart := time.Now()
|
||||
execID := generateExecutionID()
|
||||
r.transport.ExecutionId = execID
|
||||
@@ -158,7 +219,17 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
retErr == nil, time.Since(invokeStart), errCat, errReason)
|
||||
}()
|
||||
|
||||
authToken := r.resolveAuthToken(ctx)
|
||||
// Check if this product has plugin-level auth credentials registered.
|
||||
// If so, use the plugin's token instead of the default DingTalk OAuth token.
|
||||
// This allows third-party MCP servers (e.g. Bailian) to use their own API keys.
|
||||
pluginAuth, hasPluginAuth := LookupPluginAuth(invocation.CanonicalProduct)
|
||||
|
||||
authToken := ""
|
||||
if hasPluginAuth {
|
||||
authToken = pluginAuth.Token
|
||||
} else {
|
||||
authToken = r.resolveAuthToken(ctx)
|
||||
}
|
||||
|
||||
var timeoutSec int
|
||||
if r.globalFlags != nil {
|
||||
@@ -207,7 +278,15 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
)
|
||||
}
|
||||
|
||||
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
var tc *transport.Client
|
||||
if hasPluginAuth {
|
||||
// Use plugin-level auth: inject the plugin's token and trust its domains.
|
||||
tc = r.transport.WithAuth(authToken, pluginAuth.ExtraHeaders)
|
||||
tc.TrustedDomains = pluginAuth.TrustedDomains
|
||||
} else {
|
||||
// Default path: use DingTalk OAuth token with identity headers.
|
||||
tc = r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
|
||||
@@ -222,16 +301,56 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
if err != nil {
|
||||
if isAuthError(err) {
|
||||
if fn := edition.Get().OnAuthError; fn != nil {
|
||||
_ = fn(defaultConfigDir(), err)
|
||||
if overrideErr := fn(defaultConfigDir(), err); overrideErr != nil {
|
||||
captureRuntimeFailure(invocation, err, overrideErr)
|
||||
return executor.Result{}, overrideErr
|
||||
}
|
||||
}
|
||||
}
|
||||
// PAT scope error: offer human-readable output and retry after authorization
|
||||
if isPatScopeError(err) {
|
||||
scopeErr := extractPatScopeError(err)
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return executor.Result{}, err
|
||||
}
|
||||
|
||||
// ---- Edition hook gets first dibs (preserves overlay PATError passthrough) ----
|
||||
if fn := edition.Get().ClassifyToolResult; fn != nil {
|
||||
if editionErr := fn(callResult.Content); editionErr != nil {
|
||||
if patCheck := apperrors.AsPatAuthCheckError(editionErr); patCheck != nil {
|
||||
if IsPatRetrying(ctx) {
|
||||
return executor.Result{}, patCheck // already retried once, don't loop
|
||||
}
|
||||
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
return executor.Result{}, editionErr
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Structured PAT auth check (open-source fallback) ----
|
||||
if patCheck := apperrors.ClassifyPatAuthCheck(callResult.Content); patCheck != nil {
|
||||
if IsPatRetrying(ctx) {
|
||||
return executor.Result{}, patCheck // already retried once, don't loop
|
||||
}
|
||||
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
|
||||
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
|
||||
|
||||
// ClassifyToolResult hook: let the overlay intercept known error
|
||||
// patterns (PAT permission, gateway-auth) before generic handling.
|
||||
if classify := edition.Get().ClassifyToolResult; classify != nil {
|
||||
if hookErr := classify(callResult.Content); hookErr != nil {
|
||||
captureRuntimeFailure(invocation, hookErr, hookErr)
|
||||
return executor.Result{}, hookErr
|
||||
}
|
||||
}
|
||||
|
||||
mcpErr := apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
@@ -240,6 +359,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
apperrors.WithHint("MCP tool returned a business error; check tool parameters and refer to skill documentation."),
|
||||
apperrors.WithServerDiag(diag),
|
||||
)
|
||||
// PAT scope error in business response: offer human-readable output and retry
|
||||
if isPatScopeError(mcpErr) {
|
||||
scopeErr := extractPatScopeError(mcpErr)
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return executor.Result{}, mcpErr
|
||||
}
|
||||
@@ -272,12 +397,78 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
return executor.Result{Invocation: invocation, Response: response}, nil
|
||||
}
|
||||
|
||||
// executeStdioInvocation dispatches a tool call through a local StdioClient
|
||||
// subprocess instead of the HTTP transport. This is used for plugin stdio
|
||||
// servers whose endpoints use the stdio:// scheme.
|
||||
func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
if invocation.DryRun {
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
Response: map[string]any{
|
||||
"dry_run": true,
|
||||
"transport": "stdio",
|
||||
"request": executor.ToolCallRequest(invocation.Tool, invocation.Params),
|
||||
"note": "execution skipped by --dry-run",
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
client, ok := LookupStdioClient(invocation.CanonicalProduct)
|
||||
if !ok {
|
||||
return executor.Result{}, apperrors.NewInternal(
|
||||
fmt.Sprintf("stdio client not found for %q", invocation.CanonicalProduct))
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
|
||||
var cancel context.CancelFunc
|
||||
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
|
||||
defer cancel()
|
||||
}
|
||||
|
||||
callResult, err := client.CallTool(callCtx, invocation.Tool, invocation.Params)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
fmt.Sprintf("stdio call failed: %v", err),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("stdio_error"),
|
||||
)
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("mcp_tool_error"),
|
||||
apperrors.WithServerKey(invocation.CanonicalProduct),
|
||||
)
|
||||
}
|
||||
|
||||
invocation.Implemented = true
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
Response: map[string]any{
|
||||
"transport": "stdio",
|
||||
"content": callResult.Content,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
|
||||
explicitToken := ""
|
||||
if r != nil && r.globalFlags != nil {
|
||||
explicitToken = r.globalFlags.Token
|
||||
}
|
||||
return resolveRuntimeAuthToken(ctx, explicitToken)
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
}
|
||||
if tp := edition.Get().TokenProvider; tp != nil {
|
||||
token, _ := tp(ctx, func() (string, error) {
|
||||
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
|
||||
})
|
||||
return token
|
||||
}
|
||||
return getCachedRuntimeToken(ctx)
|
||||
}
|
||||
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
|
||||
@@ -302,24 +493,13 @@ func getCachedRuntimeToken(ctx context.Context) string {
|
||||
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
cachedRuntimeToken = strings.TrimSpace(token)
|
||||
return
|
||||
}
|
||||
// If the error is a decryption failure (corrupted data), log and bail out
|
||||
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return
|
||||
}
|
||||
// Try legacy manager as fallback
|
||||
manager := authpkg.NewManager(configDir, nil)
|
||||
configureLegacyAuthManagerCompatibility(manager)
|
||||
if token, _, err := manager.GetToken(); err == nil && strings.TrimSpace(token) != "" {
|
||||
cachedRuntimeToken = strings.TrimSpace(token)
|
||||
return
|
||||
if token != "" {
|
||||
cachedRuntimeToken = token
|
||||
}
|
||||
})
|
||||
return cachedRuntimeToken
|
||||
@@ -398,7 +578,7 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers = make(map[string]string)
|
||||
}
|
||||
|
||||
// Inject environment variable based headers for MCP gateway tracking
|
||||
// Inject environment variable based headers for MCP gateway tracking.
|
||||
envHeaders := map[string]string{
|
||||
"x-dingtalk-agent": os.Getenv(envDingtalkAgent),
|
||||
"x-dingtalk-trace-id": os.Getenv(envDingtalkTraceID),
|
||||
@@ -410,6 +590,12 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
// Inject third-party channel headers
|
||||
if v := os.Getenv(envDWSChannel); v != "" {
|
||||
headers["x-dws-channel"] = v
|
||||
}
|
||||
|
||||
if fn := edition.Get().MergeHeaders; fn != nil {
|
||||
headers = fn(headers)
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -25,12 +26,61 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
func setupRuntimeCommandTest(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
discoverySrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_ = json.NewEncoder(w).Encode(contactDiscoveryResponse())
|
||||
}))
|
||||
t.Cleanup(func() { discoverySrv.Close() })
|
||||
SetDiscoveryBaseURL(discoverySrv.URL)
|
||||
t.Cleanup(func() { SetDiscoveryBaseURL("") })
|
||||
}
|
||||
|
||||
func contactDiscoveryResponse() map[string]any {
|
||||
return map[string]any{
|
||||
"metadata": map[string]any{"count": 1, "nextCursor": ""},
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"server": map[string]any{
|
||||
"name": "Contact",
|
||||
"description": "通讯录",
|
||||
"remotes": []any{
|
||||
map[string]any{
|
||||
"type": "streamable-http",
|
||||
"url": "https://mcp.dingtalk.com/contact/v1",
|
||||
},
|
||||
},
|
||||
},
|
||||
"_meta": map[string]any{
|
||||
"com.dingtalk.mcp.registry/metadata": map[string]any{
|
||||
"status": "active", "isLatest": true,
|
||||
},
|
||||
"com.dingtalk.mcp.registry/cli": map[string]any{
|
||||
"id": "contact",
|
||||
"command": "contact",
|
||||
"groups": map[string]any{
|
||||
"user": map[string]any{
|
||||
"description": "用户管理",
|
||||
},
|
||||
},
|
||||
"toolOverrides": map[string]any{
|
||||
"get_current_user_profile": map[string]any{
|
||||
"cliName": "get-self",
|
||||
"group": "user",
|
||||
"flags": map[string]any{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
|
||||
@@ -596,6 +646,95 @@ func contentScanServer() *mockmcp.Server {
|
||||
return mockmcp.MustNewServer(fixture)
|
||||
}
|
||||
|
||||
func TestClassifyToolResultHookPreemptsBusinessError(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
t.Setenv("DWS_TRUSTED_DOMAINS", "*")
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, "bad request", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
method, _ := req["method"].(string)
|
||||
switch method {
|
||||
case "initialize":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": map[string]any{"tools": map[string]any{"listChanged": false}},
|
||||
"serverInfo": map[string]any{"name": "doc", "version": "1.0.0"},
|
||||
},
|
||||
})
|
||||
case "notifications/initialized":
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
case "tools/list":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"tools": []map[string]any{{
|
||||
"name": "search_documents",
|
||||
"title": "Search",
|
||||
"description": "Search documents",
|
||||
"inputSchema": map[string]any{"type": "object"},
|
||||
}},
|
||||
},
|
||||
})
|
||||
case "tools/call":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"content": map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_LOW_RISK_NO_PERMISSION",
|
||||
"data": map[string]any{"requiredScopes": []any{}},
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
hookCalled := false
|
||||
sentinelMsg := "hook-intercepted-PAT"
|
||||
edition.Override(&edition.Hooks{
|
||||
ClassifyToolResult: func(content map[string]any) error {
|
||||
if code, ok := content["code"].(string); ok && strings.Contains(code, "PAT") {
|
||||
hookCalled = true
|
||||
return fmt.Errorf("%s", sentinelMsg)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(func() { edition.Override(&edition.Hooks{}) })
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.URL, false))
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetOut(&bytes.Buffer{})
|
||||
cmd.SetErr(&bytes.Buffer{})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want hook sentinel error")
|
||||
}
|
||||
if !hookCalled {
|
||||
t.Fatal("ClassifyToolResult hook was not called")
|
||||
}
|
||||
if !strings.Contains(err.Error(), sentinelMsg) {
|
||||
t.Fatalf("error = %q, want hook sentinel %q (not generic business error)", err.Error(), sentinelMsg)
|
||||
}
|
||||
if strings.Contains(err.Error(), "business_error") || strings.Contains(err.Error(), "mcp_tool_error") {
|
||||
t.Fatalf("error = %q, should NOT contain generic framework error category", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
|
||||
+273
-18
@@ -19,7 +19,9 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -27,10 +29,24 @@ import (
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILL_API_HOST",
|
||||
Category: configmeta.CategoryNetwork,
|
||||
Description: "覆盖 Skill API 地址",
|
||||
DefaultValue: "https://mcp.dingtalk.com",
|
||||
Example: "https://custom-mcp.example.com",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// legacySkillAPIHost is the legacy skill market host used by the old cli.
|
||||
legacySkillAPIHost = "https://mcp.dingtalk.com"
|
||||
// skillDownloadEndpoint is the API endpoint for downloading skills.
|
||||
skillDownloadEndpoint = "https://aihub.dingtalk.com/cli/download"
|
||||
// skillDownloadTimeout is the timeout for skill download operations.
|
||||
@@ -51,6 +67,22 @@ type downloadSkillResult struct {
|
||||
FileName string `json:"fileName"`
|
||||
}
|
||||
|
||||
// findSkillsResponse represents the legacy skill search API response.
|
||||
type findSkillsResponse struct {
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result []CliSkillDTO `json:"result,omitempty"`
|
||||
}
|
||||
|
||||
// CliSkillDTO mirrors the old cli response payload for `skill search`.
|
||||
type CliSkillDTO struct {
|
||||
SkillID string `json:"skillId"`
|
||||
Name string `json:"name"`
|
||||
Desc string `json:"desc"`
|
||||
Icon string `json:"icon"`
|
||||
}
|
||||
|
||||
// agentSkillPaths maps target names to their relative skill installation paths.
|
||||
// These paths are relative to the user's home directory.
|
||||
var agentSkillPaths = map[string]string{
|
||||
@@ -75,7 +107,7 @@ func buildSkillCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "skill",
|
||||
Short: "技能管理",
|
||||
Long: "管理钉钉技能市场的技能。支持下载和安装技能到指定 Agent 目录。",
|
||||
Long: "管理钉钉技能市场的技能。支持搜索、下载与安装到指定 Agent 目录。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
@@ -84,13 +116,61 @@ func buildSkillCommand() *cobra.Command {
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(newSkillAddCommand())
|
||||
cmd.AddCommand(
|
||||
newSkillInstallCommand(),
|
||||
newSkillGetCommand(),
|
||||
newSkillSearchCommand(),
|
||||
newSkillFindHintCommand(),
|
||||
newSkillAddHintCommand(),
|
||||
)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddCommand() *cobra.Command {
|
||||
func newSkillGetCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "add <skillId> <target>",
|
||||
Use: "get",
|
||||
Short: "获取技能压缩文件",
|
||||
Long: "从服务端下载技能包到本地临时目录。命令执行成功后会输出临时目录路径,供调用方使用。",
|
||||
Example: " dws skill get --skill-id <skillId>",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillGet,
|
||||
}
|
||||
cmd.Flags().String("skill-id", "", "技能 ID(必填)")
|
||||
_ = cmd.MarkFlagRequired("skill-id")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillSearchCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "search",
|
||||
Short: "从钉钉技能市场搜索技能",
|
||||
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
|
||||
Example: " dws skill search --query 关键词",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillFind,
|
||||
}
|
||||
cmd.Flags().String("query", "", "搜索关键词(必填)")
|
||||
_ = cmd.MarkFlagRequired("query")
|
||||
cmd.Flags().String("scopes", "", "查询范围,空格分隔。备选值:DingtalkMarket(钉钉市场)、OrgInternal(企业内部)。为空默认查市场技能")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillFindHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "find",
|
||||
Short: "兼容旧用法,提示使用 skill search",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill search --query <关键词>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newSkillInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install <skillId> <target>",
|
||||
Short: "下载并安装技能到指定目录",
|
||||
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
|
||||
|
||||
@@ -107,9 +187,9 @@ func newSkillAddCommand() *cobra.Command {
|
||||
. -> 当前目录
|
||||
|
||||
示例:
|
||||
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
dws skill install skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill install skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill install skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillAdd,
|
||||
@@ -118,6 +198,96 @@ func newSkillAddCommand() *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "add",
|
||||
Short: "兼容旧用法,提示使用 skill install",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill install <skillId> <target>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/install?skillId=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(skillID)))
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "⬇️ 下载技能包...")
|
||||
|
||||
tmpDir, err := downloadSkillToTmpDir(cmd.Context(), apiURL, accessToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), tmpDir)
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
keyword, _ := cmd.Flags().GetString("query")
|
||||
scopes, _ := cmd.Flags().GetString("scopes")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
|
||||
if scopes != "" {
|
||||
apiURL += "&scopes=" + url.QueryEscape(scopes)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
var result findSkillsResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to parse search response: %v", err))
|
||||
}
|
||||
if !result.Success {
|
||||
errMsg := strings.TrimSpace(result.ErrorMsg)
|
||||
if errMsg == "" {
|
||||
errMsg = strings.TrimSpace(result.ErrorCode)
|
||||
}
|
||||
if errMsg == "" {
|
||||
errMsg = "unknown error"
|
||||
}
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %s", errMsg))
|
||||
}
|
||||
|
||||
if len(result.Result) == 0 {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "未找到匹配的技能")
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, skill := range result.Result {
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "SkillID: %s\n", skill.SkillID)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Name: %s\n", skill.Name)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Desc: %s\n", skill.Desc)
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "---")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
skillID := strings.TrimSpace(args[0])
|
||||
target := strings.TrimSpace(args[1])
|
||||
@@ -132,13 +302,9 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
|
||||
}
|
||||
|
||||
// Load auth token
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(cmd.Context(), skillDownloadTimeout)
|
||||
@@ -148,7 +314,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
|
||||
// Step 1: Get download URL from API
|
||||
fmt.Fprintf(w, "正在获取技能信息...\n")
|
||||
downloadResp, err := fetchSkillDownloadInfo(ctx, tokenData.AccessToken, skillID)
|
||||
downloadResp, err := fetchSkillDownloadInfo(ctx, accessToken, skillID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -189,6 +355,33 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSkillAccessToken() (string, error) {
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return "", skillAuthError()
|
||||
}
|
||||
return tokenData.AccessToken, nil
|
||||
}
|
||||
|
||||
func skillAuthError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return apperrors.NewAuth("认证信息已失效",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("请先完成钉钉账号登录后重试"))
|
||||
}
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
}
|
||||
|
||||
func skillAPIHost() string {
|
||||
if override := strings.TrimSpace(os.Getenv("DWS_SKILL_API_HOST")); override != "" {
|
||||
return strings.TrimRight(override, "/")
|
||||
}
|
||||
return legacySkillAPIHost
|
||||
}
|
||||
|
||||
// resolveSkillTargetPath resolves the target argument to an absolute path.
|
||||
func resolveSkillTargetPath(target string) (string, error) {
|
||||
target = strings.TrimSpace(target)
|
||||
@@ -236,9 +429,7 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return nil, apperrors.NewAuth("authentication failed. Please run 'dws auth login' to refresh your token",
|
||||
apperrors.WithHint("请执行 'dws auth login' 重新登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
return nil, skillAuthError()
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
@@ -259,6 +450,70 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func downloadSkillToTmpDir(ctx context.Context, apiURL, accessToken string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: skillDownloadTimeout}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill package: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
tmpDir, err := os.MkdirTemp("", "dws-skill-*")
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp dir: %v", err))
|
||||
}
|
||||
|
||||
filename := filenameFromDisposition(resp.Header.Get("Content-Disposition"))
|
||||
destPath := filepath.Join(tmpDir, filename)
|
||||
file, err := os.Create(destPath)
|
||||
if err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
if _, err := io.Copy(file, resp.Body); err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
|
||||
}
|
||||
return tmpDir, nil
|
||||
}
|
||||
|
||||
func filenameFromDisposition(cd string) string {
|
||||
if cd != "" {
|
||||
if _, params, err := mime.ParseMediaType(cd); err == nil {
|
||||
if name := strings.TrimSpace(params["filename"]); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
return "skill.zip"
|
||||
}
|
||||
|
||||
func parseLegacySkillAPIError(resp *http.Response) error {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusUnauthorized:
|
||||
return skillAuthError()
|
||||
case http.StatusBadRequest:
|
||||
return apperrors.NewValidation("request parameters are invalid")
|
||||
case http.StatusNotFound:
|
||||
return apperrors.NewValidation("skill does not exist or corresponding file was not found")
|
||||
default:
|
||||
return apperrors.NewAPI(fmt.Sprintf("skill API returned HTTP %d", resp.StatusCode),
|
||||
apperrors.WithRetryable(resp.StatusCode >= 500))
|
||||
}
|
||||
}
|
||||
|
||||
// downloadSkillFile downloads the skill zip file to a temporary location.
|
||||
func downloadSkillFile(ctx context.Context, downloadURL, fileName string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
|
||||
|
||||
@@ -319,7 +319,7 @@ func TestExtractSkillZipPreventZipSlip(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandValidation(t *testing.T) {
|
||||
func TestSkillInstallCommandValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
@@ -328,19 +328,19 @@ func TestSkillAddCommandValidation(t *testing.T) {
|
||||
}{
|
||||
{
|
||||
name: "missing arguments",
|
||||
args: []string{"skill", "add"},
|
||||
args: []string{"skill", "install"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "missing target",
|
||||
args: []string{"skill", "add", "skill-123"},
|
||||
args: []string{"skill", "install", "skill-123"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "too many arguments",
|
||||
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
|
||||
args: []string{"skill", "install", "skill-123", "qoder", "extra"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
@@ -366,7 +366,7 @@ func TestSkillAddCommandValidation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
func TestSkillInstallInvalidTarget(t *testing.T) {
|
||||
// Setup: Create config directory with valid token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
@@ -380,11 +380,11 @@ func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to save token data: %v", err)
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
|
||||
cmd.SetArgs([]string{"skill", "install", "skill-123", "invalid-target"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -399,7 +399,7 @@ func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
func TestSkillInstallRequiresAuth(t *testing.T) {
|
||||
// Setup: Create config directory without token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
@@ -411,7 +411,7 @@ func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
|
||||
cmd.SetArgs([]string{"skill", "install", "skill-123", "qoder"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -563,14 +563,16 @@ func TestSkillCommandHelp(t *testing.T) {
|
||||
if !strings.Contains(output, "技能") {
|
||||
t.Errorf("help should mention '技能', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "add") {
|
||||
t.Errorf("help should mention 'add' subcommand, got: %s", output)
|
||||
for _, subcmd := range []string{"install", "search", "get"} {
|
||||
if !strings.Contains(output, subcmd) {
|
||||
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandHelp(t *testing.T) {
|
||||
func TestSkillInstallCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "--help"})
|
||||
cmd.SetArgs([]string{"skill", "install", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -590,6 +592,56 @@ func TestSkillAddCommandHelp(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillGetCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "get"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillSearchCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "search"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillFindHintCommand(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "find"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), "dws skill search --query") {
|
||||
t.Fatalf("output = %q, want legacy hint", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSkillFileSuccess(t *testing.T) {
|
||||
// Create a mock server that returns a zip file
|
||||
expectedContent := []byte("fake zip content")
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
const stdioEndpointScheme = "stdio://"
|
||||
|
||||
var (
|
||||
stdioMu sync.RWMutex
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
)
|
||||
|
||||
// RegisterStdioClient stores a StdioClient keyed by its canonical product ID
|
||||
// (the CLI.ID used in the server descriptor). The runner looks up this client
|
||||
// when a stdio:// endpoint is resolved at execution time.
|
||||
func RegisterStdioClient(productID string, client *transport.StdioClient) {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
stdioClients[productID] = client
|
||||
}
|
||||
|
||||
// LookupStdioClient returns the StdioClient registered for the given product ID.
|
||||
// The productID can be either the full key (pluginName/serverKey) or just the serverKey.
|
||||
// This supports backward compatibility with existing CanonicalProduct values.
|
||||
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
|
||||
stdioMu.RLock()
|
||||
defer stdioMu.RUnlock()
|
||||
// Try exact match first
|
||||
if c, ok := stdioClients[productID]; ok {
|
||||
return c, true
|
||||
}
|
||||
// If not found, try matching by serverKey suffix (for backward compatibility)
|
||||
for id, c := range stdioClients {
|
||||
if idx := strings.LastIndex(id, "/"); idx >= 0 {
|
||||
if id[idx+1:] == productID {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// StdioEndpoint returns a virtual endpoint URL for a stdio-based MCP server.
|
||||
// Format: stdio://{pluginName}/{serverKey}
|
||||
func StdioEndpoint(pluginName, serverKey string) string {
|
||||
return stdioEndpointScheme + pluginName + "/" + serverKey
|
||||
}
|
||||
|
||||
// IsStdioEndpoint returns true if the endpoint uses the stdio:// scheme.
|
||||
func IsStdioEndpoint(endpoint string) bool {
|
||||
return strings.HasPrefix(endpoint, stdioEndpointScheme)
|
||||
}
|
||||
|
||||
// StopAllStdioClients stops all registered stdio clients.
|
||||
// This should be called on program exit to terminate child processes.
|
||||
func StopAllStdioClients() {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
for id, client := range stdioClients {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
}
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
}
|
||||
|
||||
// StopStdioClient stops a specific stdio client by product ID.
|
||||
// Returns true if the client was found and stopped, false otherwise.
|
||||
func StopStdioClient(productID string) bool {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
client, ok := stdioClients[productID]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", productID, "error", err)
|
||||
}
|
||||
delete(stdioClients, productID)
|
||||
return true
|
||||
}
|
||||
|
||||
// StopStdioClientsByPlugin stops all stdio clients belonging to a plugin.
|
||||
// The productID format is "pluginName/serverKey". This function stops all
|
||||
// clients whose productID has the given pluginName prefix.
|
||||
func StopStdioClientsByPlugin(pluginName string) int {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
prefix := pluginName + "/"
|
||||
count := 0
|
||||
for id, client := range stdioClients {
|
||||
if len(id) > len(prefix) && id[:len(prefix)] == prefix {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
delete(stdioClients, id)
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
// 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 app
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestStdioEndpoint(t *testing.T) {
|
||||
endpoint := StdioEndpoint("hello-plugin", "hello")
|
||||
want := "stdio://hello-plugin/hello"
|
||||
if endpoint != want {
|
||||
t.Errorf("StdioEndpoint() = %q, want %q", endpoint, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsStdioEndpoint(t *testing.T) {
|
||||
tests := []struct {
|
||||
endpoint string
|
||||
want bool
|
||||
}{
|
||||
{"stdio://hello-plugin/hello", true},
|
||||
{"stdio://conference/local", true},
|
||||
{"https://mcp.dingtalk.com", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := IsStdioEndpoint(tt.endpoint); got != tt.want {
|
||||
t.Errorf("IsStdioEndpoint(%q) = %v, want %v", tt.endpoint, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioClientRegistry(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
stdioMu.Lock()
|
||||
delete(stdioClients, "test-product")
|
||||
stdioMu.Unlock()
|
||||
}()
|
||||
|
||||
// Initially not found
|
||||
if _, ok := LookupStdioClient("test-product"); ok {
|
||||
t.Error("expected LookupStdioClient to return false for unregistered product")
|
||||
}
|
||||
|
||||
// Register a client
|
||||
client := transport.NewStdioClient("echo", nil, nil)
|
||||
RegisterStdioClient("test-product", client)
|
||||
|
||||
// Now should be found
|
||||
got, ok := LookupStdioClient("test-product")
|
||||
if !ok {
|
||||
t.Fatal("expected LookupStdioClient to return true after registration")
|
||||
}
|
||||
if got != client {
|
||||
t.Error("LookupStdioClient returned different client instance")
|
||||
}
|
||||
}
|
||||
+192
-2
@@ -15,16 +15,45 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
// PerfDebugEnv is the environment variable to enable performance timing output.
|
||||
const PerfDebugEnv = "DWS_PERF_DEBUG"
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PERF_DEBUG",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "启用性能计时输出到 stderr",
|
||||
Example: "1",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PERF_REPORT",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "JSON 性能报告输出路径 (auto=~/.dws/perf/latest.json)",
|
||||
Example: "auto",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// PerfDebugEnv is the environment variable to enable performance timing output.
|
||||
PerfDebugEnv = "DWS_PERF_DEBUG"
|
||||
|
||||
// PerfReportEnv is the environment variable to enable JSON perf report output.
|
||||
// Set to "auto" to write to ~/.dws/perf/latest.json, or a custom file path.
|
||||
PerfReportEnv = "DWS_PERF_REPORT"
|
||||
|
||||
perfReportDir = "perf"
|
||||
perfReportFile = "latest.json"
|
||||
)
|
||||
|
||||
// timingContextKey is the context key for TimingCollector.
|
||||
type timingContextKey struct{}
|
||||
@@ -189,3 +218,164 @@ func StartTiming(ctx context.Context, name string) func() {
|
||||
func IsPerfDebugEnabled() bool {
|
||||
return os.Getenv(PerfDebugEnv) != ""
|
||||
}
|
||||
|
||||
// ── Structured Performance Report ──────────────────────────────────────
|
||||
|
||||
// PerfPhase is a single phase in the performance report.
|
||||
type PerfPhase struct {
|
||||
Name string `json:"name"`
|
||||
DurationMs int64 `json:"duration_ms"`
|
||||
Seq int `json:"seq"`
|
||||
}
|
||||
|
||||
// PerfReport is the JSON-serialisable performance report.
|
||||
type PerfReport struct {
|
||||
Kind string `json:"kind"`
|
||||
Version string `json:"version"`
|
||||
CLIVersion string `json:"cli_version"`
|
||||
Command string `json:"command"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
TotalMs int64 `json:"total_ms"`
|
||||
Phases []PerfPhase `json:"phases"`
|
||||
Slowest string `json:"slowest"`
|
||||
OverheadMs int64 `json:"overhead_ms"`
|
||||
}
|
||||
|
||||
// BuildReport constructs a PerfReport from the collected timing entries.
|
||||
func (tc *TimingCollector) BuildReport(cliVersion, command string) PerfReport {
|
||||
entries := tc.Entries()
|
||||
total := tc.Total()
|
||||
totalMs := total.Milliseconds()
|
||||
|
||||
phases := make([]PerfPhase, len(entries))
|
||||
var sumMs int64
|
||||
var slowestName string
|
||||
var slowestMs int64
|
||||
|
||||
for i, e := range entries {
|
||||
ms := e.Duration.Milliseconds()
|
||||
phases[i] = PerfPhase{
|
||||
Name: e.Name,
|
||||
DurationMs: ms,
|
||||
Seq: e.Seq,
|
||||
}
|
||||
sumMs += ms
|
||||
if ms > slowestMs {
|
||||
slowestMs = ms
|
||||
slowestName = e.Name
|
||||
}
|
||||
}
|
||||
|
||||
overhead := totalMs - sumMs
|
||||
if overhead < 0 {
|
||||
overhead = 0
|
||||
}
|
||||
|
||||
return PerfReport{
|
||||
Kind: "perf_report",
|
||||
Version: "1",
|
||||
CLIVersion: cliVersion,
|
||||
Command: command,
|
||||
Timestamp: time.Now(),
|
||||
TotalMs: totalMs,
|
||||
Phases: phases,
|
||||
Slowest: slowestName,
|
||||
OverheadMs: overhead,
|
||||
}
|
||||
}
|
||||
|
||||
// WriteReportIfEnabled checks DWS_PERF_REPORT and writes a JSON report if set.
|
||||
func (tc *TimingCollector) WriteReportIfEnabled(cliVersion, command string) {
|
||||
if tc == nil {
|
||||
return
|
||||
}
|
||||
dest := os.Getenv(PerfReportEnv)
|
||||
if dest == "" {
|
||||
return
|
||||
}
|
||||
|
||||
report := tc.BuildReport(cliVersion, command)
|
||||
data, err := json.MarshalIndent(report, "", " ")
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
path := resolvePerfReportPath(dest)
|
||||
if path == "" {
|
||||
return
|
||||
}
|
||||
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
tmp := path + ".tmp"
|
||||
if err := os.WriteFile(tmp, data, 0o600); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return
|
||||
}
|
||||
_ = os.Rename(tmp, path)
|
||||
}
|
||||
|
||||
// LoadLatestReport reads the default perf report file (~/.dws/perf/latest.json).
|
||||
func LoadLatestReport() (*PerfReport, error) {
|
||||
path := defaultPerfReportPath()
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var report PerfReport
|
||||
if err := json.Unmarshal(data, &report); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &report, nil
|
||||
}
|
||||
|
||||
// resolvePerfReportPath resolves the DWS_PERF_REPORT value to an absolute path.
|
||||
func resolvePerfReportPath(dest string) string {
|
||||
if dest == "auto" {
|
||||
return defaultPerfReportPath()
|
||||
}
|
||||
return dest
|
||||
}
|
||||
|
||||
func defaultPerfReportPath() string {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return filepath.Join(home, ".dws", perfReportDir, perfReportFile)
|
||||
}
|
||||
|
||||
// sensitiveFlags are flag names whose values should be masked in commands.
|
||||
var sensitiveFlags = map[string]bool{
|
||||
"--token": true,
|
||||
"--client-secret": true,
|
||||
"--client-id": true,
|
||||
}
|
||||
|
||||
// SanitizeCommand redacts sensitive flag values from a command arg slice.
|
||||
func SanitizeCommand(args []string) string {
|
||||
sanitized := make([]string, 0, len(args))
|
||||
skipNext := false
|
||||
for _, arg := range args {
|
||||
if skipNext {
|
||||
sanitized = append(sanitized, "***")
|
||||
skipNext = false
|
||||
continue
|
||||
}
|
||||
if idx := strings.IndexByte(arg, '='); idx > 0 {
|
||||
key := arg[:idx]
|
||||
if sensitiveFlags[key] {
|
||||
sanitized = append(sanitized, key+"=***")
|
||||
continue
|
||||
}
|
||||
}
|
||||
if sensitiveFlags[arg] {
|
||||
skipNext = true
|
||||
}
|
||||
sanitized = append(sanitized, arg)
|
||||
}
|
||||
return strings.Join(sanitized, " ")
|
||||
}
|
||||
|
||||
@@ -16,7 +16,9 @@ package app
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -136,6 +138,7 @@ func TestTimingCollector_ContextIntegration(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestTimingCollectorFromContext_NilContext(t *testing.T) {
|
||||
//lint:ignore SA1012 Testing explicit nil-context guard in TimingCollectorFromContext.
|
||||
tc := TimingCollectorFromContext(nil)
|
||||
if tc != nil {
|
||||
t.Error("TimingCollectorFromContext(nil) should return nil")
|
||||
@@ -171,3 +174,280 @@ func TestIsPerfDebugEnabled(t *testing.T) {
|
||||
t.Error("IsPerfDebugEnabled should return true when env var is set")
|
||||
}
|
||||
}
|
||||
|
||||
// ── PerfReport tests ────────────────────────────────────────────────────
|
||||
|
||||
func TestBuildReport(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 45*time.Millisecond)
|
||||
tc.Record("auth_keychain", 72*time.Millisecond)
|
||||
tc.Record("mcp_call", 620*time.Millisecond)
|
||||
|
||||
report := tc.BuildReport("v1.0.8", "dws aitable list-records")
|
||||
|
||||
if report.Kind != "perf_report" {
|
||||
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
|
||||
}
|
||||
if report.Version != "1" {
|
||||
t.Errorf("expected version '1', got %q", report.Version)
|
||||
}
|
||||
if report.CLIVersion != "v1.0.8" {
|
||||
t.Errorf("expected cli_version 'v1.0.8', got %q", report.CLIVersion)
|
||||
}
|
||||
if report.Command != "dws aitable list-records" {
|
||||
t.Errorf("expected command 'dws aitable list-records', got %q", report.Command)
|
||||
}
|
||||
if len(report.Phases) != 3 {
|
||||
t.Fatalf("expected 3 phases, got %d", len(report.Phases))
|
||||
}
|
||||
if report.Phases[0].Name != "cmd_init" || report.Phases[0].DurationMs != 45 {
|
||||
t.Errorf("unexpected first phase: %+v", report.Phases[0])
|
||||
}
|
||||
if report.Slowest != "mcp_call" {
|
||||
t.Errorf("expected slowest 'mcp_call', got %q", report.Slowest)
|
||||
}
|
||||
if report.TotalMs < 0 {
|
||||
t.Errorf("total_ms should be >= 0, got %d", report.TotalMs)
|
||||
}
|
||||
if report.OverheadMs < 0 {
|
||||
t.Errorf("overhead_ms should be >= 0, got %d", report.OverheadMs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReportEmpty(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
report := tc.BuildReport("dev", "dws version")
|
||||
|
||||
if len(report.Phases) != 0 {
|
||||
t.Errorf("expected 0 phases, got %d", len(report.Phases))
|
||||
}
|
||||
if report.Slowest != "" {
|
||||
t.Errorf("expected empty slowest, got %q", report.Slowest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReportJSON(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 10*time.Millisecond)
|
||||
|
||||
report := tc.BuildReport("v1.0.0", "dws version")
|
||||
data, err := json.Marshal(report)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatalf("json.Unmarshal failed: %v", err)
|
||||
}
|
||||
|
||||
requiredKeys := []string{"kind", "version", "cli_version", "command", "timestamp", "total_ms", "phases", "slowest", "overhead_ms"}
|
||||
for _, key := range requiredKeys {
|
||||
if _, ok := parsed[key]; !ok {
|
||||
t.Errorf("missing key %q in JSON output", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
reportPath := filepath.Join(dir, "report.json")
|
||||
|
||||
t.Setenv(PerfReportEnv, reportPath)
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 50*time.Millisecond)
|
||||
tc.Record("mcp_call", 200*time.Millisecond)
|
||||
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
|
||||
data, err := os.ReadFile(reportPath)
|
||||
if err != nil {
|
||||
t.Fatalf("report file not written: %v", err)
|
||||
}
|
||||
|
||||
var report PerfReport
|
||||
if err := json.Unmarshal(data, &report); err != nil {
|
||||
t.Fatalf("invalid JSON in report: %v", err)
|
||||
}
|
||||
if report.Kind != "perf_report" {
|
||||
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
|
||||
}
|
||||
if len(report.Phases) != 2 {
|
||||
t.Errorf("expected 2 phases, got %d", len(report.Phases))
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_Auto(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
expected := filepath.Join(tmpHome, ".dws", "perf", "latest.json")
|
||||
|
||||
// Temporarily override HOME for defaultPerfReportPath
|
||||
t.Setenv("HOME", tmpHome)
|
||||
t.Setenv(PerfReportEnv, "auto")
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 10*time.Millisecond)
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
|
||||
if _, err := os.Stat(expected); err != nil {
|
||||
t.Fatalf("expected report at %s: %v", expected, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_Disabled(t *testing.T) {
|
||||
t.Setenv(PerfReportEnv, "")
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("op", 10*time.Millisecond)
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
// No file should be written; no error expected
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_NilCollector(t *testing.T) {
|
||||
t.Setenv(PerfReportEnv, "/tmp/should-not-exist.json")
|
||||
var tc *TimingCollector
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
}
|
||||
|
||||
func TestLoadLatestReport(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
t.Setenv("HOME", tmpHome)
|
||||
|
||||
perfDir := filepath.Join(tmpHome, ".dws", "perf")
|
||||
if err := os.MkdirAll(perfDir, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
report := PerfReport{
|
||||
Kind: "perf_report",
|
||||
Version: "1",
|
||||
CLIVersion: "v1.0.0",
|
||||
Command: "dws version",
|
||||
TotalMs: 100,
|
||||
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}},
|
||||
Slowest: "cmd_init",
|
||||
OverheadMs: 50,
|
||||
}
|
||||
data, _ := json.MarshalIndent(report, "", " ")
|
||||
if err := os.WriteFile(filepath.Join(perfDir, "latest.json"), data, 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
loaded, err := LoadLatestReport()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadLatestReport failed: %v", err)
|
||||
}
|
||||
if loaded.CLIVersion != "v1.0.0" {
|
||||
t.Errorf("expected cli_version 'v1.0.0', got %q", loaded.CLIVersion)
|
||||
}
|
||||
if len(loaded.Phases) != 1 {
|
||||
t.Errorf("expected 1 phase, got %d", len(loaded.Phases))
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadLatestReport_NotFound(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
t.Setenv("HOME", tmpHome)
|
||||
|
||||
_, err := LoadLatestReport()
|
||||
if err == nil {
|
||||
t.Error("expected error when report file does not exist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeCommand(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "no sensitive flags",
|
||||
args: []string{"dws", "aitable", "list-records"},
|
||||
want: "dws aitable list-records",
|
||||
},
|
||||
{
|
||||
name: "token with space-separated value",
|
||||
args: []string{"dws", "--token", "secret123", "version"},
|
||||
want: "dws --token *** version",
|
||||
},
|
||||
{
|
||||
name: "token with equals sign",
|
||||
args: []string{"dws", "--token=secret123", "version"},
|
||||
want: "dws --token=*** version",
|
||||
},
|
||||
{
|
||||
name: "client-secret space-separated",
|
||||
args: []string{"dws", "--client-secret", "mysecret", "--client-id", "myid", "auth"},
|
||||
want: "dws --client-secret *** --client-id *** auth",
|
||||
},
|
||||
{
|
||||
name: "client-id with equals",
|
||||
args: []string{"dws", "--client-id=abc123"},
|
||||
want: "dws --client-id=***",
|
||||
},
|
||||
{
|
||||
name: "empty args",
|
||||
args: []string{},
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := SanitizeCommand(tt.args)
|
||||
if got != tt.want {
|
||||
t.Errorf("SanitizeCommand(%v) = %q, want %q", tt.args, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePerfReportPath_Auto(t *testing.T) {
|
||||
p := resolvePerfReportPath("auto")
|
||||
if p == "" {
|
||||
t.Skip("HOME not available")
|
||||
}
|
||||
if !strings.HasSuffix(p, filepath.Join("perf", "latest.json")) {
|
||||
t.Errorf("expected path ending in perf/latest.json, got %q", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePerfReportPath_Custom(t *testing.T) {
|
||||
p := resolvePerfReportPath("/tmp/my-report.json")
|
||||
if p != "/tmp/my-report.json" {
|
||||
t.Errorf("expected '/tmp/my-report.json', got %q", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPerfReportSummary(t *testing.T) {
|
||||
report := &PerfReport{
|
||||
Command: "dws version",
|
||||
Timestamp: time.Now(),
|
||||
TotalMs: 300,
|
||||
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}, {Name: "mcp_call", DurationMs: 200, Seq: 1}},
|
||||
Slowest: "mcp_call",
|
||||
OverheadMs: 50,
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
printPerfReportSummary(&buf, report)
|
||||
out := buf.String()
|
||||
|
||||
if !strings.Contains(out, "cmd_init") {
|
||||
t.Error("output should contain 'cmd_init'")
|
||||
}
|
||||
if !strings.Contains(out, "mcp_call") {
|
||||
t.Error("output should contain 'mcp_call'")
|
||||
}
|
||||
if !strings.Contains(out, "← 最慢") {
|
||||
t.Error("output should contain '← 最慢' marker")
|
||||
}
|
||||
if !strings.Contains(out, "总耗时") {
|
||||
t.Error("output should contain '总耗时'")
|
||||
}
|
||||
if !strings.Contains(out, "框架开销") {
|
||||
t.Error("output should contain '框架开销'")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
// 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 "testing"
|
||||
|
||||
func TestClassifyDenialReason(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
status *CLIAuthStatus
|
||||
currentChannel string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "error CHANNEL_REQUIRED",
|
||||
status: &CLIAuthStatus{
|
||||
ErrorCode: "CHANNEL_REQUIRED",
|
||||
},
|
||||
want: "channel_required",
|
||||
},
|
||||
{
|
||||
name: "error NO_AUTH",
|
||||
status: &CLIAuthStatus{
|
||||
ErrorCode: "NO_AUTH",
|
||||
},
|
||||
want: "no_auth",
|
||||
},
|
||||
{
|
||||
name: "success false or nil result → unknown",
|
||||
status: &CLIAuthStatus{
|
||||
Success: false,
|
||||
},
|
||||
want: "unknown",
|
||||
},
|
||||
{
|
||||
name: "cliAuthEnabled true → no denial",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "userScope forbidden wins over channel",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "forbidden",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "channel-b",
|
||||
want: "user_forbidden",
|
||||
},
|
||||
{
|
||||
// Real-world case reported: user is in allowedUsers but the current
|
||||
// DWS_CHANNEL is not in allowedChannels. Reason must be channel,
|
||||
// NOT user.
|
||||
name: "user allowed but channel not in allowedChannels → channel_not_allowed",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
AllowedUsers: []string{"014566033934857460"},
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"2a4a658e467998befb7fa333c19ba2b3a3bacfa4"},
|
||||
},
|
||||
},
|
||||
currentChannel: "different-channel",
|
||||
want: "channel_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "channelScope specified but current channel empty → channel_required",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "",
|
||||
want: "channel_required",
|
||||
},
|
||||
{
|
||||
name: "channel matches allowedChannels → fall back to user denial",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "channel-a",
|
||||
want: "user_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "only userScope=specified, no channel restriction → user_not_allowed",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
},
|
||||
},
|
||||
currentChannel: "",
|
||||
want: "user_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "no user or channel restriction → cli_not_enabled",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
},
|
||||
want: "cli_not_enabled",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := classifyDenialReason(tc.status, tc.currentChannel)
|
||||
if got != tc.want {
|
||||
t.Fatalf("classifyDenialReason() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -154,9 +154,7 @@ func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
@@ -169,7 +167,7 @@ func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !status.Success || !status.Result.CLIAuthEnabled {
|
||||
if !status.Success || status.Result == nil || !status.Result.CLIAuthEnabled {
|
||||
t.Fatalf("expected CLIAuthEnabled=true, got %+v", status)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
@@ -190,9 +188,7 @@ func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
@@ -204,7 +200,7 @@ func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !status.Result.CLIAuthEnabled {
|
||||
if status.Result == nil || !status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=true")
|
||||
}
|
||||
t.Logf("✅ Normal enabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
@@ -215,9 +211,7 @@ func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
@@ -229,7 +223,7 @@ func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status.Result.CLIAuthEnabled {
|
||||
if status.Result == nil || status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=false")
|
||||
}
|
||||
t.Logf("✅ Normal disabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
@@ -262,12 +256,10 @@ func TestOAuthCallback_CLIAuthEnabled_ShowsSuccessPage(t *testing.T) {
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && authStatus.Result.CLIAuthEnabled
|
||||
if !cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be true when API returns enabled")
|
||||
}
|
||||
@@ -280,12 +272,10 @@ func TestOAuthCallback_CLIAuthDisabledByServer_ShowsNotEnabledPage(t *testing.T)
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && authStatus.Result.CLIAuthEnabled
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when server says disabled")
|
||||
}
|
||||
@@ -314,12 +304,18 @@ func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
@@ -340,13 +336,14 @@ func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
@@ -378,12 +375,18 @@ func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
@@ -396,9 +399,7 @@ func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, SuperAdminPath):
|
||||
@@ -415,13 +416,14 @@ func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
@@ -450,12 +452,18 @@ func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
@@ -468,9 +476,7 @@ func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
|
||||
default:
|
||||
@@ -481,13 +487,14 @@ func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
token, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
+213
-39
@@ -33,32 +33,37 @@ import (
|
||||
|
||||
const (
|
||||
// defaultPollInterval is the default seconds between device token polls.
|
||||
defaultPollInterval = 5
|
||||
// The server-side Redis TTL is 10 minutes; a 2-second interval keeps the
|
||||
// user-perceived latency low while staying well within rate limits.
|
||||
defaultPollInterval = 2
|
||||
// maxPollInterval caps the polling interval to prevent DoS via slow_down.
|
||||
maxPollInterval = 30
|
||||
// maxPollTotalWait caps the total wait time for device authorization.
|
||||
maxPollTotalWait = 15 * time.Minute
|
||||
// Aligned with the server-side Redis TTL (10 minutes).
|
||||
maxPollTotalWait = 10 * time.Minute
|
||||
)
|
||||
|
||||
type DeviceFlowProvider struct {
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
terminalBaseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewDeviceFlowProvider(configDir string, logger *slog.Logger) *DeviceFlowProvider {
|
||||
return &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: ClientID(),
|
||||
scope: DefaultScopes,
|
||||
baseURL: DefaultDeviceBaseURL,
|
||||
logger: logger,
|
||||
Output: os.Stderr,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
configDir: configDir,
|
||||
clientID: ClientID(),
|
||||
scope: DefaultScopes,
|
||||
baseURL: DefaultDeviceBaseURL,
|
||||
terminalBaseURL: GetMCPBaseURL(),
|
||||
logger: logger,
|
||||
Output: os.Stderr,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -66,6 +71,18 @@ func (p *DeviceFlowProvider) SetBaseURL(baseURL string) {
|
||||
p.baseURL = strings.TrimRight(baseURL, "/")
|
||||
}
|
||||
|
||||
// SetTerminalBaseURL sets the terminal API base URL for device flow polling.
|
||||
func (p *DeviceFlowProvider) SetTerminalBaseURL(baseURL string) {
|
||||
p.terminalBaseURL = strings.TrimRight(baseURL, "/")
|
||||
}
|
||||
|
||||
// SetScope overrides the OAuth scope for the device flow.
|
||||
func (p *DeviceFlowProvider) SetScope(scope string) {
|
||||
if p != nil {
|
||||
p.scope = scope
|
||||
}
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) output() io.Writer {
|
||||
if p != nil && p.Output != nil {
|
||||
return p.Output
|
||||
@@ -80,6 +97,7 @@ type DeviceAuthResponse struct {
|
||||
VerificationURIComplete string `json:"verificationUriComplete"`
|
||||
ExpiresIn int `json:"expiresIn"`
|
||||
Interval int `json:"interval"`
|
||||
FlowID string `json:"flowId"`
|
||||
}
|
||||
|
||||
type DeviceTokenResponse struct {
|
||||
@@ -88,6 +106,20 @@ type DeviceTokenResponse struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
// DevicePollResponse represents the response from the terminal API poll endpoint.
|
||||
type DevicePollResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Code string `json:"code,omitempty"`
|
||||
Message string `json:"message,omitempty"`
|
||||
Data DevicePollData `json:"data"`
|
||||
}
|
||||
|
||||
type DevicePollData struct {
|
||||
Status string `json:"status"`
|
||||
AuthCode string `json:"authCode,omitempty"`
|
||||
FlowID string `json:"flowId,omitempty"`
|
||||
}
|
||||
|
||||
type serviceResult struct {
|
||||
Success bool `json:"success"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
@@ -176,33 +208,61 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请检查网络连接后重试。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("检查 CLI 授权状态失败"), authErr)
|
||||
} else if authStatus.Success && !authStatus.Result.CLIAuthEnabled {
|
||||
// CLI auth is disabled - show detailed error with admin info
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
|
||||
}
|
||||
denialReason := classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
|
||||
if denialReason != "" {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
switch denialReason {
|
||||
case "user_forbidden":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织已禁止所有成员使用 CLI")))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织已禁止所有成员使用 CLI"))
|
||||
case "user_not_allowed":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 您不在该组织的 CLI 授权人员范围内")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员将您加入 CLI 授权人员名单。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员"))
|
||||
case "channel_not_allowed":
|
||||
ch := os.Getenv("DWS_CHANNEL")
|
||||
_, _ = fmt.Fprintf(p.output(), dfRed(i18n.T("⚠️ 当前渠道 %s 未获得该组织授权"))+"\n", ch)
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员开通该渠道的访问权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf(i18n.T("当前渠道 %s 未获得该组织授权,请联系组织管理员"), ch)
|
||||
case "channel_required":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 当前组织已开启渠道管控")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请升级到最新版本的 CLI 后重试。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("当前组织已开启渠道管控,请升级到最新版本的 CLI 后重试"))
|
||||
case "no_auth":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 认证已失效")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请执行 dws auth 重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("认证已失效,请执行 dws auth 重新登录"))
|
||||
default:
|
||||
// cli_not_enabled or unknown — show existing admin-apply flow
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
|
||||
// Try to get super admin list
|
||||
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
|
||||
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
|
||||
// Show up to 3 admins
|
||||
maxAdmins := 3
|
||||
if len(admins.Result) < maxAdmins {
|
||||
maxAdmins = len(admins.Result)
|
||||
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
|
||||
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
|
||||
maxAdmins := 3
|
||||
if len(admins.Result) < maxAdmins {
|
||||
maxAdmins = len(admins.Result)
|
||||
}
|
||||
var adminNames []string
|
||||
for i := 0; i < maxAdmins; i++ {
|
||||
adminNames = append(adminNames, admins.Result[i].Name)
|
||||
}
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
|
||||
}
|
||||
var adminNames []string
|
||||
for i := 0; i < maxAdmins; i++ {
|
||||
adminNames = append(adminNames, admins.Result[i].Name)
|
||||
}
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("管理员操作入口:"), config.GetDeveloperSettingsURL())
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
|
||||
}
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
@@ -211,6 +271,13 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Always persist clientId to app.json so future process startups
|
||||
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
|
||||
if p.clientID != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
|
||||
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
oauthProvider.persistAppConfigIfNeeded()
|
||||
|
||||
@@ -279,7 +346,91 @@ func (p *DeviceFlowProvider) pollDeviceToken(ctx context.Context, deviceCode str
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
// pollDeviceStatus polls the terminal API for device authorization status.
|
||||
//
|
||||
// Note: The server returns success=false for REJECTED and EXPIRED terminal
|
||||
// states (with a valid data.Status value). These are normal business outcomes,
|
||||
// not transport errors, so we return the response to the caller and let the
|
||||
// status-switch handle them.
|
||||
func (p *DeviceFlowProvider) pollDeviceStatus(ctx context.Context, flowID string) (*DevicePollResponse, error) {
|
||||
endpoint := fmt.Sprintf("%s%s?flowId=%s", p.terminalBaseURL, DevicePollPath, url.QueryEscape(flowID))
|
||||
body, err := p.doGet(ctx, endpoint)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var resp DevicePollResponse
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("解析响应失败"), err)
|
||||
}
|
||||
// REJECTED/EXPIRED carry success=false but have a valid data.Status;
|
||||
// only treat as a real server error when data.Status is empty.
|
||||
if !resp.Success && resp.Data.Status == "" {
|
||||
return nil, fmt.Errorf("%s: [%s] %s", i18n.T("服务端返回错误"), resp.Code, resp.Message)
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorization(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
if auth.FlowID == "" {
|
||||
// Keep the pre-flowId device-code polling path for regular device flow
|
||||
// login responses that do not include terminal polling metadata.
|
||||
return p.waitForAuthorizationByDeviceCode(ctx, auth)
|
||||
}
|
||||
return p.waitForAuthorizationByFlowID(ctx, auth)
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorizationByFlowID(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
startTime := time.Now()
|
||||
interval := time.Duration(auth.Interval) * time.Second
|
||||
deadline := time.Duration(auth.ExpiresIn) * time.Second
|
||||
pollCount := 0
|
||||
|
||||
for {
|
||||
elapsed := time.Since(startTime)
|
||||
if elapsed >= maxPollTotalWait || elapsed >= deadline {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf("%s", i18n.Tf("设备授权码已过期(%d 秒),请重试", auth.ExpiresIn))
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-time.After(interval):
|
||||
}
|
||||
|
||||
pollCount++
|
||||
elapsedSec := int(time.Since(startTime).Seconds())
|
||||
dfPrintPollStatus(p.output(), pollCount, elapsedSec)
|
||||
|
||||
pollResp, err := p.pollDeviceStatus(ctx, auth.FlowID)
|
||||
if err != nil {
|
||||
dfPrintPollResult(p.output(), "network_error", i18n.T("网络错误,继续重试..."))
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("poll error", "error", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
switch pollResp.Data.Status {
|
||||
case StatusApproved:
|
||||
dfPrintPollResult(p.output(), "authorized", i18n.T("授权成功!"))
|
||||
return &DeviceTokenResponse{AuthCode: pollResp.Data.AuthCode}, nil
|
||||
case StatusPending:
|
||||
dfPrintPollResult(p.output(), "pending", i18n.T("等待用户授权..."))
|
||||
case StatusRejected:
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("用户拒绝了授权请求"))
|
||||
case StatusExpired:
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("设备授权码已过期"))
|
||||
default:
|
||||
dfPrintPollResult(p.output(), "unknown", fmt.Sprintf(i18n.T("未知状态: %s"), pollResp.Data.Status))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorizationByDeviceCode(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
startTime := time.Now()
|
||||
interval := time.Duration(auth.Interval) * time.Second
|
||||
deadline := time.Duration(auth.ExpiresIn) * time.Second
|
||||
@@ -359,6 +510,29 @@ func (p *DeviceFlowProvider) postForm(ctx context.Context, endpoint string, para
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// doGet performs an HTTP GET request and returns the response body.
|
||||
func (p *DeviceFlowProvider) doGet(ctx context.Context, endpoint string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("创建请求失败"), err)
|
||||
}
|
||||
|
||||
resp, err := p.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("发送请求失败"), err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("读取响应失败"), err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncateBody(body, 200))
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// truncateBody returns a string of at most maxLen bytes from body, appending
|
||||
// "...(truncated)" when the content exceeds the limit. This prevents leaking
|
||||
// potentially sensitive response payloads in error messages.
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
// 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
|
||||
|
||||
// Device flow authorization status constants.
|
||||
// Shared across device_flow.go and pat_auth_retry.go to avoid maintaining
|
||||
// string literals in multiple places.
|
||||
const (
|
||||
StatusPending = "PENDING"
|
||||
StatusApproved = "APPROVED"
|
||||
StatusRejected = "REJECTED"
|
||||
StatusExpired = "EXPIRED"
|
||||
StatusCancelled = "CANCELLED"
|
||||
)
|
||||
|
||||
// ParseDeviceFlowStatus normalizes a raw status string from the device flow
|
||||
// poll response into a canonical status constant. When the server returns an
|
||||
// empty status with success=false, it falls back to StatusExpired (server
|
||||
// error / flow not found).
|
||||
func ParseDeviceFlowStatus(rawStatus string, success bool) string {
|
||||
switch rawStatus {
|
||||
case StatusApproved, StatusRejected, StatusExpired, StatusPending, StatusCancelled:
|
||||
return rawStatus
|
||||
default:
|
||||
if rawStatus == "" && !success {
|
||||
return StatusExpired
|
||||
}
|
||||
return rawStatus
|
||||
}
|
||||
}
|
||||
@@ -14,15 +14,21 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
)
|
||||
|
||||
func newDeviceFlowTestLogger() *slog.Logger {
|
||||
@@ -89,22 +95,42 @@ func TestWaitForAuthorizationSucceedsAfterPending(t *testing.T) {
|
||||
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// New terminal API uses GET method
|
||||
if r.Method != http.MethodGet {
|
||||
t.Fatalf("method = %s, want GET", r.Method)
|
||||
}
|
||||
if !strings.Contains(r.URL.RawQuery, "flowId=") {
|
||||
t.Fatal("flowId query parameter should be present")
|
||||
}
|
||||
if calls.Add(1) <= 2 {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
// Return PENDING status
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{"status": "PENDING"},
|
||||
})
|
||||
return
|
||||
}
|
||||
writeServiceResult(w, true, DeviceTokenResponse{AuthCode: "final-auth-code"}, "", "")
|
||||
// Return APPROVED status with authCode
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "final-auth-code",
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
provider.Output = io.Discard
|
||||
provider.SetBaseURL(server.URL)
|
||||
provider.SetTerminalBaseURL(server.URL)
|
||||
|
||||
resp, err := provider.waitForAuthorization(context.Background(), &DeviceAuthResponse{
|
||||
DeviceCode: "dc-1",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("waitForAuthorization() error = %v", err)
|
||||
@@ -117,9 +143,87 @@ func TestWaitForAuthorizationSucceedsAfterPending(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationFallsBackToDeviceCodeWhenFlowIDMissing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
t.Fatalf("method = %s, want POST", r.Method)
|
||||
}
|
||||
if err := r.ParseForm(); err != nil {
|
||||
t.Fatalf("ParseForm() error = %v", err)
|
||||
}
|
||||
if got := r.FormValue("device_code"); got != "legacy-device-code" {
|
||||
t.Fatalf("device_code = %q, want legacy-device-code", got)
|
||||
}
|
||||
if calls.Add(1) <= 2 {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
return
|
||||
}
|
||||
writeServiceResult(w, true, DeviceTokenResponse{AuthCode: "legacy-auth-code"}, "", "")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
var output bytes.Buffer
|
||||
provider.Output = &output
|
||||
provider.SetBaseURL(server.URL)
|
||||
|
||||
resp, err := provider.waitForAuthorization(context.Background(), &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("waitForAuthorization() error = %v", err)
|
||||
}
|
||||
if resp.AuthCode != "legacy-auth-code" {
|
||||
t.Fatalf("auth code = %q, want legacy-auth-code", resp.AuthCode)
|
||||
}
|
||||
if calls.Load() != 3 {
|
||||
t.Fatalf("poll calls = %d, want 3", calls.Load())
|
||||
}
|
||||
if !strings.Contains(output.String(), i18n.T("等待用户授权...")) {
|
||||
t.Fatalf("expected device-code path to emit pending output, got %q", output.String())
|
||||
}
|
||||
if !strings.Contains(output.String(), i18n.T("授权成功!")) {
|
||||
t.Fatalf("expected device-code path to emit success output, got %q", output.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationHonorsContextCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// New terminal API uses GET method
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{"status": "PENDING"},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
provider.Output = io.Discard
|
||||
provider.SetTerminalBaseURL(server.URL)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
if _, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
FlowID: "test-flow-id-2",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}); err == nil {
|
||||
t.Fatal("waitForAuthorization() error = nil, want context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationByDeviceCodeHonorsContextCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
}))
|
||||
@@ -132,11 +236,89 @@ func TestWaitForAuthorizationHonorsContextCancellation(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
if _, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "dc-2",
|
||||
_, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}); err == nil {
|
||||
t.Fatal("waitForAuthorization() error = nil, want context cancellation")
|
||||
})
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want context deadline exceeded", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationByDeviceCodeErrorStates(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
timeout time.Duration
|
||||
responses []DeviceTokenResponse
|
||||
wantErr string
|
||||
wantErrIs error
|
||||
wantOutput string
|
||||
}{
|
||||
{
|
||||
name: "slow_down_then_context_cancelled",
|
||||
timeout: 1500 * time.Millisecond,
|
||||
responses: []DeviceTokenResponse{{Error: "slow_down"}},
|
||||
wantErrIs: context.DeadlineExceeded,
|
||||
wantOutput: fmt.Sprintf(i18n.T("轮询过快,间隔增加至 %ds"), 6),
|
||||
},
|
||||
{
|
||||
name: "access_denied",
|
||||
timeout: 5 * time.Second,
|
||||
responses: []DeviceTokenResponse{{Error: "access_denied"}},
|
||||
wantErr: i18n.T("用户拒绝了授权请求"),
|
||||
wantOutput: fmt.Sprintf(i18n.T("[%d] 轮询中... (%ds)"), 1, 1),
|
||||
},
|
||||
{
|
||||
name: "expired_token",
|
||||
timeout: 5 * time.Second,
|
||||
responses: []DeviceTokenResponse{{Error: "expired_token"}},
|
||||
wantErr: i18n.T("设备授权码已过期"),
|
||||
wantOutput: fmt.Sprintf(i18n.T("[%d] 轮询中... (%ds)"), 1, 1),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
idx := int(calls.Add(1)) - 1
|
||||
if idx >= len(tt.responses) {
|
||||
idx = len(tt.responses) - 1
|
||||
}
|
||||
writeServiceResult(w, true, tt.responses[idx], "", "")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
var output bytes.Buffer
|
||||
provider.Output = &output
|
||||
provider.SetBaseURL(server.URL)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), tt.timeout)
|
||||
defer cancel()
|
||||
|
||||
_, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
})
|
||||
|
||||
if tt.wantErrIs != nil {
|
||||
if !errors.Is(err, tt.wantErrIs) {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want %v", err, tt.wantErrIs)
|
||||
}
|
||||
} else if err == nil || err.Error() != tt.wantErr {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want %q", err, tt.wantErr)
|
||||
}
|
||||
|
||||
if !strings.Contains(output.String(), tt.wantOutput) {
|
||||
t.Fatalf("expected output to contain %q, got %q", tt.wantOutput, output.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,8 +18,34 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_ID",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppKey (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CHANNEL",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "第三方渠道编码 (channelCode),如 Qoderwork",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// AuthorizeURL is the DingTalk OAuth authorization page.
|
||||
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
|
||||
@@ -58,6 +84,14 @@ const (
|
||||
// DeviceGrantType is the grant_type value defined by RFC 8628.
|
||||
DeviceGrantType = "urn:ietf:params:oauth:grant-type:device_code"
|
||||
|
||||
// Terminal API base URL for developer settings page.
|
||||
DefaultTerminalBaseURL = "https://open-dev.dingtalk.com"
|
||||
// DevicePollPath is the device flow polling path (used with MCP base URL).
|
||||
DevicePollPath = "/cli/oauth/device/poll"
|
||||
|
||||
// DeveloperSettingsPath is the path to the organization developer settings page.
|
||||
DeveloperSettingsPath = "/fe/old#/developerSettings"
|
||||
|
||||
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
|
||||
LogoutContinueURL = "https://login.dingtalk.com"
|
||||
|
||||
@@ -74,6 +108,19 @@ const (
|
||||
MCPRevokeTokenPath = "/oauth2/revokeToken"
|
||||
)
|
||||
|
||||
// GetTerminalBaseURL returns the terminal base URL with priority:
|
||||
// 1. ~/.dws/terminal_url file content (for pre-release environment)
|
||||
// 2. Default value (https://open-dev.dingtalk.com)
|
||||
func GetTerminalBaseURL() string {
|
||||
return config.GetTerminalBaseURL()
|
||||
}
|
||||
|
||||
// GetDeveloperSettingsURL returns the full URL to the organization developer
|
||||
// settings page, derived from the terminal base URL.
|
||||
func GetDeveloperSettingsURL() string {
|
||||
return config.GetDeveloperSettingsURL()
|
||||
}
|
||||
|
||||
// GetMCPBaseURL returns the MCP base URL with priority:
|
||||
// 1. ~/.dws/mcp_url file content (for pre-release environment)
|
||||
// 2. Default value (https://mcp.dingtalk.com)
|
||||
@@ -110,7 +157,7 @@ func SetClientIDFromMCP(id string) {
|
||||
func IsClientIDFromMCP() bool {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return clientIDFromMCP
|
||||
return clientIDFromMCP || edition.Get().AuthClientFromMCP
|
||||
}
|
||||
|
||||
// GetUserAccessTokenURL returns the appropriate token exchange URL.
|
||||
@@ -189,6 +236,9 @@ func ClientID() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
if id := edition.Get().AuthClientID; id != "" {
|
||||
return id
|
||||
}
|
||||
// Try loading from persisted app config
|
||||
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
|
||||
return id
|
||||
|
||||
@@ -21,6 +21,8 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
@@ -59,6 +61,19 @@ func (p *OAuthProvider) exchangeCode(ctx context.Context, code string) (*TokenDa
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// ExchangeCodeForToken exchanges an authorization code for token data using
|
||||
// the currently configured client credentials. This is a convenience wrapper
|
||||
// around OAuthProvider.exchangeCode for callers outside the auth package.
|
||||
func ExchangeCodeForToken(ctx context.Context, configDir, code string) (*TokenData, error) {
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
clientID: ClientID(),
|
||||
Output: io.Discard,
|
||||
httpClient: oauthHTTPClient,
|
||||
}
|
||||
return p.exchangeCode(ctx, code)
|
||||
}
|
||||
|
||||
// exchangeCodeViaMCP exchanges auth code for token via MCP proxy.
|
||||
// This is used when client secret is not available (server-side secret management).
|
||||
func (p *OAuthProvider) exchangeCodeViaMCP(ctx context.Context, code string) (*TokenData, error) {
|
||||
@@ -919,14 +934,254 @@ const notEnabledHTML = `<!doctype html>
|
||||
</body>
|
||||
</html>`
|
||||
|
||||
const accessDeniedHTML = `<!doctype html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1" />
|
||||
<title>钉钉 CLI</title>
|
||||
<style>
|
||||
body {
|
||||
font-family:
|
||||
-apple-system, BlinkMacSystemFont, "Segoe UI", Roboto,
|
||||
"Helvetica Neue", Arial, sans-serif;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
min-height: 100vh;
|
||||
margin: 0;
|
||||
background: #f5f5f5;
|
||||
padding: 20px;
|
||||
}
|
||||
.card {
|
||||
height: 600px;
|
||||
width: 480px;
|
||||
border-radius: 16px;
|
||||
background: #ffffff;
|
||||
box-sizing: border-box;
|
||||
border: 1px solid #f2f2f6;
|
||||
box-shadow: 0px 2px 4px 0px rgba(0, 0, 0, 0.12);
|
||||
padding: 32px 24px 24px;
|
||||
text-align: center;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
flex-direction: column;
|
||||
}
|
||||
.lock-icon {
|
||||
width: 120px;
|
||||
height: 120px;
|
||||
margin: 0 auto;
|
||||
object-fit: contain;
|
||||
display: block;
|
||||
}
|
||||
h1 {
|
||||
margin: 8px 0 0;
|
||||
font-family:
|
||||
"PingFang SC",
|
||||
-apple-system,
|
||||
BlinkMacSystemFont,
|
||||
"Segoe UI",
|
||||
Roboto,
|
||||
"Helvetica Neue",
|
||||
Arial,
|
||||
sans-serif;
|
||||
font-size: 18px;
|
||||
font-weight: 600;
|
||||
line-height: 44px;
|
||||
text-align: center;
|
||||
letter-spacing: normal;
|
||||
color: #181c1f;
|
||||
}
|
||||
p {
|
||||
margin: 0;
|
||||
font-family:
|
||||
"PingFang SC",
|
||||
-apple-system,
|
||||
BlinkMacSystemFont,
|
||||
"Segoe UI",
|
||||
Roboto,
|
||||
"Helvetica Neue",
|
||||
Arial,
|
||||
sans-serif;
|
||||
font-size: 14px;
|
||||
font-weight: normal;
|
||||
line-height: 21px;
|
||||
text-align: center;
|
||||
letter-spacing: normal;
|
||||
color: rgba(24, 28, 31, 0.6);
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="card">
|
||||
<img
|
||||
class="lock-icon"
|
||||
src="https://img.alicdn.com/imgextra/i4/O1CN01fS3xxz1vbzZSGjbe0_!!6000000006192-2-tps-480-480.png"
|
||||
alt="lock icon"
|
||||
/>
|
||||
<h1>无权限访问</h1>
|
||||
<p>您不在该组织的 CLI 授权人员范围内。请联系组织管理员将您加入授权名单。此页面可以关闭。</p>
|
||||
</div>
|
||||
</body>
|
||||
</html>`
|
||||
|
||||
const channelDeniedHTML = `<!doctype html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="utf-8" />
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1" />
|
||||
<title>钉钉 CLI</title>
|
||||
<style>
|
||||
body {
|
||||
font-family:
|
||||
-apple-system, BlinkMacSystemFont, "Segoe UI", Roboto,
|
||||
"Helvetica Neue", Arial, sans-serif;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
min-height: 100vh;
|
||||
margin: 0;
|
||||
background: #f5f5f5;
|
||||
padding: 20px;
|
||||
}
|
||||
.card {
|
||||
height: 600px;
|
||||
width: 480px;
|
||||
border-radius: 16px;
|
||||
background: #ffffff;
|
||||
box-sizing: border-box;
|
||||
border: 1px solid #f2f2f6;
|
||||
box-shadow: 0px 2px 4px 0px rgba(0, 0, 0, 0.12);
|
||||
padding: 32px 24px 24px;
|
||||
text-align: center;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
flex-direction: column;
|
||||
}
|
||||
.lock-icon {
|
||||
width: 120px;
|
||||
height: 120px;
|
||||
margin: 0 auto;
|
||||
object-fit: contain;
|
||||
display: block;
|
||||
}
|
||||
h1 {
|
||||
margin: 8px 0 0;
|
||||
font-family:
|
||||
"PingFang SC",
|
||||
-apple-system,
|
||||
BlinkMacSystemFont,
|
||||
"Segoe UI",
|
||||
Roboto,
|
||||
"Helvetica Neue",
|
||||
Arial,
|
||||
sans-serif;
|
||||
font-size: 18px;
|
||||
font-weight: 600;
|
||||
line-height: 44px;
|
||||
text-align: center;
|
||||
letter-spacing: normal;
|
||||
color: #181c1f;
|
||||
}
|
||||
p {
|
||||
margin: 0;
|
||||
font-family:
|
||||
"PingFang SC",
|
||||
-apple-system,
|
||||
BlinkMacSystemFont,
|
||||
"Segoe UI",
|
||||
Roboto,
|
||||
"Helvetica Neue",
|
||||
Arial,
|
||||
sans-serif;
|
||||
font-size: 14px;
|
||||
font-weight: normal;
|
||||
line-height: 21px;
|
||||
text-align: center;
|
||||
letter-spacing: normal;
|
||||
color: rgba(24, 28, 31, 0.6);
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="card">
|
||||
<img
|
||||
class="lock-icon"
|
||||
src="https://img.alicdn.com/imgextra/i4/O1CN01fS3xxz1vbzZSGjbe0_!!6000000006192-2-tps-480-480.png"
|
||||
alt="lock icon"
|
||||
/>
|
||||
<h1>渠道未授权</h1>
|
||||
<p>当前渠道未获得该组织授权,或组织已开启渠道管控。请联系组织管理员开通渠道访问权限,或升级到最新版本的 CLI。此页面可以关闭。</p>
|
||||
</div>
|
||||
</body>
|
||||
</html>`
|
||||
|
||||
// CLIAuthStatus represents the response from /cli/cliAuthEnabled API.
|
||||
type CLIAuthStatus struct {
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
} `json:"result"`
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result *CLIAuthResult `json:"result"`
|
||||
}
|
||||
|
||||
// CLIAuthResult holds the business data returned by /cli/cliAuthEnabled.
|
||||
// The server computes cliAuthEnabled by considering the org switch, userScope,
|
||||
// and channelScope together; the CLI uses it as-is.
|
||||
type CLIAuthResult struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
UserScope string `json:"userScope,omitempty"` // "all" | "specified" | "forbidden"
|
||||
AllowedUsers []string `json:"allowedUsers,omitempty"` // staffId list when userScope="specified"
|
||||
ChannelScope string `json:"channelScope,omitempty"` // "all" | "specified"
|
||||
AllowedChannels []string `json:"allowedChannels,omitempty"` // channelCode list when channelScope="specified"
|
||||
ChannelConfigEnabled bool `json:"channelConfigEnabled,omitempty"` // whether org has any channel restriction configured
|
||||
}
|
||||
|
||||
// classifyDenialReason inspects a CLIAuthStatus response and returns a machine-readable
|
||||
// denial reason string. Returns "" when access is granted.
|
||||
//
|
||||
// Priority rationale:
|
||||
// 1. Explicit org-wide ban (userScope=forbidden) always wins.
|
||||
// 2. Channel scope is evaluated BEFORE user scope because the CLI has
|
||||
// authoritative knowledge of DWS_CHANNEL and can verify membership against
|
||||
// allowedChannels. This avoids falsely blaming the user when the real
|
||||
// denial cause is a channel mismatch (e.g. user is in allowedUsers but the
|
||||
// current channel is not in allowedChannels).
|
||||
// 3. Only when the channel is unrestricted or matches do we attribute the
|
||||
// denial to the user scope.
|
||||
func classifyDenialReason(status *CLIAuthStatus, currentChannel string) string {
|
||||
if status.ErrorCode == "CHANNEL_REQUIRED" {
|
||||
return "channel_required"
|
||||
}
|
||||
if status.ErrorCode == "NO_AUTH" {
|
||||
return "no_auth"
|
||||
}
|
||||
if status.Result == nil || !status.Success {
|
||||
return "unknown"
|
||||
}
|
||||
r := status.Result
|
||||
if r.CLIAuthEnabled {
|
||||
return ""
|
||||
}
|
||||
|
||||
if r.UserScope == "forbidden" {
|
||||
return "user_forbidden"
|
||||
}
|
||||
|
||||
if r.ChannelScope == "specified" {
|
||||
if currentChannel == "" {
|
||||
return "channel_required"
|
||||
}
|
||||
if !slices.Contains(r.AllowedChannels, currentChannel) {
|
||||
return "channel_not_allowed"
|
||||
}
|
||||
}
|
||||
|
||||
if r.UserScope == "specified" {
|
||||
return "user_not_allowed"
|
||||
}
|
||||
return "cli_not_enabled"
|
||||
}
|
||||
|
||||
// SuperAdmin represents a corp super admin.
|
||||
@@ -985,6 +1240,9 @@ func (p *OAuthProvider) doCheckCLIAuthEnabled(ctx context.Context, accessToken s
|
||||
return nil, fmt.Errorf("creating request: %w", err)
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
if ch := os.Getenv("DWS_CHANNEL"); ch != "" {
|
||||
req.Header.Set("x-dws-channel", ch)
|
||||
}
|
||||
|
||||
client := p.httpClient
|
||||
if client == nil {
|
||||
|
||||
@@ -124,6 +124,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
token *TokenData
|
||||
err error
|
||||
cliAuthDisabled bool
|
||||
denialReason string
|
||||
}
|
||||
resultCh := make(chan callbackResult, 1)
|
||||
errCh := make(chan error, 1)
|
||||
@@ -236,19 +237,30 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
|
||||
// Check CLI auth enabled status (fail-closed: treat errors as disabled)
|
||||
authStatus, statusErr := p.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
var denialReason string
|
||||
if statusErr != nil {
|
||||
denialReason = "unknown"
|
||||
} else {
|
||||
denialReason = classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
|
||||
}
|
||||
cliAuthEnabled := denialReason == ""
|
||||
|
||||
// Update CLI auth disabled state
|
||||
callbackTokenMu.Lock()
|
||||
callbackAuthDisabled = !cliAuthEnabled
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Display appropriate HTML based on CLI auth status
|
||||
// Display appropriate HTML based on auth status and denial reason
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if !cliAuthEnabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
switch {
|
||||
case cliAuthEnabled:
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
case denialReason == "user_forbidden" || denialReason == "user_not_allowed":
|
||||
_, _ = fmt.Fprint(w, accessDeniedHTML)
|
||||
case denialReason == "channel_not_allowed" || denialReason == "channel_required":
|
||||
_, _ = fmt.Fprint(w, channelDeniedHTML)
|
||||
default:
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
}
|
||||
// Ensure response is flushed to client
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
@@ -256,7 +268,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
}
|
||||
// Notify main goroutine with full result
|
||||
select {
|
||||
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled}:
|
||||
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled, denialReason: denialReason}:
|
||||
default:
|
||||
}
|
||||
})
|
||||
@@ -395,8 +407,18 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), result.err)
|
||||
}
|
||||
|
||||
// Handle CLI auth disabled - keep server running for user to apply
|
||||
// Handle CLI auth disabled - for terminal denial reasons, exit immediately
|
||||
// (page shows accessDeniedHTML/channelDeniedHTML with no apply button,
|
||||
// so polling for apply submission would hang forever).
|
||||
// Error messages are kept consistent with the text shown on the HTML pages.
|
||||
if result.cliAuthDisabled {
|
||||
switch result.denialReason {
|
||||
case "user_forbidden", "user_not_allowed":
|
||||
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员将您加入授权名单"))
|
||||
case "channel_not_allowed", "channel_required":
|
||||
return nil, errors.New(i18n.T("当前渠道未获得该组织授权,或组织已开启渠道管控,请联系组织管理员开通渠道访问权限,或升级到最新版本的 CLI"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
|
||||
|
||||
@@ -435,7 +457,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
// Check if CLI auth is now enabled (admin approved)
|
||||
if currentToken != nil {
|
||||
authStatus, err := p.CheckCLIAuthEnabled(ctx, currentToken.AccessToken)
|
||||
if err == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled {
|
||||
if err == nil && classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL")) == "" {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
|
||||
time.Sleep(2 * time.Second)
|
||||
result.token = currentToken
|
||||
@@ -463,6 +485,13 @@ continueLogin:
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Always persist clientId to app.json so future process startups
|
||||
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
|
||||
if p.clientID != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
|
||||
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
p.persistAppConfigIfNeeded()
|
||||
|
||||
|
||||
+65
-17
@@ -20,7 +20,11 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// TokenData holds the OAuth token set persisted to disk.
|
||||
@@ -61,44 +65,88 @@ func (t *TokenData) HasPersistentCode() bool {
|
||||
return t != nil && t.PersistentCode != ""
|
||||
}
|
||||
|
||||
// SaveTokenData saves TokenData to the platform keychain.
|
||||
// Uses the new keychain-based storage with random master key for better security.
|
||||
const tokenJSONFile = "token.json"
|
||||
|
||||
// TokenMarker is a lightweight file the host application reads to detect
|
||||
// whether the CLI has a valid token without accessing the keychain.
|
||||
type TokenMarker struct {
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// WriteTokenMarker writes a token.json marker containing only an updated_at
|
||||
// timestamp. The host application uses this file's presence and mtime to
|
||||
// decide whether it needs to trigger a new auth exchange.
|
||||
func WriteTokenMarker(configDir string) error {
|
||||
marker := TokenMarker{UpdatedAt: time.Now().Format(time.RFC3339)}
|
||||
data, _ := json.MarshalIndent(marker, "", " ")
|
||||
if err := os.MkdirAll(configDir, 0o700); err != nil {
|
||||
return err
|
||||
}
|
||||
tmp := filepath.Join(configDir, tokenJSONFile+".tmp")
|
||||
if err := os.WriteFile(tmp, data, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmp, filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// DeleteTokenMarker removes the token.json marker file.
|
||||
func DeleteTokenMarker(configDir string) error {
|
||||
return os.Remove(filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to the default keychain-based storage.
|
||||
func SaveTokenData(configDir string, data *TokenData) error {
|
||||
if h := edition.Get(); h.SaveToken != nil {
|
||||
jsonData, err := json.MarshalIndent(data, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling token data for hook: %w", err)
|
||||
}
|
||||
return h.SaveToken(configDir, jsonData)
|
||||
}
|
||||
return SaveTokenDataKeychain(data)
|
||||
}
|
||||
|
||||
// LoadTokenData reads TokenData from the platform keychain.
|
||||
// On first call, it attempts to migrate legacy .data file if present.
|
||||
// LoadTokenData reads TokenData. When an edition hook (LoadToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain with legacy .data migration.
|
||||
func LoadTokenData(configDir string) (*TokenData, error) {
|
||||
// Try loading from new keychain first
|
||||
if h := edition.Get(); h.LoadToken != nil {
|
||||
jsonData, err := h.LoadToken(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var td TokenData
|
||||
if err := json.Unmarshal(jsonData, &td); err != nil {
|
||||
return nil, fmt.Errorf("parsing token data from hook: %w", err)
|
||||
}
|
||||
return &td, nil
|
||||
}
|
||||
|
||||
// Default: keychain with legacy .data migration
|
||||
if TokenDataExistsKeychain() {
|
||||
return LoadTokenDataKeychain()
|
||||
}
|
||||
|
||||
// Fallback: try legacy .data file and migrate
|
||||
data, err := LoadSecureTokenData(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Migrate to keychain for future use
|
||||
if err := SaveTokenDataKeychain(data); err == nil {
|
||||
// Successfully migrated, delete legacy file
|
||||
_ = DeleteSecureData(configDir)
|
||||
}
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes token data from both keychain and legacy storage.
|
||||
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain + legacy cleanup.
|
||||
func DeleteTokenData(configDir string) error {
|
||||
// Delete from keychain
|
||||
if h := edition.Get(); h.DeleteToken != nil {
|
||||
return h.DeleteToken(configDir)
|
||||
}
|
||||
keychainErr := DeleteTokenDataKeychain()
|
||||
|
||||
// Also clean up any legacy .data file
|
||||
legacyErr := DeleteSecureData(configDir)
|
||||
|
||||
// Return keychain error if any, otherwise legacy error
|
||||
if keychainErr != nil {
|
||||
return keychainErr
|
||||
}
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -121,6 +122,25 @@ func NewSchemaCommand(loader CatalogLoader) *cobra.Command {
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
catalog, err := loader.Load(cmd.Context())
|
||||
if err != nil {
|
||||
var degraded *CatalogDegraded
|
||||
if errors.As(err, °raded) {
|
||||
fmt.Fprintf(cmd.ErrOrStderr(), "hint: %s\n", degraded.Hint)
|
||||
payload := map[string]any{
|
||||
"kind": "schema",
|
||||
"count": 0,
|
||||
"products": []any{},
|
||||
"degraded": true,
|
||||
"reason": string(degraded.Reason),
|
||||
"hint": degraded.Hint,
|
||||
}
|
||||
return output.WriteFiltered(
|
||||
cmd.OutOrStdout(),
|
||||
output.ResolveFormat(cmd, output.FormatJSON),
|
||||
payload,
|
||||
output.ResolveFields(cmd),
|
||||
output.ResolveJQ(cmd),
|
||||
)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -213,6 +233,27 @@ func newProductCommand(product ir.CanonicalProduct, runner executor.Runner, engi
|
||||
for _, tool := range product.Tools {
|
||||
cmd.AddCommand(newToolCommand(product, tool, runner, engine))
|
||||
}
|
||||
|
||||
// Register phase: notify the pipeline that a product and its
|
||||
// tools have been added to the command tree. This runs once at
|
||||
// startup (not per-request) and enables handlers to inspect or
|
||||
// enrich the registered command surface.
|
||||
if engine != nil && engine.HasHandlers(pipeline.Register) {
|
||||
pctx := &pipeline.Context{
|
||||
Command: product.ID,
|
||||
}
|
||||
// Best-effort — registration errors are logged but do not
|
||||
// prevent the CLI from starting.
|
||||
if pipeErr := engine.RunPhase(pipeline.Register, pctx); pipeErr != nil {
|
||||
slog.Debug("pipeline register phase", "product", product.ID, "error", pipeErr)
|
||||
} else {
|
||||
slog.Debug("pipeline register",
|
||||
"product", product.ID,
|
||||
"tool_count", len(product.Tools),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -368,6 +409,16 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
params = pctx.Params
|
||||
for _, c := range pctx.Corrections {
|
||||
slog.Debug("pipeline correction",
|
||||
"phase", "post-parse",
|
||||
"handler", c.Handler,
|
||||
"kind", c.Kind,
|
||||
"field", c.Field,
|
||||
"original", c.Original,
|
||||
"corrected", c.Corrected,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if err := ValidateInputSchema(params, tool.InputSchema); err != nil {
|
||||
@@ -392,6 +443,10 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
params = pctx.Params
|
||||
slog.Debug("pipeline pre-request",
|
||||
"command", tool.CanonicalPath,
|
||||
"param_count", len(params),
|
||||
)
|
||||
}
|
||||
|
||||
invocation := executor.NewInvocation(product, tool, params)
|
||||
@@ -414,6 +469,10 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
result.Response = pctx.Response
|
||||
slog.Debug("pipeline post-response",
|
||||
"command", tool.CanonicalPath,
|
||||
"has_response", result.Response != nil,
|
||||
)
|
||||
}
|
||||
|
||||
if warning := lifecycleWarning(product); warning != "" {
|
||||
|
||||
@@ -1033,6 +1033,83 @@ func newTestMCPCommand(t *testing.T, catalog ir.Catalog, runner executor.Runner)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func TestSchemaCommandOutputsDegradedOnUnauthenticated(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
degradedErr := &CatalogDegraded{
|
||||
Reason: DegradedUnauthenticated,
|
||||
Hint: "未登录,无法发现 MCP 服务。请先执行: dws auth login",
|
||||
}
|
||||
cmd := NewSchemaCommand(errorLoader{err: degradedErr})
|
||||
|
||||
var out, errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v, want nil (degraded handled gracefully)", err)
|
||||
}
|
||||
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if payload["degraded"] != true {
|
||||
t.Fatalf("payload[degraded] = %v, want true", payload["degraded"])
|
||||
}
|
||||
if payload["reason"] != "unauthenticated" {
|
||||
t.Fatalf("payload[reason] = %v, want unauthenticated", payload["reason"])
|
||||
}
|
||||
if payload["count"] != float64(0) {
|
||||
t.Fatalf("payload[count] = %v, want 0", payload["count"])
|
||||
}
|
||||
if !strings.Contains(errOut.String(), "hint:") {
|
||||
t.Fatalf("stderr = %q, want hint message", errOut.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaCommandOutputsDegradedOnMarketUnreachable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
degradedErr := &CatalogDegraded{
|
||||
Reason: DegradedMarketUnreachable,
|
||||
Hint: "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络",
|
||||
}
|
||||
cmd := NewSchemaCommand(errorLoader{err: degradedErr})
|
||||
|
||||
var out, errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if payload["reason"] != "market_unreachable" {
|
||||
t.Fatalf("payload[reason] = %v, want market_unreachable", payload["reason"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaCommandPropagatesNonDegradedError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
wantErr := errors.New("unexpected failure")
|
||||
cmd := NewSchemaCommand(errorLoader{err: wantErr})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("Execute() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
type errorLoader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
+93
-8
@@ -28,15 +28,96 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CACHE_DIR",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "覆盖缓存目录",
|
||||
DefaultValue: "~/.dws/cache",
|
||||
Example: "/tmp/dws-cache",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CATALOG_FIXTURE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "使用本地 JSON 文件替代在线目录发现",
|
||||
Example: "/path/to/catalog.json",
|
||||
Hidden: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PLUGIN_COLD_TIMEOUT",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "插件 MCP 冷启动发现的超时时长(Go duration 格式,如 2s / 1500ms)。设置后同时覆盖 HTTP 与 stdio 插件的冷启动预算;未设置时使用内置默认值(HTTP 无鉴权 1s / 有鉴权 1.5s / stdio 2s)。",
|
||||
DefaultValue: "",
|
||||
Example: "3s",
|
||||
})
|
||||
}
|
||||
|
||||
// CatalogDegradedReason identifies why catalog discovery returned empty.
|
||||
type CatalogDegradedReason string
|
||||
|
||||
const (
|
||||
DegradedUnauthenticated CatalogDegradedReason = "unauthenticated"
|
||||
DegradedMarketUnreachable CatalogDegradedReason = "market_unreachable"
|
||||
DegradedRuntimeAllFailed CatalogDegradedReason = "runtime_all_failed"
|
||||
)
|
||||
|
||||
// CatalogDegraded is returned by EnvironmentLoader.Load when discovery
|
||||
// fails for a diagnosable reason. Callers that need graceful degradation
|
||||
// (e.g. the runtime runner) can check errors.As and fall back to an
|
||||
// empty catalog; callers like the schema command can surface the hint.
|
||||
type CatalogDegraded struct {
|
||||
Reason CatalogDegradedReason
|
||||
Hint string
|
||||
ServerCount int // number of servers discovered (only set for runtime_all_failed)
|
||||
}
|
||||
|
||||
func (e *CatalogDegraded) Error() string { return string(e.Reason) + ": " + e.Hint }
|
||||
|
||||
func degradedHint(reason CatalogDegradedReason, serverCount int) string {
|
||||
embedded := edition.Get().IsEmbedded
|
||||
switch reason {
|
||||
case DegradedUnauthenticated:
|
||||
if embedded {
|
||||
return "未登录,请重新认证"
|
||||
}
|
||||
return "未登录,无法发现 MCP 服务。请先执行: dws auth login"
|
||||
case DegradedMarketUnreachable:
|
||||
if embedded {
|
||||
return "无法连接 MCP 市场,请检查网络"
|
||||
}
|
||||
return "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络"
|
||||
case DegradedRuntimeAllFailed:
|
||||
if embedded {
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试", serverCount)
|
||||
}
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试或执行: dws cache refresh", serverCount)
|
||||
default:
|
||||
return "MCP 服务发现失败"
|
||||
}
|
||||
}
|
||||
|
||||
func newCatalogDegraded(reason CatalogDegradedReason, serverCount int) *CatalogDegraded {
|
||||
return &CatalogDegraded{
|
||||
Reason: reason,
|
||||
Hint: degradedHint(reason, serverCount),
|
||||
ServerCount: serverCount,
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
CatalogFixtureEnv = "DWS_CATALOG_FIXTURE"
|
||||
CacheDirEnv = "DWS_CACHE_DIR"
|
||||
PluginColdTimeoutEnv = "DWS_PLUGIN_COLD_TIMEOUT"
|
||||
DefaultMarketBaseURL = "https://mcp.dingtalk.com"
|
||||
|
||||
// defaultDiscoveryTimeout bounds the time spent on live registry discovery.
|
||||
defaultDiscoveryTimeout = 10 * time.Second
|
||||
// Tightened to 4s so a slow/unreachable discovery endpoint cannot block
|
||||
// every CLI command invocation. See issue #119.
|
||||
defaultDiscoveryTimeout = 4 * time.Second
|
||||
)
|
||||
|
||||
type CatalogLoader interface {
|
||||
@@ -128,17 +209,23 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
// Startup command construction should not block on synchronous discovery
|
||||
// just because the cache has aged past the short revalidation window.
|
||||
cached := l.loadFromCache(store)
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
|
||||
transportClient := transport.NewClient(nil)
|
||||
hasAuth := false
|
||||
if l.AuthTokenFunc != nil {
|
||||
if token := l.AuthTokenFunc(ctx); token != "" {
|
||||
transportClient = transportClient.WithAuth(token, nil)
|
||||
hasAuth = true
|
||||
}
|
||||
}
|
||||
|
||||
if !hasAuth {
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedUnauthenticated, 0)
|
||||
}
|
||||
|
||||
// Use a bounded context so discovery doesn't hang in test or CI environments.
|
||||
timeout := defaultDiscoveryTimeout
|
||||
if l.DiscoveryTimeout > 0 {
|
||||
@@ -157,12 +244,10 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
}
|
||||
response, err := service.MarketClient.FetchServers(discoverCtx, 200)
|
||||
if err != nil {
|
||||
// Graceful degradation: return empty catalog on discovery failure.
|
||||
// The runtime runner will fall back to EchoRunner for unknown products.
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedMarketUnreachable, 0)
|
||||
}
|
||||
|
||||
servers := market.NormalizeServers(response, "live_market")
|
||||
@@ -192,10 +277,10 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
|
||||
refreshed, failures := service.DiscoverAllRuntime(discoverCtx, toRefresh)
|
||||
if len(unchangedRuntime) == 0 && len(refreshed) == 0 && len(failures) > 0 {
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedRuntimeAllFailed, len(servers))
|
||||
}
|
||||
|
||||
refreshedByKey := make(map[string]discovery.RuntimeServer, len(refreshed))
|
||||
|
||||
@@ -95,9 +95,13 @@ func BuildDynamicCommands(servers []market.ServerDescriptor, runner executor.Run
|
||||
|
||||
bindings, normalizer := buildOverrideBindings(override)
|
||||
|
||||
// Resolve Short/Long from Detail API toolTitle/toolDesc; fallback to generic.
|
||||
// Resolve Short/Long from Detail API toolTitle/toolDesc;
|
||||
// fallback to overlay description; then to generic cmdName/cliName.
|
||||
short := fmt.Sprintf("%s/%s", cmdName, cliName)
|
||||
long := ""
|
||||
if desc := strings.TrimSpace(override.Description); desc != "" {
|
||||
short = desc
|
||||
}
|
||||
if dt, ok := detailIndex[toolName]; ok {
|
||||
if title := strings.TrimSpace(dt.ToolTitle); title != "" {
|
||||
short = title
|
||||
|
||||
@@ -31,6 +31,7 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/convert"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
)
|
||||
|
||||
type ValueKind string
|
||||
@@ -139,6 +140,11 @@ func NewDirectCommand(route Route, runner executor.Runner) *cobra.Command {
|
||||
for key, value := range bindingParams {
|
||||
params[key] = value
|
||||
}
|
||||
|
||||
// Collect schema-derived flags (from buildFlagsFromDetailSchema)
|
||||
// that are not covered by explicit bindings.
|
||||
collectSchemaFlags(cmd, route.Bindings, params)
|
||||
|
||||
if route.Normalizer != nil {
|
||||
if err := route.Normalizer(cmd, params); err != nil {
|
||||
return err
|
||||
@@ -246,6 +252,70 @@ func ApplyBindings(cmd *cobra.Command, bindings []FlagBinding) {
|
||||
_ = cmd.Flags().MarkHidden("params")
|
||||
}
|
||||
|
||||
// collectSchemaFlags picks up flags created by buildFlagsFromDetailSchema that
|
||||
// have no explicit FlagBinding. This bridges the gap for plugin-defined tools
|
||||
// whose parameters come from the MCP inputSchema rather than CLIToolOverride.Flags.
|
||||
func collectSchemaFlags(cmd *cobra.Command, bindings []FlagBinding, params map[string]any) {
|
||||
// Build a set of flag names already covered by bindings.
|
||||
bound := make(map[string]bool, len(bindings)*2)
|
||||
for _, b := range bindings {
|
||||
if n := strings.TrimSpace(b.FlagName); n != "" {
|
||||
bound[n] = true
|
||||
}
|
||||
if a := strings.TrimSpace(b.Alias); a != "" {
|
||||
bound[a] = true
|
||||
}
|
||||
}
|
||||
|
||||
// Reserved/internal flags that should never be forwarded as tool params.
|
||||
skip := map[string]bool{
|
||||
"json": true, "params": true, "help": true,
|
||||
"format": true, "fields": true, "jq": true,
|
||||
"debug": true, "verbose": true, "dry-run": true,
|
||||
"yes": true, "mock": true, "timeout": true,
|
||||
"client-id": true, "client-secret": true,
|
||||
}
|
||||
|
||||
cmd.Flags().Visit(func(f *pflag.Flag) {
|
||||
if bound[f.Name] || skip[f.Name] {
|
||||
return
|
||||
}
|
||||
// Convert flag name back to the original parameter name (kebab → snake/camel)
|
||||
// For simplicity, use the flag name as-is since MCP tools typically
|
||||
// use snake_case which maps to kebab-case flags.
|
||||
paramName := toOriginalParamName(f.Name)
|
||||
if _, exists := params[paramName]; exists {
|
||||
return // already set by --json/--params
|
||||
}
|
||||
|
||||
switch f.Value.Type() {
|
||||
case "int":
|
||||
if v, err := cmd.Flags().GetInt(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
case "bool":
|
||||
if v, err := cmd.Flags().GetBool(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
case "stringSlice":
|
||||
if v, err := cmd.Flags().GetStringSlice(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
default:
|
||||
if v, err := cmd.Flags().GetString(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// toOriginalParamName converts a kebab-case flag name back to the original
|
||||
// MCP parameter name. Since toKebabCase converts both camelCase and snake_case
|
||||
// to kebab-case, we default to snake_case (the MCP convention).
|
||||
func toOriginalParamName(flagName string) string {
|
||||
return strings.ReplaceAll(flagName, "-", "_")
|
||||
}
|
||||
|
||||
func CollectBindings(cmd *cobra.Command, bindings []FlagBinding, existing map[string]any) (map[string]any, error) {
|
||||
if existing == nil {
|
||||
existing = map[string]any{}
|
||||
|
||||
@@ -146,3 +146,102 @@ func TestCollectBindingsParsesJSONFlagValue(t *testing.T) {
|
||||
t.Fatalf("config.options = %#v, want array of 1", config["options"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsPicksUpUnboundFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Simulate a plugin command with schema-generated flags but no bindings.
|
||||
cmd := &cobra.Command{Use: "greet"}
|
||||
cmd.Flags().String("name", "", "Name of person")
|
||||
cmd.Flags().String("language", "en", "Language")
|
||||
cmd.Flags().Int("count", 0, "Repeat count")
|
||||
cmd.Flags().Bool("loud", false, "Loud mode")
|
||||
cmd.Flags().String("json", "", "")
|
||||
cmd.Flags().String("params", "", "")
|
||||
|
||||
// User sets --name and --count but not --language
|
||||
_ = cmd.Flags().Set("name", "Alice")
|
||||
_ = cmd.Flags().Set("count", "3")
|
||||
_ = cmd.Flags().Set("loud", "true")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, nil, params)
|
||||
|
||||
if params["name"] != "Alice" {
|
||||
t.Errorf("name = %v, want Alice", params["name"])
|
||||
}
|
||||
if params["count"] != 3 {
|
||||
t.Errorf("count = %v, want 3", params["count"])
|
||||
}
|
||||
if params["loud"] != true {
|
||||
t.Errorf("loud = %v, want true", params["loud"])
|
||||
}
|
||||
// language was not set by user, should not appear
|
||||
if _, exists := params["language"]; exists {
|
||||
t.Errorf("language should not be in params (not set by user)")
|
||||
}
|
||||
// json/params are reserved, should not appear
|
||||
if _, exists := params["json"]; exists {
|
||||
t.Error("json should be skipped")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsSkipsBoundFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
bindings := []FlagBinding{
|
||||
{FlagName: "dept-id", Property: "deptId", Kind: ValueString},
|
||||
}
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
ApplyBindings(cmd, bindings)
|
||||
// Also add a schema-generated flag
|
||||
cmd.Flags().String("title", "", "Title")
|
||||
|
||||
_ = cmd.Flags().Set("dept-id", "D001")
|
||||
_ = cmd.Flags().Set("title", "Hello")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, bindings, params)
|
||||
|
||||
// dept-id is bound, should NOT be collected by collectSchemaFlags
|
||||
if _, exists := params["dept_id"]; exists {
|
||||
t.Error("dept-id should be skipped (already has binding)")
|
||||
}
|
||||
// title is unbound, should be collected
|
||||
if params["title"] != "Hello" {
|
||||
t.Errorf("title = %v, want Hello", params["title"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsSkipsGlobalFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("name", "", "Name")
|
||||
cmd.Flags().Bool("debug", false, "Debug")
|
||||
cmd.Flags().Bool("verbose", false, "Verbose")
|
||||
cmd.Flags().Bool("dry-run", false, "Dry run")
|
||||
cmd.Flags().String("format", "json", "Format")
|
||||
cmd.Flags().String("json", "", "")
|
||||
cmd.Flags().String("params", "", "")
|
||||
|
||||
_ = cmd.Flags().Set("name", "Bob")
|
||||
_ = cmd.Flags().Set("debug", "true")
|
||||
_ = cmd.Flags().Set("verbose", "true")
|
||||
_ = cmd.Flags().Set("dry-run", "true")
|
||||
_ = cmd.Flags().Set("format", "table")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, nil, params)
|
||||
|
||||
if params["name"] != "Bob" {
|
||||
t.Errorf("name = %v, want Bob", params["name"])
|
||||
}
|
||||
// Global flags should be skipped
|
||||
for _, skip := range []string{"debug", "verbose", "dry_run", "format"} {
|
||||
if _, exists := params[skip]; exists {
|
||||
t.Errorf("%s should be skipped (global flag)", skip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -27,8 +27,24 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_TENANT",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "缓存分区的租户标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_AUTH_IDENTITY",
|
||||
Category: configmeta.CategorySecurity,
|
||||
Description: "缓存分区的认证身份标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
tenantEnv = "DWS_TENANT"
|
||||
authIdentityEnv = "DWS_AUTH_IDENTITY"
|
||||
@@ -43,6 +59,11 @@ type Service struct {
|
||||
Tenant string
|
||||
AuthIdentity string
|
||||
Logger *slog.Logger
|
||||
// PerServerTimeout overrides the default per-server discovery timeout
|
||||
// when greater than zero. Useful for tests and for callers that need a
|
||||
// tighter or looser bound. When zero, defaultPerServerDiscoveryTimeout
|
||||
// applies.
|
||||
PerServerTimeout time.Duration
|
||||
}
|
||||
|
||||
type RuntimeServer struct {
|
||||
@@ -154,7 +175,11 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
|
||||
}, nil
|
||||
}
|
||||
|
||||
const perServerDiscoveryTimeout = 5 * time.Second
|
||||
// defaultPerServerDiscoveryTimeout bounds the time spent discovering tools on
|
||||
// a single registry-listed server. Tightened to 2s so a slow/unreachable
|
||||
// server cannot stall every CLI command — a healthy MCP endpoint negotiates
|
||||
// well under a second. See issue #119.
|
||||
const defaultPerServerDiscoveryTimeout = 2 * time.Second
|
||||
|
||||
func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.ServerDescriptor) ([]RuntimeServer, []RuntimeFailure) {
|
||||
type discoveryResult struct {
|
||||
@@ -162,6 +187,11 @@ func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.Serve
|
||||
failure *RuntimeFailure
|
||||
}
|
||||
|
||||
perServerTimeout := defaultPerServerDiscoveryTimeout
|
||||
if s.PerServerTimeout > 0 {
|
||||
perServerTimeout = s.PerServerTimeout
|
||||
}
|
||||
|
||||
filtered := make([]market.ServerDescriptor, 0, len(servers))
|
||||
for _, srv := range servers {
|
||||
if !srv.CLI.Skip {
|
||||
@@ -178,7 +208,7 @@ func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.Serve
|
||||
wg.Add(1)
|
||||
go func(server market.ServerDescriptor) {
|
||||
defer wg.Done()
|
||||
serverCtx, cancel := context.WithTimeout(ctx, perServerDiscoveryTimeout)
|
||||
serverCtx, cancel := context.WithTimeout(ctx, perServerTimeout)
|
||||
defer cancel()
|
||||
start := time.Now()
|
||||
rs, err := s.DiscoverServerRuntime(serverCtx, server)
|
||||
|
||||
@@ -20,6 +20,8 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
// Category represents a stable error class with a documented exit code.
|
||||
@@ -198,12 +200,31 @@ func NewInternal(message string, opts ...Option) error {
|
||||
return newError(CategoryInternal, message, opts...)
|
||||
}
|
||||
|
||||
// ExitCoder is implemented by errors that provide their own exit code.
|
||||
// Edition-specific error types (e.g. PATError, CLIError) implement this
|
||||
// so the framework can resolve exit codes without importing edition packages.
|
||||
type ExitCoder interface {
|
||||
ExitCode() int
|
||||
}
|
||||
|
||||
// RawStderrError is implemented by errors that must output raw content
|
||||
// directly to stderr, bypassing all CLI formatting (e.g. "Error:" prefix).
|
||||
// PAT authorization errors use this to pass JSON through to the desktop runtime.
|
||||
type RawStderrError interface {
|
||||
error
|
||||
RawStderr() string
|
||||
}
|
||||
|
||||
// ExitCode maps any error to a stable exit code.
|
||||
func ExitCode(err error) int {
|
||||
var typed *Error
|
||||
if stderrors.As(err, &typed) {
|
||||
return typed.ExitCode()
|
||||
}
|
||||
var ec ExitCoder
|
||||
if stderrors.As(err, &ec) {
|
||||
return ec.ExitCode()
|
||||
}
|
||||
return 5
|
||||
}
|
||||
|
||||
@@ -257,7 +278,7 @@ func PrintJSON(w io.Writer, err error) error {
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
|
||||
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
|
||||
errorPayload["action_url"] = config.GetDeveloperSettingsURL()
|
||||
}
|
||||
}
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
@@ -323,7 +344,7 @@ func PrintHumanAt(w io.Writer, err error, v Verbosity) error {
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
|
||||
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
|
||||
lines = append(lines, "Action: 开启地址: "+config.GetDeveloperSettingsURL())
|
||||
}
|
||||
|
||||
if len(typed.Actions) > 0 {
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package errors
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type stubExitCoder struct{ code int }
|
||||
|
||||
func (s *stubExitCoder) Error() string { return "stub" }
|
||||
func (s *stubExitCoder) ExitCode() int { return s.code }
|
||||
|
||||
type stubRawStderr struct{ raw string }
|
||||
|
||||
func (s *stubRawStderr) Error() string { return s.raw }
|
||||
func (s *stubRawStderr) RawStderr() string { return s.raw }
|
||||
|
||||
func TestExitCode_ExitCoderInterface(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
err error
|
||||
want int
|
||||
}{
|
||||
{"exit code 4 via interface", &stubExitCoder{code: 4}, 4},
|
||||
{"exit code 1 via interface", &stubExitCoder{code: 1}, 1},
|
||||
{"framework Error takes precedence", NewAPI("api"), 1},
|
||||
{"plain error falls back to 5", stderrors.New("plain"), 5},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := ExitCode(tc.err); got != tc.want {
|
||||
t.Errorf("ExitCode() = %d, want %d", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExitCode_WrappedExitCoder(t *testing.T) {
|
||||
t.Parallel()
|
||||
wrapped := stderrors.Join(stderrors.New("context"), &stubExitCoder{code: 4})
|
||||
if got := ExitCode(wrapped); got != 4 {
|
||||
t.Errorf("ExitCode(wrapped) = %d, want 4", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRawStderrError_Interface(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &stubRawStderr{raw: `{"code":"PAT_LOW_RISK_NO_PERMISSION"}`}
|
||||
var raw RawStderrError
|
||||
if !stderrors.As(err, &raw) {
|
||||
t.Fatal("expected errors.As to match RawStderrError")
|
||||
}
|
||||
if !strings.Contains(raw.RawStderr(), "PAT_LOW_RISK_NO_PERMISSION") {
|
||||
t.Errorf("RawStderr() = %q, want PAT code", raw.RawStderr())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
// 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 errors
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ExitCodePermission is the process exit code for PAT authorisation failures.
|
||||
const ExitCodePermission = 4
|
||||
|
||||
// PATError represents a PAT (Personal Action Token) authorization failure
|
||||
// that should be passed through to stderr as raw JSON without any CLI-layer
|
||||
// wrapping. The host application (e.g. RewindDesktop) parses the JSON to
|
||||
// display its own authorisation UI.
|
||||
type PATError struct {
|
||||
RawJSON string
|
||||
}
|
||||
|
||||
func (e *PATError) Error() string { return e.RawJSON }
|
||||
|
||||
// ExitCode returns the documented exit code for PAT permission errors (4).
|
||||
func (e *PATError) ExitCode() int { return ExitCodePermission }
|
||||
|
||||
// RawStderr returns the raw JSON to be written directly to stderr.
|
||||
func (e *PATError) RawStderr() string { return e.RawJSON }
|
||||
|
||||
// patNoPermissionCodes are PAT error codes that should be passed through
|
||||
// as transparent PATError without CLI-level wrapping.
|
||||
var patNoPermissionCodes = map[string]bool{
|
||||
"PAT_NO_PERMISSION": true,
|
||||
"PAT_LOW_RISK_NO_PERMISSION": true,
|
||||
"PAT_MEDIUM_RISK_NO_PERMISSION": true,
|
||||
"PAT_HIGH_RISK_NO_PERMISSION": true,
|
||||
}
|
||||
|
||||
// patAuthRequiredCodes are error codes that trigger the PAT authorization
|
||||
// flow (e.g. the server auto-created a CLI app and returned auth details).
|
||||
var patAuthRequiredCodes = map[string]bool{
|
||||
"AGENT_CODE_NOT_EXISTS": true,
|
||||
}
|
||||
|
||||
// IsPATError reports whether err is a *PATError.
|
||||
func IsPATError(err error) bool {
|
||||
_, ok := err.(*PATError)
|
||||
return ok
|
||||
}
|
||||
|
||||
// IsPATNoPermissionCode reports whether code is a known PAT permission error code.
|
||||
func IsPATNoPermissionCode(code string) bool {
|
||||
return patNoPermissionCodes[code]
|
||||
}
|
||||
|
||||
// ---- DWS gateway auth errors (shared between PAT & general auth) ----------
|
||||
|
||||
// dwsGatewayErrors is the set of DWS gateway-level auth error codes.
|
||||
var dwsGatewayErrors = map[string]bool{
|
||||
"DWS_SERVICE_UNAUTHORIZED": true,
|
||||
"DWS_AUTH_SERVICE_FAILED": true,
|
||||
}
|
||||
|
||||
// getDWSGatewayErrorCode extracts a DWS gateway error code from errBody
|
||||
// (supports both errorCode and error_code field names).
|
||||
func getDWSGatewayErrorCode(errBody map[string]any) (string, bool) {
|
||||
for _, key := range []string{"errorCode", "error_code"} {
|
||||
if code, ok := errBody[key].(string); ok && dwsGatewayErrors[code] {
|
||||
return code, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// isNotLoggedInError checks if the error body indicates missing authentication.
|
||||
func isNotLoggedInError(body map[string]any) bool {
|
||||
if errMsg, ok := body["error"].(string); ok {
|
||||
if strings.Contains(errMsg, "Missing service_id or access_key") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isBusinessError checks if a parsed JSON body represents a business-level error.
|
||||
func isBusinessError(body map[string]any) bool {
|
||||
if _, ok := body["error"].(string); ok {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(bool); ok && !v {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(string); ok && strings.EqualFold(v, "false") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ---- Classification functions -----------------------------------------------
|
||||
|
||||
// ClassifyToolResultContent checks a raw MCP tool result content map for
|
||||
// DWS gateway auth errors and PAT permission error codes. This is intended
|
||||
// for use as the edition.Hooks.ClassifyToolResult callback so the framework's
|
||||
// runner returns a typed error before its generic business-error classification.
|
||||
//
|
||||
// Check order: DWS gateway auth > PAT permission.
|
||||
func ClassifyToolResultContent(content map[string]any) error {
|
||||
if _, ok := getDWSGatewayErrorCode(content); ok {
|
||||
raw, _ := json.Marshal(content)
|
||||
return NewAuth(string(raw),
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClassifyMCPResponseText classifies a text response returned by an MCP tool call.
|
||||
// Returns a typed error for known gateway auth failures, PAT interceptions,
|
||||
// and business-level errors embedded in HTTP-200 JSON bodies.
|
||||
//
|
||||
// Check order: DWS gateway > PAT permission > generic business error.
|
||||
func ClassifyMCPResponseText(text string) error {
|
||||
var body map[string]any
|
||||
if json.Unmarshal([]byte(text), &body) != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, ok := getDWSGatewayErrorCode(body); ok {
|
||||
return NewAuth(text,
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
|
||||
if isNotLoggedInError(body) {
|
||||
return NewAuth("当前未登录",
|
||||
WithReason("not_configured"),
|
||||
WithHint(notLoggedInHint()),
|
||||
WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := body[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(body, code)}
|
||||
}
|
||||
}
|
||||
|
||||
if isBusinessError(body) {
|
||||
return NewAPI(text,
|
||||
WithReason("business_error"),
|
||||
WithHint(suggestForBusinessErrorText(body)),
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ---- Hints -----------------------------------------------------------------
|
||||
|
||||
func authExpiredHint() string {
|
||||
return "Re-authenticate: dws auth login"
|
||||
}
|
||||
|
||||
func notLoggedInHint() string {
|
||||
return "请先登录:dws auth login"
|
||||
}
|
||||
|
||||
func suggestForBusinessErrorText(body map[string]any) string {
|
||||
msg := ""
|
||||
if v, ok := body["errorMsg"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["message"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["error"].(string); ok {
|
||||
msg = v
|
||||
}
|
||||
switch {
|
||||
case strings.Contains(msg, "搜索内容不能为空"):
|
||||
return "请提供非空搜索关键词: dws doc search --query \"关键词\""
|
||||
case strings.Contains(msg, "User has no permission to access this email"):
|
||||
return "请确认邮箱地址正确,查看可用邮箱: dws mail mailbox list"
|
||||
case strings.Contains(msg, "频率超限") || strings.Contains(msg, "rate limit"):
|
||||
return "API rate limit exceeded, wait a moment and retry"
|
||||
case strings.Contains(msg, "参数错误") || strings.Contains(msg, "param error"):
|
||||
return "Check input parameters. Use --help for available flags"
|
||||
default:
|
||||
return "MCP tool returned a business error; check parameters and refer to skill documentation."
|
||||
}
|
||||
}
|
||||
|
||||
// ---- PAT JSON helpers ------------------------------------------------------
|
||||
|
||||
var patTopLevelStrip = map[string]bool{
|
||||
"success": true, "code": true, "errorCode": true, "error_code": true,
|
||||
"message": true, "error": true, "trace_id": true, "class": true,
|
||||
}
|
||||
|
||||
func cleanPATJSON(body map[string]any, code string) string {
|
||||
out := map[string]any{
|
||||
"success": false,
|
||||
"code": code,
|
||||
}
|
||||
if data, ok := body["data"]; ok {
|
||||
out["data"] = stripClassFields(data)
|
||||
} else {
|
||||
fallback := map[string]any{}
|
||||
for k, v := range body {
|
||||
if !patTopLevelStrip[k] {
|
||||
fallback[k] = v
|
||||
}
|
||||
}
|
||||
if len(fallback) > 0 {
|
||||
out["data"] = stripClassFields(fallback)
|
||||
}
|
||||
}
|
||||
b, err := json.MarshalIndent(out, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Sprintf(`{"success":false,"code":"%s"}`, code)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// ---- Runner adapter functions ------------------------------------------------
|
||||
// These match the function signatures referenced by runner.go's PAT check
|
||||
// framework (ClassifyPatAuthCheck / AsPatAuthCheckError).
|
||||
|
||||
// ClassifyPatAuthCheck is the open-source fallback that checks a tool-call
|
||||
// Content map for PAT permission codes and auth-required codes. Returns a
|
||||
// non-nil *PATError when the content carries a recognised PAT/auth error.
|
||||
func ClassifyPatAuthCheck(content map[string]any) *PATError {
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok {
|
||||
if patNoPermissionCodes[code] || patAuthRequiredCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AsPatAuthCheckError extracts a *PATError from an error chain.
|
||||
func AsPatAuthCheckError(err error) *PATError {
|
||||
var patErr *PATError
|
||||
if stderrors.As(err, &patErr) {
|
||||
return patErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func stripClassFields(v any) any {
|
||||
switch val := v.(type) {
|
||||
case map[string]any:
|
||||
clean := make(map[string]any, len(val))
|
||||
for k, item := range val {
|
||||
if k == "class" {
|
||||
continue
|
||||
}
|
||||
clean[k] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
case []any:
|
||||
clean := make([]any, len(val))
|
||||
for i, item := range val {
|
||||
clean[i] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,521 @@
|
||||
// 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 errors
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// PATError basic behaviour
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPATError_Implements(t *testing.T) {
|
||||
t.Parallel()
|
||||
raw := `{"success":false,"code":"PAT_NO_PERMISSION"}`
|
||||
pe := &PATError{RawJSON: raw}
|
||||
|
||||
if pe.Error() != raw {
|
||||
t.Errorf("Error() = %q, want %q", pe.Error(), raw)
|
||||
}
|
||||
if pe.ExitCode() != ExitCodePermission {
|
||||
t.Errorf("ExitCode() = %d, want %d", pe.ExitCode(), ExitCodePermission)
|
||||
}
|
||||
if pe.RawStderr() != raw {
|
||||
t.Errorf("RawStderr() = %q, want %q", pe.RawStderr(), raw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPATError_True(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PATError{RawJSON: "{}"}
|
||||
if !IsPATError(err) {
|
||||
t.Fatal("expected IsPATError to return true for *PATError")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPATError_False(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := stderrors.New("some other error")
|
||||
if IsPATError(err) {
|
||||
t.Fatal("expected IsPATError to return false for non-PATError")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// IsPATNoPermissionCode
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsPATNoPermissionCode(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
code string
|
||||
want bool
|
||||
}{
|
||||
{"PAT_NO_PERMISSION", true},
|
||||
{"PAT_LOW_RISK_NO_PERMISSION", true},
|
||||
{"PAT_MEDIUM_RISK_NO_PERMISSION", true},
|
||||
{"PAT_HIGH_RISK_NO_PERMISSION", true},
|
||||
{"AGENT_CODE_NOT_EXISTS", false},
|
||||
{"UNKNOWN_CODE", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if got := IsPATNoPermissionCode(tc.code); got != tc.want {
|
||||
t.Errorf("IsPATNoPermissionCode(%q) = %v, want %v", tc.code, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// getDWSGatewayErrorCode
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetDWSGatewayErrorCode_ErrorCode(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED"}
|
||||
code, ok := getDWSGatewayErrorCode(body)
|
||||
if !ok || code != "DWS_SERVICE_UNAUTHORIZED" {
|
||||
t.Errorf("got (%q, %v), want (DWS_SERVICE_UNAUTHORIZED, true)", code, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_ErrorCodeUnderscore(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error_code": "DWS_AUTH_SERVICE_FAILED"}
|
||||
code, ok := getDWSGatewayErrorCode(body)
|
||||
if !ok || code != "DWS_AUTH_SERVICE_FAILED" {
|
||||
t.Errorf("got (%q, %v), want (DWS_AUTH_SERVICE_FAILED, true)", code, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_Unknown(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"errorCode": "SOME_OTHER_ERROR"}
|
||||
_, ok := getDWSGatewayErrorCode(body)
|
||||
if ok {
|
||||
t.Fatal("expected ok=false for unknown error code")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{}
|
||||
_, ok := getDWSGatewayErrorCode(body)
|
||||
if ok {
|
||||
t.Fatal("expected ok=false for empty body")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// isNotLoggedInError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsNotLoggedInError_True(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "Missing service_id or access_key in request headers"}
|
||||
if !isNotLoggedInError(body) {
|
||||
t.Fatal("expected true for Missing service_id message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNotLoggedInError_False(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "something else happened"}
|
||||
if isNotLoggedInError(body) {
|
||||
t.Fatal("expected false for unrelated error message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNotLoggedInError_NoErrorField(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"message": "Missing service_id or access_key"}
|
||||
if isNotLoggedInError(body) {
|
||||
t.Fatal("expected false when error field is absent")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// isBusinessError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsBusinessError_ErrorField(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "some error message"}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when 'error' field is present")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessBoolFalse(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": false}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when success=false (bool)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessStringFalse(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": "False"}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when success=\"False\" (string)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessTrue(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": true, "data": "ok"}
|
||||
if isBusinessError(body) {
|
||||
t.Fatal("expected false when success=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_EmptyBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"data": "hello"}
|
||||
if isBusinessError(body) {
|
||||
t.Fatal("expected false for body without error indicators")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyToolResultContent
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyToolResultContent_GatewayAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED", "message": "expired"}
|
||||
err := ClassifyToolResultContent(content)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error for gateway auth")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Category != CategoryAuth {
|
||||
t.Errorf("Category = %v, want %v", typed.Category, CategoryAuth)
|
||||
}
|
||||
if typed.Reason != "gateway_auth_expired" {
|
||||
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyToolResultContent_PATPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"data": map[string]any{"desc": "需要授权"},
|
||||
}
|
||||
err := ClassifyToolResultContent(content)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error for PAT permission")
|
||||
}
|
||||
var patErr *PATError
|
||||
if !stderrors.As(err, &patErr) {
|
||||
t.Fatalf("expected *PATError, got %T", err)
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain PAT_NO_PERMISSION, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyToolResultContent_NoError(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"success": true, "data": "ok"}
|
||||
if err := ClassifyToolResultContent(content); err != nil {
|
||||
t.Fatalf("expected nil error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyMCPResponseText
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyMCPResponseText_GatewayAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"errorCode":"DWS_SERVICE_UNAUTHORIZED","message":"token expired"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "gateway_auth_expired" {
|
||||
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_NotLoggedIn(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"error":"Missing service_id or access_key in headers"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "not_configured" {
|
||||
t.Errorf("Reason = %q, want not_configured", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_PATPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"code":"PAT_HIGH_RISK_NO_PERMISSION","data":{"desc":"high risk"}}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var patErr *PATError
|
||||
if !stderrors.As(err, &patErr) {
|
||||
t.Fatalf("expected *PATError, got %T", err)
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_HIGH_RISK_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_BusinessError(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"success":false,"errorMsg":"搜索内容不能为空"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "business_error" {
|
||||
t.Errorf("Reason = %q, want business_error", typed.Reason)
|
||||
}
|
||||
if !strings.Contains(typed.Hint, "搜索关键词") {
|
||||
t.Errorf("Hint should contain search suggestion, got: %s", typed.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_InvalidJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := "not json at all"
|
||||
if err := ClassifyMCPResponseText(text); err != nil {
|
||||
t.Fatalf("expected nil for invalid JSON, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_NoError(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"success":true,"data":"hello"}`
|
||||
if err := ClassifyMCPResponseText(text); err != nil {
|
||||
t.Fatalf("expected nil for success response, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyPatAuthCheck
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyPatAuthCheck_PATNoPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"code": "PAT_NO_PERMISSION", "data": map[string]any{"flowId": "f1"}}
|
||||
patErr := ClassifyPatAuthCheck(content)
|
||||
if patErr == nil {
|
||||
t.Fatal("expected non-nil *PATError")
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_AgentCodeNotExists(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"errorCode": "AGENT_CODE_NOT_EXISTS", "data": map[string]any{"clientId": "c1"}}
|
||||
patErr := ClassifyPatAuthCheck(content)
|
||||
if patErr == nil {
|
||||
t.Fatal("expected non-nil *PATError for AGENT_CODE_NOT_EXISTS")
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "AGENT_CODE_NOT_EXISTS") {
|
||||
t.Errorf("RawJSON should contain AGENT_CODE_NOT_EXISTS, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_NoMatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"code": "SOME_BUSINESS_ERROR", "message": "oops"}
|
||||
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
|
||||
t.Fatalf("expected nil, got %v", patErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_EmptyContent(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{}
|
||||
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
|
||||
t.Fatalf("expected nil for empty content, got %v", patErr)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AsPatAuthCheckError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAsPatAuthCheckError_Wrapped(t *testing.T) {
|
||||
t.Parallel()
|
||||
inner := &PATError{RawJSON: `{"code":"PAT_NO_PERMISSION"}`}
|
||||
wrapped := stderrors.Join(stderrors.New("context"), inner)
|
||||
got := AsPatAuthCheckError(wrapped)
|
||||
if got == nil {
|
||||
t.Fatal("expected non-nil *PATError from wrapped error")
|
||||
}
|
||||
if got.RawJSON != inner.RawJSON {
|
||||
t.Errorf("RawJSON = %q, want %q", got.RawJSON, inner.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAsPatAuthCheckError_NotPAT(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := stderrors.New("just a plain error")
|
||||
if got := AsPatAuthCheckError(err); got != nil {
|
||||
t.Fatalf("expected nil for non-PAT error, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// cleanPATJSON
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCleanPATJSON_WithData(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"data": map[string]any{
|
||||
"desc": "需要授权",
|
||||
"flowId": "f123",
|
||||
"class": "com.foo.Bar",
|
||||
},
|
||||
}
|
||||
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
|
||||
if !strings.Contains(result, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("expected code in output, got: %s", result)
|
||||
}
|
||||
if !strings.Contains(result, "flowId") {
|
||||
t.Errorf("expected flowId in data, got: %s", result)
|
||||
}
|
||||
if strings.Contains(result, "class") {
|
||||
t.Errorf("expected class field to be stripped, got: %s", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanPATJSON_WithoutData(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"message": "no permission",
|
||||
"extra": "value",
|
||||
}
|
||||
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
|
||||
if !strings.Contains(result, "extra") {
|
||||
t.Errorf("expected extra field in fallback data, got: %s", result)
|
||||
}
|
||||
// Top-level stripped fields should not appear
|
||||
if strings.Contains(result, `"message"`) {
|
||||
t.Errorf("expected message to be stripped from top level, got: %s", result)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// stripClassFields
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStripClassFields_Map(t *testing.T) {
|
||||
t.Parallel()
|
||||
input := map[string]any{
|
||||
"name": "test",
|
||||
"class": "com.foo.Bar",
|
||||
"nested": map[string]any{
|
||||
"value": 42,
|
||||
"class": "com.baz.Qux",
|
||||
},
|
||||
}
|
||||
result := stripClassFields(input).(map[string]any)
|
||||
if _, ok := result["class"]; ok {
|
||||
t.Error("top-level class should be removed")
|
||||
}
|
||||
nested := result["nested"].(map[string]any)
|
||||
if _, ok := nested["class"]; ok {
|
||||
t.Error("nested class should be removed")
|
||||
}
|
||||
if nested["value"] != 42 {
|
||||
t.Errorf("nested value should be preserved, got %v", nested["value"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripClassFields_Array(t *testing.T) {
|
||||
t.Parallel()
|
||||
input := []any{
|
||||
map[string]any{"id": 1, "class": "Foo"},
|
||||
map[string]any{"id": 2},
|
||||
}
|
||||
result := stripClassFields(input).([]any)
|
||||
first := result[0].(map[string]any)
|
||||
if _, ok := first["class"]; ok {
|
||||
t.Error("class in array element should be removed")
|
||||
}
|
||||
if first["id"] != 1 {
|
||||
t.Error("other fields in array element should be preserved")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripClassFields_Scalar(t *testing.T) {
|
||||
t.Parallel()
|
||||
if stripClassFields("hello") != "hello" {
|
||||
t.Error("scalar string should pass through unchanged")
|
||||
}
|
||||
if stripClassFields(42) != 42 {
|
||||
t.Error("scalar int should pass through unchanged")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// suggestForBusinessErrorText
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSuggestForBusinessErrorText(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
body map[string]any
|
||||
contains string
|
||||
}{
|
||||
{map[string]any{"errorMsg": "搜索内容不能为空"}, "搜索关键词"},
|
||||
{map[string]any{"message": "User has no permission to access this email"}, "邮箱"},
|
||||
{map[string]any{"error": "频率超限"}, "rate limit"},
|
||||
{map[string]any{"errorMsg": "参数错误"}, "parameters"},
|
||||
{map[string]any{"error": "unknown"}, "business error"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
hint := suggestForBusinessErrorText(tc.body)
|
||||
if !strings.Contains(strings.ToLower(hint), strings.ToLower(tc.contains)) {
|
||||
t.Errorf("suggestForBusinessErrorText(%v) = %q, want to contain %q", tc.body, hint, tc.contains)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -20,9 +20,27 @@ import (
|
||||
"strings"
|
||||
|
||||
registryassets "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/registry"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILLS_PERSONAS_FILE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "覆盖内置 personas.yaml 的本地文件路径",
|
||||
Example: "/path/to/personas.yaml",
|
||||
Hidden: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILLS_RECIPES_FILE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "覆盖内置 recipes.yaml 的本地文件路径",
|
||||
Example: "/path/to/recipes.yaml",
|
||||
Hidden: true,
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
PersonaRegistryPathEnv = "DWS_SKILLS_PERSONAS_FILE"
|
||||
RecipeRegistryPathEnv = "DWS_SKILLS_RECIPES_FILE"
|
||||
|
||||
@@ -37,9 +37,20 @@ import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"golang.org/x/text/language"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_LANG",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "界面语言 (en/zh),回退到 LANG",
|
||||
DefaultValue: "en",
|
||||
Example: "zh",
|
||||
})
|
||||
}
|
||||
|
||||
//go:embed locales/*.json
|
||||
var localeFS embed.FS
|
||||
|
||||
|
||||
@@ -162,7 +162,7 @@
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " The organization admin has not enabled \"Allow members to access their personal data via CLI\".",
|
||||
" 组织主管理员:": " Organization super admins: ",
|
||||
" 请联系组织主管理员开启后重新登录。": " Please contact the organization super admin to enable it and re-login.",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " Admin settings: https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"管理员操作入口:": "Admin settings: ",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "CLI data access is not enabled for this organization, please contact admin to enable it",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ CLI data access is not enabled for this organization, please submit an authorization request in the browser...",
|
||||
"✅ 权限已开启,继续登录...": "✅ Permission enabled, continuing login...",
|
||||
@@ -170,5 +170,25 @@
|
||||
"等待提交申请中": "Waiting to submit request",
|
||||
"操作超时,请重新登录": "Operation timed out, please re-login",
|
||||
"检查组织 CLI 授权状态...": "Checking organization CLI auth status...",
|
||||
"🔐 登录钉钉": "🔐 Login to DingTalk"
|
||||
"🔐 登录钉钉": "🔐 Login to DingTalk",
|
||||
"插件管理": "Manage plugins",
|
||||
"列出已安装的插件": "List installed plugins",
|
||||
"安装插件": "Install a plugin",
|
||||
"查看插件详情": "Show plugin details",
|
||||
"启用插件": "Enable a plugin",
|
||||
"禁用插件": "Disable a plugin",
|
||||
"卸载已安装的插件": "Remove an installed plugin",
|
||||
"校验 plugin.json": "Validate a plugin.json",
|
||||
"脚手架生成新插件目录": "Scaffold a new plugin directory",
|
||||
"将本地目录注册为开发态插件": "Register a local directory as a dev plugin",
|
||||
"管理插件配置": "Manage plugin configuration",
|
||||
"设置插件配置项": "Set a plugin config value",
|
||||
"读取插件配置项": "Get a plugin config value",
|
||||
"列出插件所有配置项": "List all config values for a plugin",
|
||||
"删除插件配置项": "Remove a plugin config value",
|
||||
"将插件 stdio server 编译为原生二进制": "Build plugin's stdio server into a native binary",
|
||||
"覆盖 OAuth 客户端 ID (钉钉 AppKey)": "Override OAuth client ID (DingTalk AppKey)",
|
||||
"覆盖 OAuth 客户端密钥 (钉钉 AppSecret)": "Override OAuth client secret (DingTalk AppSecret)",
|
||||
"查看任意命令的帮助信息": "Help about any command",
|
||||
"显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。": "Help provides help for any command in the application.\nSimply type dws help [path to command] for full details."
|
||||
}
|
||||
|
||||
@@ -162,7 +162,7 @@
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。",
|
||||
" 组织主管理员:": " 组织主管理员:",
|
||||
" 请联系组织主管理员开启后重新登录。": " 请联系组织主管理员开启后重新登录。",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"管理员操作入口:": "管理员操作入口:",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "该组织尚未开启 CLI 数据访问权限,请联系管理员开启",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...",
|
||||
"✅ 权限已开启,继续登录...": "✅ 权限已开启,继续登录...",
|
||||
@@ -170,5 +170,25 @@
|
||||
"等待提交申请中": "等待提交申请中",
|
||||
"操作超时,请重新登录": "操作超时,请重新登录",
|
||||
"检查组织 CLI 授权状态...": "检查组织 CLI 授权状态...",
|
||||
"🔐 登录钉钉": "🔐 登录钉钉"
|
||||
"🔐 登录钉钉": "🔐 登录钉钉",
|
||||
"插件管理": "插件管理",
|
||||
"列出已安装的插件": "列出已安装的插件",
|
||||
"安装插件": "安装插件",
|
||||
"查看插件详情": "查看插件详情",
|
||||
"启用插件": "启用插件",
|
||||
"禁用插件": "禁用插件",
|
||||
"卸载已安装的插件": "卸载已安装的插件",
|
||||
"校验 plugin.json": "校验 plugin.json",
|
||||
"脚手架生成新插件目录": "脚手架生成新插件目录",
|
||||
"将本地目录注册为开发态插件": "将本地目录注册为开发态插件",
|
||||
"管理插件配置": "管理插件配置",
|
||||
"设置插件配置项": "设置插件配置项",
|
||||
"读取插件配置项": "读取插件配置项",
|
||||
"列出插件所有配置项": "列出插件所有配置项",
|
||||
"删除插件配置项": "删除插件配置项",
|
||||
"将插件 stdio server 编译为原生二进制": "将插件 stdio server 编译为原生二进制",
|
||||
"覆盖 OAuth 客户端 ID (钉钉 AppKey)": "覆盖 OAuth 客户端 ID (钉钉 AppKey)",
|
||||
"覆盖 OAuth 客户端密钥 (钉钉 AppSecret)": "覆盖 OAuth 客户端密钥 (钉钉 AppSecret)",
|
||||
"查看任意命令的帮助信息": "查看任意命令的帮助信息",
|
||||
"显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。": "显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。"
|
||||
}
|
||||
|
||||
+18
-16
@@ -107,6 +107,7 @@ type CLIGroupDef struct {
|
||||
// CLIToolOverride maps an MCP tool to a CLI command with flag aliases and transforms.
|
||||
type CLIToolOverride struct {
|
||||
CLIName string `json:"cliName"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Group string `json:"group,omitempty"`
|
||||
IsSensitive bool `json:"isSensitive,omitempty"`
|
||||
Hidden bool `json:"hidden,omitempty"`
|
||||
@@ -180,22 +181,23 @@ type DetailLocator struct {
|
||||
}
|
||||
|
||||
type ServerDescriptor struct {
|
||||
Key string `json:"key"`
|
||||
SourceServerID string `json:"source_server_id,omitempty"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
SchemaURI string `json:"schema_uri,omitempty"`
|
||||
NegotiatedProtocolVersion string `json:"negotiated_protocol_version,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
PublishedAt time.Time `json:"published_at,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Source string `json:"source"`
|
||||
Degraded bool `json:"degraded"`
|
||||
DetailLocator DetailLocator `json:"detail_locator,omitempty"`
|
||||
Lifecycle LifecycleInfo `json:"lifecycle,omitempty"`
|
||||
CLI CLIOverlay `json:"cli,omitempty"`
|
||||
HasCLIMeta bool `json:"has_cli_meta,omitempty"`
|
||||
Key string `json:"key"`
|
||||
SourceServerID string `json:"source_server_id,omitempty"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
SchemaURI string `json:"schema_uri,omitempty"`
|
||||
NegotiatedProtocolVersion string `json:"negotiated_protocol_version,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
PublishedAt time.Time `json:"published_at,omitempty"`
|
||||
Status string `json:"status,omitempty"`
|
||||
Source string `json:"source"`
|
||||
Degraded bool `json:"degraded"`
|
||||
DetailLocator DetailLocator `json:"detail_locator,omitempty"`
|
||||
Lifecycle LifecycleInfo `json:"lifecycle,omitempty"`
|
||||
CLI CLIOverlay `json:"cli,omitempty"`
|
||||
HasCLIMeta bool `json:"has_cli_meta,omitempty"`
|
||||
AuthHeaders map[string]string `json:"auth_headers,omitempty"` // plugin-level auth headers for third-party MCP servers
|
||||
}
|
||||
|
||||
func NewClient(baseURL string, httpClient *http.Client) *Client {
|
||||
|
||||
@@ -6,33 +6,33 @@ import (
|
||||
)
|
||||
|
||||
func TestResolveFieldsShadowing(t *testing.T) {
|
||||
rootCmd := &cobra.Command{Use: "dws"}
|
||||
var globalFields string
|
||||
// Register the global persistent flag.
|
||||
rootCmd.PersistentFlags().StringVar(&globalFields, "fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
t.Run("global persistent flag propagates", func(t *testing.T) {
|
||||
rootCmd := &cobra.Command{Use: "dws"}
|
||||
rootCmd.PersistentFlags().String("fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
|
||||
// 1. Normal command that relies on the global output filter
|
||||
normalCmd := &cobra.Command{Use: "normal"}
|
||||
rootCmd.AddCommand(normalCmd)
|
||||
rootCmd.SetArgs([]string{"normal", "--fields", "data,status"})
|
||||
rootCmd.Execute()
|
||||
normalCmd := &cobra.Command{Use: "normal"}
|
||||
rootCmd.AddCommand(normalCmd)
|
||||
rootCmd.SetArgs([]string{"normal", "--fields", "data,status"})
|
||||
rootCmd.Execute()
|
||||
|
||||
if fields := ResolveFields(normalCmd); fields != "data,status" {
|
||||
t.Errorf("expected 'data,status' for normal cmd, got %q", fields)
|
||||
}
|
||||
if fields := ResolveFields(normalCmd); fields != "data,status" {
|
||||
t.Errorf("expected 'data,status' for normal cmd, got %q", fields)
|
||||
}
|
||||
})
|
||||
|
||||
// 2. Command that shadows the global format flag with its own local business logic
|
||||
bizCmd := &cobra.Command{Use: "biz"}
|
||||
var localFields string
|
||||
bizCmd.Flags().StringVar(&localFields, "fields", "", "JSON string array of objects")
|
||||
rootCmd.AddCommand(bizCmd)
|
||||
t.Run("shadowed local flag is ignored", func(t *testing.T) {
|
||||
rootCmd := &cobra.Command{Use: "dws"}
|
||||
rootCmd.PersistentFlags().String("fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
|
||||
// Reset
|
||||
rootCmd.SetArgs([]string{"biz", "--fields", "[\"fake\"]"})
|
||||
rootCmd.Execute()
|
||||
bizCmd := &cobra.Command{Use: "biz"}
|
||||
bizCmd.Flags().String("fields", "", "JSON string array of objects")
|
||||
rootCmd.AddCommand(bizCmd)
|
||||
|
||||
// It should now correctly ignore the localized fields parameter!
|
||||
if fields := ResolveFields(bizCmd); fields != "" {
|
||||
t.Errorf("expected empty fields for shadowed cmd since it's a business param, got %q", fields)
|
||||
}
|
||||
rootCmd.SetArgs([]string{"biz", "--fields", "[\"fake\"]"})
|
||||
rootCmd.Execute()
|
||||
|
||||
if fields := ResolveFields(bizCmd); fields != "" {
|
||||
t.Errorf("expected empty fields for shadowed cmd since it's a business param, got %q", fields)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
// 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 pat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/fatih/color"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
var validGrantTypes = map[string]bool{
|
||||
"once": true,
|
||||
"session": true,
|
||||
"permanent": true,
|
||||
}
|
||||
|
||||
func newChmodCommand(caller edition.ToolCaller) *cobra.Command {
|
||||
chmodCmd := &cobra.Command{
|
||||
Use: "chmod <scope>...",
|
||||
Short: "授予指定权限",
|
||||
Long: `授予指定 scope 的操作权限。
|
||||
|
||||
scope 格式: <product>.<entity>:<permission>
|
||||
例: aitable.record:read chat.group:write calendar.event:read
|
||||
|
||||
grantType 规则:
|
||||
once 一次性,执行一次后自动失效
|
||||
session 当前会话有效(默认),需要 --session-id
|
||||
permanent 永久有效`,
|
||||
Args: cobra.MinimumNArgs(1),
|
||||
Example: ` dws pat chmod aitable.record:read --agentCode agt-xxxx --grant-type session --session-id session-xxx
|
||||
dws pat chmod chat.message:list --grant-type once --agentCode agt-xxxx
|
||||
dws pat chmod aitable.record:read aitable.record:write --agentCode agt-xxxx --grant-type permanent`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
agentCode, _ := cmd.Flags().GetString("agentCode")
|
||||
if agentCode == "" {
|
||||
return fmt.Errorf("flag --agentCode is required\n hint: dws pat chmod <scope>... --agentCode <id>")
|
||||
}
|
||||
scopes := args
|
||||
grantType, _ := cmd.Flags().GetString("grant-type")
|
||||
sessionID, _ := cmd.Flags().GetString("session-id")
|
||||
|
||||
if !validGrantTypes[grantType] {
|
||||
return fmt.Errorf("invalid --grant-type %q, must be one of: once, session, permanent", grantType)
|
||||
}
|
||||
|
||||
if grantType == "session" && sessionID == "" && os.Getenv("DWS_SESSION_ID") == "" {
|
||||
return fmt.Errorf("--session-id is required when --grant-type is session\n hint: dws pat chmod <scope> --agentCode <id> --grant-type session --session-id <id>")
|
||||
}
|
||||
|
||||
if caller != nil && caller.DryRun() {
|
||||
bold := color.New(color.FgYellow, color.Bold)
|
||||
bold.Println("[DRY-RUN] Preview only, not executed:")
|
||||
fmt.Printf("%-16s%s\n", "Tool:", "个人授权")
|
||||
fmt.Printf("%-16s%s\n", "AgentCode:", agentCode)
|
||||
fmt.Printf("%-16s%v\n", "Scope:", scopes)
|
||||
fmt.Printf("%-16s%s\n", "GrantType:", grantType)
|
||||
if sessionID != "" {
|
||||
fmt.Printf("%-16s%s\n", "SessionID:", sessionID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if caller == nil {
|
||||
return fmt.Errorf("internal error: tool runtime not initialized")
|
||||
}
|
||||
|
||||
toolArgs := map[string]any{
|
||||
"agentCode": agentCode,
|
||||
"scope": scopes,
|
||||
"grantType": grantType,
|
||||
}
|
||||
if sessionID == "" {
|
||||
sessionID = os.Getenv("DWS_SESSION_ID")
|
||||
}
|
||||
if sessionID != "" {
|
||||
toolArgs["sessionId"] = sessionID
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
result, err := caller.CallTool(ctx, "pat", "个人授权", toolArgs)
|
||||
if err != nil {
|
||||
return fmt.Errorf("pat chmod failed: %w", err)
|
||||
}
|
||||
|
||||
return handleToolResult(result)
|
||||
},
|
||||
}
|
||||
|
||||
chmodCmd.Flags().String("agentCode", "", "Agent 唯一标识(必填)")
|
||||
_ = chmodCmd.MarkFlagRequired("agentCode")
|
||||
chmodCmd.Flags().String("grant-type", "session", "授权策略: once|session|permanent")
|
||||
chmodCmd.Flags().String("session-id", "", "会话标识(session 模式下必填)")
|
||||
|
||||
return chmodCmd
|
||||
}
|
||||
|
||||
// handleToolResult processes a ToolResult and writes output to stdout.
|
||||
func handleToolResult(result *edition.ToolResult) error {
|
||||
if result == nil {
|
||||
return fmt.Errorf("empty tool result")
|
||||
}
|
||||
for _, c := range result.Content {
|
||||
if c.Type != "text" || c.Text == "" {
|
||||
continue
|
||||
}
|
||||
if respErr := apperrors.ClassifyMCPResponseText(c.Text); respErr != nil {
|
||||
return respErr
|
||||
}
|
||||
fmt.Println(c.Text)
|
||||
return nil
|
||||
}
|
||||
data, err := json.MarshalIndent(result, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal result: %w", err)
|
||||
}
|
||||
fmt.Println(string(data))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// 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 pat implements the "dws pat" command group for PAT (Personal Action
|
||||
// Token) authorization management.
|
||||
package pat
|
||||
|
||||
import (
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// RegisterCommands adds the pat command tree to rootCmd.
|
||||
func RegisterCommands(root *cobra.Command, c edition.ToolCaller) {
|
||||
patCmd := &cobra.Command{
|
||||
Use: "pat",
|
||||
Short: "行为授权管理",
|
||||
Long: `管理行为授权(PAT)。
|
||||
|
||||
命令结构:
|
||||
dws pat chmod <scope>... 授予指定权限`,
|
||||
RunE: cmdutil.GroupRunE,
|
||||
}
|
||||
|
||||
patCmd.AddCommand(newChmodCommand(c))
|
||||
root.AddCommand(patCmd)
|
||||
}
|
||||
@@ -244,6 +244,194 @@ func TestFullPipelineEndToEnd(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestFullFivePhasePipeline exercises all five phases in order:
|
||||
// Register → PreParse → PostParse → PreRequest → PostResponse,
|
||||
// simulating a complete command lifecycle from registration through
|
||||
// response output.
|
||||
func TestFullFivePhasePipeline(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
RegisterHandler{},
|
||||
AliasHandler{},
|
||||
StickyHandler{},
|
||||
ParamNameHandler{},
|
||||
ParamValueHandler{},
|
||||
PreRequestHandler{},
|
||||
PostResponseHandler{},
|
||||
)
|
||||
|
||||
// Verify all five phases have handlers.
|
||||
for _, phase := range []pipeline.Phase{
|
||||
pipeline.Register,
|
||||
pipeline.PreParse,
|
||||
pipeline.PostParse,
|
||||
pipeline.PreRequest,
|
||||
pipeline.PostResponse,
|
||||
} {
|
||||
if !engine.HasHandlers(phase) {
|
||||
t.Fatalf("engine missing handlers for phase %v", phase)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 1: Register — command tree being built.
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable",
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.Register, ctx); err != nil {
|
||||
t.Fatalf("Register error: %v", err)
|
||||
}
|
||||
|
||||
// Phase 2: PreParse — fix raw argv.
|
||||
ctx.Args = []string{
|
||||
"--userId", "u001",
|
||||
"--pageSize50",
|
||||
"--verbosetrue",
|
||||
}
|
||||
ctx.FlagSpecs = flagSpecs("user-id", "page-size", "verbose")
|
||||
|
||||
if err := engine.RunPhase(pipeline.PreParse, ctx); err != nil {
|
||||
t.Fatalf("PreParse error: %v", err)
|
||||
}
|
||||
|
||||
want := "--user-id u001 --page-size 50 --verbose true"
|
||||
got := strings.Join(ctx.Args, " ")
|
||||
if got != want {
|
||||
t.Errorf("after PreParse: Args = %q, want %q", got, want)
|
||||
}
|
||||
preParseCorrections := len(ctx.Corrections)
|
||||
|
||||
// Phase 3: PostParse — simulate Cobra having parsed the corrected
|
||||
// args into structured params, then normalise values.
|
||||
ctx.Command = "aitable.query_records"
|
||||
ctx.Params = map[string]any{
|
||||
"user_id": "u001",
|
||||
"page_size": "1,000",
|
||||
"verbose": "yes",
|
||||
}
|
||||
ctx.Schema = map[string]any{
|
||||
"properties": map[string]any{
|
||||
"user_id": map[string]any{"type": "string"},
|
||||
"page_size": map[string]any{"type": "integer"},
|
||||
"verbose": map[string]any{"type": "boolean"},
|
||||
},
|
||||
}
|
||||
|
||||
if err := engine.RunPhase(pipeline.PostParse, ctx); err != nil {
|
||||
t.Fatalf("PostParse error: %v", err)
|
||||
}
|
||||
|
||||
if got := ctx.Params["verbose"]; got != true {
|
||||
t.Errorf("verbose = %v (%T), want true (bool)", got, got)
|
||||
}
|
||||
if got := ctx.Params["page_size"]; got != int64(1000) {
|
||||
t.Errorf("page_size = %v, want 1000", got)
|
||||
}
|
||||
postParseCorrections := len(ctx.Corrections) - preParseCorrections
|
||||
if postParseCorrections != 2 {
|
||||
t.Errorf("PostParse corrections = %d, want 2", postParseCorrections)
|
||||
}
|
||||
|
||||
// Phase 4: PreRequest — inspect final payload before dispatch.
|
||||
ctx.Payload = ctx.Params
|
||||
if err := engine.RunPhase(pipeline.PreRequest, ctx); err != nil {
|
||||
t.Fatalf("PreRequest error: %v", err)
|
||||
}
|
||||
// Verify payload was not corrupted.
|
||||
if ctx.Payload["user_id"] != "u001" {
|
||||
t.Error("PreRequest corrupted Payload")
|
||||
}
|
||||
|
||||
// Phase 5: PostResponse — process response before output.
|
||||
ctx.Response = map[string]any{
|
||||
"records": []any{
|
||||
map[string]any{"id": "rec001", "fields": map[string]any{"name": "test"}},
|
||||
},
|
||||
"total": 1,
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.PostResponse, ctx); err != nil {
|
||||
t.Fatalf("PostResponse error: %v", err)
|
||||
}
|
||||
// Verify response was not corrupted.
|
||||
if ctx.Response["total"] != 1 {
|
||||
t.Error("PostResponse corrupted Response")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFullFivePhasePipelineWithEngineRun exercises all five phases
|
||||
// using Engine.Run (single shot) to verify the ordering is correct
|
||||
// end-to-end.
|
||||
func TestFullFivePhasePipelineWithEngineRun(t *testing.T) {
|
||||
var seq []string
|
||||
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
&phaseTracker{name: "reg", phase: pipeline.Register, seq: &seq},
|
||||
&phaseTracker{name: "pre-parse", phase: pipeline.PreParse, seq: &seq},
|
||||
&phaseTracker{name: "post-parse", phase: pipeline.PostParse, seq: &seq},
|
||||
&phaseTracker{name: "pre-req", phase: pipeline.PreRequest, seq: &seq},
|
||||
&phaseTracker{name: "post-resp", phase: pipeline.PostResponse, seq: &seq},
|
||||
)
|
||||
|
||||
ctx := &pipeline.Context{Command: "test.tool"}
|
||||
if err := engine.Run(ctx); err != nil {
|
||||
t.Fatalf("Engine.Run error: %v", err)
|
||||
}
|
||||
|
||||
want := "reg,pre-parse,post-parse,pre-req,post-resp"
|
||||
got := strings.Join(seq, ",")
|
||||
if got != want {
|
||||
t.Errorf("phase execution order = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFivePhasePipelineCorrectHandlerCounts verifies that the
|
||||
// production-equivalent engine has the expected handler distribution.
|
||||
func TestFivePhasePipelineCorrectHandlerCounts(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
RegisterHandler{},
|
||||
AliasHandler{},
|
||||
StickyHandler{},
|
||||
ParamNameHandler{},
|
||||
ParamValueHandler{},
|
||||
PreRequestHandler{},
|
||||
PostResponseHandler{},
|
||||
)
|
||||
|
||||
tests := []struct {
|
||||
phase pipeline.Phase
|
||||
want int
|
||||
}{
|
||||
{pipeline.Register, 1},
|
||||
{pipeline.PreParse, 3},
|
||||
{pipeline.PostParse, 1},
|
||||
{pipeline.PreRequest, 1},
|
||||
{pipeline.PostResponse, 1},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := len(engine.Handlers(tt.phase)); got != tt.want {
|
||||
t.Errorf("Handlers(%v) = %d, want %d", tt.phase, got, tt.want)
|
||||
}
|
||||
}
|
||||
if got := engine.HandlerCount(); got != 7 {
|
||||
t.Errorf("HandlerCount = %d, want 7", got)
|
||||
}
|
||||
}
|
||||
|
||||
// phaseTracker is a test helper that records its name when Handle is called.
|
||||
type phaseTracker struct {
|
||||
name string
|
||||
phase pipeline.Phase
|
||||
seq *[]string
|
||||
}
|
||||
|
||||
func (h *phaseTracker) Name() string { return h.name }
|
||||
func (h *phaseTracker) Phase() pipeline.Phase { return h.phase }
|
||||
func (h *phaseTracker) Handle(_ *pipeline.Context) error {
|
||||
*h.seq = append(*h.seq, h.name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestPreParseDoesNotBreakValidArgs verifies that valid, correctly
|
||||
// formatted args pass through the pipeline without modification.
|
||||
func TestPreParseDoesNotBreakValidArgs(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
// PostResponseHandler runs in the PostResponse phase — after the
|
||||
// transport returns a result and before the output is written to
|
||||
// stdout. It receives the raw response and can mutate it.
|
||||
//
|
||||
// Default behaviour: no-op pass-through. This establishes the
|
||||
// extension point for:
|
||||
// - Output format transformation (e.g. table, CSV, YAML renderers)
|
||||
// - Response field filtering or redaction
|
||||
// - Pagination metadata injection
|
||||
// - Response caching or analytics collection
|
||||
//
|
||||
// Logging is handled at the integration point in canonical.go,
|
||||
// consistent with how other phases log at their call sites.
|
||||
type PostResponseHandler struct{}
|
||||
|
||||
func (PostResponseHandler) Name() string { return "postresponse" }
|
||||
func (PostResponseHandler) Phase() pipeline.Phase { return pipeline.PostResponse }
|
||||
|
||||
func (PostResponseHandler) Handle(ctx *pipeline.Context) error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
func TestPostResponseHandlerMeta(t *testing.T) {
|
||||
h := PostResponseHandler{}
|
||||
if got := h.Name(); got != "postresponse" {
|
||||
t.Errorf("Name() = %q, want %q", got, "postresponse")
|
||||
}
|
||||
if got := h.Phase(); got != pipeline.PostResponse {
|
||||
t.Errorf("Phase() = %v, want %v", got, pipeline.PostResponse)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostResponseHandlerEmptyContext(t *testing.T) {
|
||||
h := PostResponseHandler{}
|
||||
ctx := &pipeline.Context{}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostResponseHandlerNoSideEffects(t *testing.T) {
|
||||
h := PostResponseHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable.query_records",
|
||||
Response: map[string]any{
|
||||
"records": []any{
|
||||
map[string]any{"id": "rec001"},
|
||||
},
|
||||
"total": 1,
|
||||
},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
if ctx.Response["total"] != 1 {
|
||||
t.Error("PostResponseHandler should not mutate Response")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostResponseHandlerNilResponse(t *testing.T) {
|
||||
h := PostResponseHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "todo.list",
|
||||
Response: nil,
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPostResponseHandlerInEngine(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.Register(PostResponseHandler{})
|
||||
|
||||
if !engine.HasHandlers(pipeline.PostResponse) {
|
||||
t.Fatal("engine should have PostResponse handler")
|
||||
}
|
||||
|
||||
ctx := &pipeline.Context{
|
||||
Command: "calendar.list_events",
|
||||
Response: map[string]any{"events": []any{}},
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.PostResponse, ctx); err != nil {
|
||||
t.Fatalf("RunPhase(PostResponse) returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
// PreRequestHandler runs in the PreRequest phase — after parameter
|
||||
// validation succeeds and just before the JSON-RPC call is dispatched.
|
||||
// It receives the final payload and can inspect or mutate it.
|
||||
//
|
||||
// Default behaviour: no-op pass-through. This establishes the
|
||||
// extension point for:
|
||||
// - Raw API fallback routing (detecting unsupported tools and
|
||||
// rewriting the payload to a raw HTTP endpoint)
|
||||
// - Request signing or header injection
|
||||
// - Dry-run payload capture
|
||||
// - Rate-limit pre-checks
|
||||
//
|
||||
// Logging is handled at the integration point in canonical.go,
|
||||
// consistent with how other phases log at their call sites.
|
||||
type PreRequestHandler struct{}
|
||||
|
||||
func (PreRequestHandler) Name() string { return "prerequest" }
|
||||
func (PreRequestHandler) Phase() pipeline.Phase { return pipeline.PreRequest }
|
||||
|
||||
func (PreRequestHandler) Handle(ctx *pipeline.Context) error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
func TestPreRequestHandlerMeta(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
if got := h.Name(); got != "prerequest" {
|
||||
t.Errorf("Name() = %q, want %q", got, "prerequest")
|
||||
}
|
||||
if got := h.Phase(); got != pipeline.PreRequest {
|
||||
t.Errorf("Phase() = %v, want %v", got, pipeline.PreRequest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerEmptyContext(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerNoSideEffects(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable.query_records",
|
||||
Params: map[string]any{
|
||||
"spaceId": "sp001",
|
||||
"datasheetId": "ds001",
|
||||
},
|
||||
Payload: map[string]any{
|
||||
"spaceId": "sp001",
|
||||
"datasheetId": "ds001",
|
||||
},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
if ctx.Params["spaceId"] != "sp001" {
|
||||
t.Error("PreRequestHandler should not mutate Params")
|
||||
}
|
||||
if ctx.Payload["spaceId"] != "sp001" {
|
||||
t.Error("PreRequestHandler should not mutate Payload")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerNilPayload(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "chat.send_message",
|
||||
Params: map[string]any{"userId": "u001"},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerInEngine(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.Register(PreRequestHandler{})
|
||||
|
||||
if !engine.HasHandlers(pipeline.PreRequest) {
|
||||
t.Fatal("engine should have PreRequest handler")
|
||||
}
|
||||
|
||||
ctx := &pipeline.Context{
|
||||
Command: "todo.create",
|
||||
Params: map[string]any{"subject": "test"},
|
||||
Payload: map[string]any{"subject": "test"},
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.PreRequest, ctx); err != nil {
|
||||
t.Fatalf("RunPhase(PreRequest) returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
// RegisterHandler runs during the Register phase — the first stage
|
||||
// in the pipeline, executed while the Cobra command tree is being
|
||||
// built. It validates that the registration context carries a
|
||||
// non-empty command identifier.
|
||||
//
|
||||
// The handler is intentionally lightweight and side-effect free.
|
||||
// This provides the structural hook for future extensions (e.g.
|
||||
// dynamic command injection, feature gating, or Raw API fallback
|
||||
// command registration) without adding any runtime overhead to
|
||||
// the default path. Logging is handled at the call site in
|
||||
// canonical.go, consistent with how PreParse logging is done
|
||||
// in cobra.go.
|
||||
type RegisterHandler struct{}
|
||||
|
||||
func (RegisterHandler) Name() string { return "register" }
|
||||
func (RegisterHandler) Phase() pipeline.Phase { return pipeline.Register }
|
||||
|
||||
func (RegisterHandler) Handle(ctx *pipeline.Context) error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
// 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 handlers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
func TestRegisterHandlerMeta(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
if got := h.Name(); got != "register" {
|
||||
t.Errorf("Name() = %q, want %q", got, "register")
|
||||
}
|
||||
if got := h.Phase(); got != pipeline.Register {
|
||||
t.Errorf("Phase() = %v, want %v", got, pipeline.Register)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerEmptyContext(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerWithCommand(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable",
|
||||
Schema: map[string]any{
|
||||
"properties": map[string]any{
|
||||
"spaceId": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerNoSideEffects(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "todo",
|
||||
Params: map[string]any{"key": "value"},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
if ctx.Params["key"] != "value" {
|
||||
t.Error("RegisterHandler should not mutate Params")
|
||||
}
|
||||
if ctx.Command != "todo" {
|
||||
t.Error("RegisterHandler should not mutate Command")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerInEngine(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.Register(RegisterHandler{})
|
||||
|
||||
if !engine.HasHandlers(pipeline.Register) {
|
||||
t.Fatal("engine should have Register handler")
|
||||
}
|
||||
|
||||
ctx := &pipeline.Context{Command: "calendar"}
|
||||
if err := engine.RunPhase(pipeline.Register, ctx); err != nil {
|
||||
t.Fatalf("RunPhase(Register) returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,156 @@
|
||||
// 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 plugin
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
// UserContext holds the minimal user identity fields injected into
|
||||
// stdio plugin subprocesses via environment variables.
|
||||
type UserContext struct {
|
||||
UserID string
|
||||
CorpID string
|
||||
}
|
||||
|
||||
// StdioServerClient pairs a transport.StdioClient with its server key.
|
||||
type StdioServerClient struct {
|
||||
Key string
|
||||
Client *transport.StdioClient
|
||||
}
|
||||
|
||||
// StdioClients returns StdioClient instances for all stdio-type MCP
|
||||
// servers declared by this plugin. uc is the current user's identity;
|
||||
// if non-nil, DWS_USER_ID and DWS_CORP_ID are injected as environment
|
||||
// variables so that the subprocess can identify the caller without
|
||||
// implementing its own auth.
|
||||
func (p *Plugin) StdioClients(uc *UserContext) []StdioServerClient {
|
||||
var clients []StdioServerClient
|
||||
for key, srv := range p.Manifest.MCPServers {
|
||||
if srv.Type != "stdio" {
|
||||
continue
|
||||
}
|
||||
|
||||
command := srv.Command
|
||||
if command == "" {
|
||||
slog.Warn("plugin: stdio server missing command",
|
||||
"plugin", p.Manifest.Name, "server", key)
|
||||
continue
|
||||
}
|
||||
|
||||
// Expand ${DWS_PLUGIN_ROOT} in command and args.
|
||||
command = expandPluginVars(command, p.Root)
|
||||
args := make([]string, len(srv.Args))
|
||||
for i, a := range srv.Args {
|
||||
args[i] = expandPluginVars(a, p.Root)
|
||||
}
|
||||
|
||||
env := make(map[string]string)
|
||||
for k, v := range srv.Env {
|
||||
env[k] = expandPluginVars(v, p.Root)
|
||||
}
|
||||
env["DWS_PLUGIN_ROOT"] = p.Root
|
||||
env["DWS_PLUGIN_DATA"] = filepath.Join(filepath.Dir(filepath.Dir(p.Root)), "data", p.Manifest.Name)
|
||||
|
||||
// Inject user identity so the subprocess knows who is calling.
|
||||
if uc != nil {
|
||||
if uc.UserID != "" {
|
||||
env["DWS_USER_ID"] = uc.UserID
|
||||
}
|
||||
if uc.CorpID != "" {
|
||||
env["DWS_CORP_ID"] = uc.CorpID
|
||||
}
|
||||
}
|
||||
|
||||
sc := transport.NewStdioClient(command, args, env)
|
||||
clients = append(clients, StdioServerClient{Key: key, Client: sc})
|
||||
}
|
||||
return clients
|
||||
}
|
||||
|
||||
// expandPluginVars replaces ${DWS_PLUGIN_ROOT} with the actual plugin
|
||||
// root path and ${DWS_PLUGIN_DATA} with the data directory.
|
||||
func expandPluginVars(s, root string) string {
|
||||
s = strings.ReplaceAll(s, "${DWS_PLUGIN_ROOT}", root)
|
||||
dataDir := filepath.Join(filepath.Dir(filepath.Dir(root)), "data")
|
||||
s = strings.ReplaceAll(s, "${DWS_PLUGIN_DATA}", dataDir)
|
||||
return os.Expand(s, os.Getenv)
|
||||
}
|
||||
|
||||
// ToServerDescriptors converts a loaded plugin's MCP servers into
|
||||
// market.ServerDescriptor values suitable for SetDynamicServers.
|
||||
// Only streamable-http servers are converted; stdio servers are
|
||||
// skipped (they require the stdio transport extension).
|
||||
func (p *Plugin) ToServerDescriptors() []market.ServerDescriptor {
|
||||
var descriptors []market.ServerDescriptor
|
||||
for key, srv := range p.Manifest.MCPServers {
|
||||
if srv.Type != "streamable-http" {
|
||||
slog.Debug("plugin: skipping non-http server",
|
||||
"plugin", p.Manifest.Name,
|
||||
"server", key,
|
||||
"type", srv.Type,
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
overlay := market.CLIOverlay{}
|
||||
if len(srv.CLI) > 0 {
|
||||
if err := json.Unmarshal(srv.CLI, &overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLIOverlay",
|
||||
"plugin", p.Manifest.Name,
|
||||
"server", key,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// Ensure the overlay has an ID — fall back to server key.
|
||||
if overlay.ID == "" {
|
||||
overlay.ID = key
|
||||
}
|
||||
if overlay.Command == "" {
|
||||
overlay.Command = key
|
||||
}
|
||||
|
||||
source := "plugin"
|
||||
|
||||
// Resolve headers: expand environment variable references (e.g. ${DASHSCOPE_API_KEY}).
|
||||
var resolvedHeaders map[string]string
|
||||
if len(srv.Headers) > 0 {
|
||||
resolvedHeaders = make(map[string]string, len(srv.Headers))
|
||||
for headerKey, headerVal := range srv.Headers {
|
||||
resolvedHeaders[headerKey] = expandPluginVars(headerVal, p.Root)
|
||||
}
|
||||
}
|
||||
|
||||
descriptors = append(descriptors, market.ServerDescriptor{
|
||||
Key: key,
|
||||
DisplayName: p.Manifest.Name + "/" + key,
|
||||
Description: p.Manifest.Description,
|
||||
Endpoint: srv.Endpoint,
|
||||
Source: source,
|
||||
CLI: overlay,
|
||||
HasCLIMeta: len(srv.CLI) > 0,
|
||||
AuthHeaders: resolvedHeaders,
|
||||
})
|
||||
}
|
||||
return descriptors
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
// 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 plugin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
const defaultHookTimeout = 30 * time.Second
|
||||
|
||||
// HookAdapter wraps a plugin hook entry as a pipeline.Handler.
|
||||
type HookAdapter struct {
|
||||
pluginName string
|
||||
entry HookEntry
|
||||
phase pipeline.Phase
|
||||
timeout time.Duration
|
||||
}
|
||||
|
||||
// NewHookAdapter creates a pipeline handler from a plugin hook entry.
|
||||
func NewHookAdapter(pluginName string, entry HookEntry) *HookAdapter {
|
||||
phase := parsePhase(entry.Phase)
|
||||
timeout := defaultHookTimeout
|
||||
if entry.Timeout > 0 {
|
||||
timeout = time.Duration(entry.Timeout) * time.Second
|
||||
}
|
||||
return &HookAdapter{
|
||||
pluginName: pluginName,
|
||||
entry: entry,
|
||||
phase: phase,
|
||||
timeout: timeout,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Name() string {
|
||||
return fmt.Sprintf("plugin-hook:%s/%s", h.pluginName, h.entry.Phase)
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Phase() pipeline.Phase {
|
||||
return h.phase
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Handle(ctx *pipeline.Context) error {
|
||||
// Check matcher: if set, only run for matching commands.
|
||||
if h.entry.Matcher != "" {
|
||||
matched, err := filepath.Match(h.entry.Matcher, ctx.Command)
|
||||
if err != nil || !matched {
|
||||
return nil // skip silently
|
||||
}
|
||||
}
|
||||
|
||||
// Serialize context to JSON for the hook's stdin.
|
||||
input, err := json.Marshal(map[string]any{
|
||||
"command": ctx.Command,
|
||||
"params": ctx.Params,
|
||||
"args": ctx.Args,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Warn("plugin hook: failed to serialize context",
|
||||
"plugin", h.pluginName, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
timeoutCtx, cancel := context.WithTimeout(context.Background(), h.timeout)
|
||||
defer cancel()
|
||||
|
||||
cmd := exec.CommandContext(timeoutCtx, "sh", "-c", h.entry.Command)
|
||||
cmd.Stdin = strings.NewReader(string(input))
|
||||
output, err := cmd.CombinedOutput()
|
||||
|
||||
if err != nil {
|
||||
if exitErr, ok := err.(*exec.ExitError); ok {
|
||||
code := exitErr.ExitCode()
|
||||
if code == 2 {
|
||||
// Exit 2 = abort pipeline.
|
||||
return fmt.Errorf("plugin hook %s/%s aborted: %s",
|
||||
h.pluginName, h.entry.Phase, strings.TrimSpace(string(output)))
|
||||
}
|
||||
}
|
||||
slog.Warn("plugin hook failed",
|
||||
"plugin", h.pluginName,
|
||||
"phase", h.entry.Phase,
|
||||
"error", err,
|
||||
"output", string(output),
|
||||
)
|
||||
return nil // non-fatal: log warning and continue
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func parsePhase(s string) pipeline.Phase {
|
||||
switch strings.TrimSpace(strings.ToLower(s)) {
|
||||
case "pre-parse":
|
||||
return pipeline.PreParse
|
||||
case "post-parse":
|
||||
return pipeline.PostParse
|
||||
case "pre-request":
|
||||
return pipeline.PreRequest
|
||||
case "post-response":
|
||||
return pipeline.PostResponse
|
||||
default:
|
||||
return pipeline.PreRequest
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,869 @@
|
||||
// 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 plugin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
// Loader scans plugin directories and returns loaded, validated plugins.
|
||||
type Loader struct {
|
||||
// PluginsDir is the root directory for all plugins.
|
||||
// Defaults to ~/.dws/plugins/.
|
||||
PluginsDir string
|
||||
|
||||
// CLIVersion is the current CLI version, used for
|
||||
// minCLIVersion compatibility checks.
|
||||
CLIVersion string
|
||||
}
|
||||
|
||||
// NewLoader creates a Loader with default paths.
|
||||
func NewLoader(cliVersion string) *Loader {
|
||||
home, _ := os.UserHomeDir()
|
||||
return &Loader{
|
||||
PluginsDir: filepath.Join(home, ".dws", "plugins"),
|
||||
CLIVersion: cliVersion,
|
||||
}
|
||||
}
|
||||
|
||||
// Settings holds user preferences for plugin management.
|
||||
type Settings struct {
|
||||
EnabledPlugins map[string]bool `json:"enabledPlugins,omitempty"`
|
||||
PluginConfigs map[string]map[string]any `json:"pluginConfigs,omitempty"`
|
||||
PluginAutoUpdate bool `json:"pluginAutoUpdate,omitempty"`
|
||||
DevPlugins map[string]string `json:"devPlugins,omitempty"` // name → absolute path
|
||||
}
|
||||
|
||||
// LoadUser scans ~/.dws/plugins/user/ and returns enabled user plugins.
|
||||
func (l *Loader) LoadUser() []*Plugin {
|
||||
userDir := filepath.Join(l.PluginsDir, "user")
|
||||
settings := l.loadSettings()
|
||||
|
||||
var plugins []*Plugin
|
||||
// User plugins may be nested: user/{workspace}/{name}/
|
||||
entries, err := os.ReadDir(userDir)
|
||||
if err != nil {
|
||||
if !os.IsNotExist(err) {
|
||||
slog.Debug("plugin: cannot read user dir", "path", userDir, "error", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
entryPath := filepath.Join(userDir, entry.Name())
|
||||
|
||||
// Check if this is a direct plugin directory (has plugin.json)
|
||||
if _, err := os.Stat(filepath.Join(entryPath, "plugin.json")); err == nil {
|
||||
p := l.loadPlugin(entryPath)
|
||||
if p != nil && isPluginEnabled(settings, p.Manifest.Name) {
|
||||
plugins = append(plugins, p)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// Otherwise treat as workspace directory: user/{workspace}/{name}/
|
||||
subEntries, err := os.ReadDir(entryPath)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
for _, sub := range subEntries {
|
||||
if !sub.IsDir() {
|
||||
continue
|
||||
}
|
||||
subPath := filepath.Join(entryPath, sub.Name())
|
||||
p := l.loadPlugin(subPath)
|
||||
if p != nil {
|
||||
qualifiedName := entry.Name() + "/" + p.Manifest.Name
|
||||
if isPluginEnabled(settings, qualifiedName) {
|
||||
plugins = append(plugins, p)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return plugins
|
||||
}
|
||||
|
||||
// LoadAll loads user + dev plugins.
|
||||
func (l *Loader) LoadAll() []*Plugin {
|
||||
user := l.LoadUser()
|
||||
dev := l.LoadDev()
|
||||
return append(user, dev...)
|
||||
}
|
||||
|
||||
// loadPlugin reads and validates a single plugin directory.
|
||||
func (l *Loader) loadPlugin(dir string) *Plugin {
|
||||
manifestPath := filepath.Join(dir, "plugin.json")
|
||||
manifest, err := ParseManifest(manifestPath)
|
||||
if err != nil {
|
||||
slog.Warn("plugin: failed to parse manifest",
|
||||
"path", manifestPath, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := manifest.Validate(l.CLIVersion); err != nil {
|
||||
slog.Warn("plugin: validation failed",
|
||||
"plugin", manifest.Name, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
return &Plugin{
|
||||
Manifest: *manifest,
|
||||
Root: dir,
|
||||
}
|
||||
}
|
||||
|
||||
// settingsPath returns the path to settings.json.
|
||||
// Uses PluginsDir's parent (~/.dws/) for production, PluginsDir itself for tests.
|
||||
func (l *Loader) settingsPath() string {
|
||||
// If PluginsDir ends with "plugins", go up one level to ~/.dws/
|
||||
if filepath.Base(l.PluginsDir) == "plugins" {
|
||||
return filepath.Join(filepath.Dir(l.PluginsDir), "settings.json")
|
||||
}
|
||||
// For test temp dirs, use PluginsDir directly
|
||||
return filepath.Join(l.PluginsDir, "settings.json")
|
||||
}
|
||||
|
||||
// loadSettings reads settings.json from the parent of PluginsDir.
|
||||
func (l *Loader) loadSettings() *Settings {
|
||||
settingsPath := l.settingsPath()
|
||||
data, err := os.ReadFile(settingsPath)
|
||||
if err != nil {
|
||||
return &Settings{}
|
||||
}
|
||||
var s Settings
|
||||
if err := json.Unmarshal(data, &s); err != nil {
|
||||
slog.Debug("plugin: failed to parse settings.json", "error", err)
|
||||
return &Settings{}
|
||||
}
|
||||
return &s
|
||||
}
|
||||
|
||||
func isPluginEnabled(s *Settings, name string) bool {
|
||||
if s == nil || s.EnabledPlugins == nil {
|
||||
return true // default: enabled
|
||||
}
|
||||
enabled, exists := s.EnabledPlugins[name]
|
||||
if !exists {
|
||||
return true // not in list = enabled
|
||||
}
|
||||
return enabled
|
||||
}
|
||||
|
||||
// InstalledPlugins returns the list of all installed plugins with their
|
||||
// status info. Used by `dws plugin list`.
|
||||
type PluginInfo struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Type string `json:"type"` // "user" or "dev"
|
||||
Enabled bool `json:"enabled"`
|
||||
Path string `json:"path"`
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
|
||||
// ListInstalled returns info about all installed plugins.
|
||||
func (l *Loader) ListInstalled() []PluginInfo {
|
||||
var result []PluginInfo
|
||||
settings := l.loadSettings()
|
||||
|
||||
// User plugins
|
||||
userDir := filepath.Join(l.PluginsDir, "user")
|
||||
if entries, err := os.ReadDir(userDir); err == nil {
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() {
|
||||
continue
|
||||
}
|
||||
l.collectUserPluginInfos(filepath.Join(userDir, entry.Name()), entry.Name(), settings, &result)
|
||||
}
|
||||
}
|
||||
|
||||
// Dev plugins
|
||||
for name, dir := range settings.DevPlugins {
|
||||
m, err := ParseManifest(filepath.Join(dir, "plugin.json"))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
result = append(result, PluginInfo{
|
||||
Name: name,
|
||||
Version: m.Version,
|
||||
Type: "dev",
|
||||
Enabled: true,
|
||||
Path: dir,
|
||||
Description: m.Description,
|
||||
})
|
||||
}
|
||||
|
||||
return result
|
||||
}
|
||||
|
||||
func (l *Loader) collectUserPluginInfos(dir, prefix string, settings *Settings, result *[]PluginInfo) {
|
||||
// Direct plugin
|
||||
if m, err := ParseManifest(filepath.Join(dir, "plugin.json")); err == nil {
|
||||
qualName := prefix
|
||||
*result = append(*result, PluginInfo{
|
||||
Name: qualName,
|
||||
Version: m.Version,
|
||||
Type: "user",
|
||||
Enabled: isPluginEnabled(settings, qualName),
|
||||
Path: dir,
|
||||
Description: m.Description,
|
||||
})
|
||||
return
|
||||
}
|
||||
// Workspace: dir is a workspace, iterate sub-plugins
|
||||
subEntries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
for _, sub := range subEntries {
|
||||
if !sub.IsDir() {
|
||||
continue
|
||||
}
|
||||
subDir := filepath.Join(dir, sub.Name())
|
||||
m, err := ParseManifest(filepath.Join(subDir, "plugin.json"))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
qualName := prefix + "/" + m.Name
|
||||
*result = append(*result, PluginInfo{
|
||||
Name: qualName,
|
||||
Version: m.Version,
|
||||
Type: "user",
|
||||
Enabled: isPluginEnabled(settings, qualName),
|
||||
Path: subDir,
|
||||
Description: m.Description,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// InstallFromDir copies a plugin from a source directory to the user
|
||||
// plugins directory.
|
||||
func (l *Loader) InstallFromDir(srcDir string) (*Plugin, error) {
|
||||
manifestPath := filepath.Join(srcDir, "plugin.json")
|
||||
manifest, err := ParseManifest(manifestPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid plugin: %w", err)
|
||||
}
|
||||
if err := manifest.Validate(l.CLIVersion); err != nil {
|
||||
return nil, fmt.Errorf("plugin validation failed: %w", err)
|
||||
}
|
||||
|
||||
destDir := filepath.Join(l.PluginsDir, "user", manifest.Name)
|
||||
if err := copyDir(srcDir, destDir); err != nil {
|
||||
return nil, fmt.Errorf("install failed: %w", err)
|
||||
}
|
||||
|
||||
// Remove stale files in destDir that no longer exist in srcDir.
|
||||
removeStaleFiles(srcDir, destDir)
|
||||
|
||||
// Run build if configured (compile server to binary).
|
||||
if manifest.Build != nil {
|
||||
if err := runBuild(destDir, manifest.Build); err != nil {
|
||||
// Clean up on build failure.
|
||||
_ = os.RemoveAll(destDir)
|
||||
return nil, fmt.Errorf("plugin build failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Enable by default in settings
|
||||
l.setPluginEnabled(manifest.Name, true)
|
||||
|
||||
return &Plugin{
|
||||
Manifest: *manifest,
|
||||
Root: destDir,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// InstallFromGit clones a git repository and installs the plugin.
|
||||
// The workspace is extracted from the git URL (e.g. github.com/{workspace}/{name}).
|
||||
func (l *Loader) InstallFromGit(gitURL string) (*Plugin, error) {
|
||||
workspace, repoName, err := parseGitURL(gitURL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid git URL: %w", err)
|
||||
}
|
||||
|
||||
// Clone to temp directory.
|
||||
tmpDir, err := os.MkdirTemp("", "dws-plugin-git-*")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create temp dir: %w", err)
|
||||
}
|
||||
defer os.RemoveAll(tmpDir)
|
||||
|
||||
cloneDir := filepath.Join(tmpDir, repoName)
|
||||
cmd := exec.Command("git", "clone", "--depth", "1", gitURL, cloneDir)
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
if err := cmd.Run(); err != nil {
|
||||
return nil, fmt.Errorf("git clone failed: %w", err)
|
||||
}
|
||||
|
||||
// Parse and validate manifest.
|
||||
manifest, err := ParseManifest(filepath.Join(cloneDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid plugin: %w", err)
|
||||
}
|
||||
if err := manifest.Validate(l.CLIVersion); err != nil {
|
||||
return nil, fmt.Errorf("plugin validation failed: %w", err)
|
||||
}
|
||||
|
||||
// All plugins install to the user directory with workspace nesting:
|
||||
// ~/.dws/plugins/user/{workspace}/{name}/. There is no privileged
|
||||
// workspace — every plugin is third-party.
|
||||
destDir := filepath.Join(l.PluginsDir, config.PluginUserDir, workspace, manifest.Name)
|
||||
|
||||
// Remove .git directory before copying.
|
||||
_ = os.RemoveAll(filepath.Join(cloneDir, ".git"))
|
||||
|
||||
if err := copyDir(cloneDir, destDir); err != nil {
|
||||
return nil, fmt.Errorf("install failed: %w", err)
|
||||
}
|
||||
|
||||
// Run build if configured (compile server to binary).
|
||||
if manifest.Build != nil {
|
||||
if err := runBuild(destDir, manifest.Build); err != nil {
|
||||
// Clean up on build failure.
|
||||
_ = os.RemoveAll(destDir)
|
||||
return nil, fmt.Errorf("plugin build failed: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
qualifiedName := workspace + "/" + manifest.Name
|
||||
l.setPluginEnabled(qualifiedName, true)
|
||||
|
||||
return &Plugin{
|
||||
Manifest: *manifest,
|
||||
Root: destDir,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// parseGitURL extracts workspace and repo name from a git URL.
|
||||
// Supports: https://github.com/org/repo.git, git@github.com:org/repo.git
|
||||
// Rejects file:// and other local protocols to prevent reading local files.
|
||||
func parseGitURL(gitURL string) (workspace, repoName string, err error) {
|
||||
gitURL = strings.TrimSpace(gitURL)
|
||||
|
||||
// Reject dangerous protocols that could read local files.
|
||||
lower := strings.ToLower(gitURL)
|
||||
if strings.HasPrefix(lower, "file://") || strings.HasPrefix(lower, "/") || strings.HasPrefix(lower, ".") {
|
||||
return "", "", fmt.Errorf("local paths and file:// URLs are not allowed: %q", gitURL)
|
||||
}
|
||||
|
||||
// Handle SSH format: git@github.com:org/repo.git
|
||||
if strings.HasPrefix(gitURL, "git@") {
|
||||
parts := strings.SplitN(gitURL, ":", 2)
|
||||
if len(parts) != 2 {
|
||||
return "", "", fmt.Errorf("cannot parse SSH URL %q", gitURL)
|
||||
}
|
||||
path := strings.TrimSuffix(parts[1], ".git")
|
||||
segments := strings.Split(path, "/")
|
||||
if len(segments) < 2 {
|
||||
return "", "", fmt.Errorf("SSH URL %q must have org/repo format", gitURL)
|
||||
}
|
||||
return segments[len(segments)-2], segments[len(segments)-1], nil
|
||||
}
|
||||
|
||||
// Handle HTTPS format.
|
||||
u, err := url.Parse(gitURL)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("cannot parse URL %q: %w", gitURL, err)
|
||||
}
|
||||
|
||||
// Only allow https:// and http:// schemes.
|
||||
if u.Scheme != "https" && u.Scheme != "http" {
|
||||
return "", "", fmt.Errorf("unsupported URL scheme %q: only https and ssh are allowed", u.Scheme)
|
||||
}
|
||||
|
||||
path := strings.TrimSuffix(strings.Trim(u.Path, "/"), ".git")
|
||||
segments := strings.Split(path, "/")
|
||||
if len(segments) < 2 {
|
||||
return "", "", fmt.Errorf("URL %q must have org/repo format", gitURL)
|
||||
}
|
||||
|
||||
return segments[len(segments)-2], segments[len(segments)-1], nil
|
||||
}
|
||||
|
||||
// RemovePlugin removes an installed plugin by name.
|
||||
func (l *Loader) RemovePlugin(name string, keepData bool) error {
|
||||
pluginDir := l.findUserPluginDir(name)
|
||||
if pluginDir == "" {
|
||||
return fmt.Errorf("plugin %q not found", name)
|
||||
}
|
||||
|
||||
if err := os.RemoveAll(pluginDir); err != nil {
|
||||
return fmt.Errorf("failed to remove plugin: %w", err)
|
||||
}
|
||||
|
||||
if !keepData {
|
||||
dataDir := filepath.Join(l.PluginsDir, config.PluginDataDir, name)
|
||||
_ = os.RemoveAll(dataDir)
|
||||
}
|
||||
|
||||
l.purgePluginFromSettings(name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// purgePluginFromSettings removes all traces of a plugin from settings.json:
|
||||
// its enabled flag and any persisted pluginConfigs entry. Called after
|
||||
// RemovePlugin succeeds so settings.json does not retain dangling state for
|
||||
// plugins that no longer exist on disk.
|
||||
func (l *Loader) purgePluginFromSettings(name string) {
|
||||
settings := l.loadSettings()
|
||||
changed := false
|
||||
if _, ok := settings.EnabledPlugins[name]; ok {
|
||||
delete(settings.EnabledPlugins, name)
|
||||
changed = true
|
||||
}
|
||||
if _, ok := settings.PluginConfigs[name]; ok {
|
||||
delete(settings.PluginConfigs, name)
|
||||
changed = true
|
||||
}
|
||||
if !changed {
|
||||
return
|
||||
}
|
||||
l.saveSettings(settings)
|
||||
}
|
||||
|
||||
// SetEnabled enables or disables a plugin in settings.json.
|
||||
func (l *Loader) SetEnabled(name string, enabled bool) error {
|
||||
if l.findUserPluginDir(name) == "" {
|
||||
return fmt.Errorf("plugin %q not found", name)
|
||||
}
|
||||
l.setPluginEnabled(name, enabled)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *Loader) findUserPluginDir(name string) string {
|
||||
// Try direct: user/{name}/
|
||||
dir := filepath.Join(l.PluginsDir, "user", name)
|
||||
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err == nil {
|
||||
return dir
|
||||
}
|
||||
// Try workspace: user/{workspace}/{plugin}/
|
||||
parts := strings.SplitN(name, "/", 2)
|
||||
if len(parts) == 2 {
|
||||
dir = filepath.Join(l.PluginsDir, "user", parts[0], parts[1])
|
||||
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err == nil {
|
||||
return dir
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func (l *Loader) setPluginEnabled(name string, enabled bool) {
|
||||
settings := l.loadSettings()
|
||||
if settings.EnabledPlugins == nil {
|
||||
settings.EnabledPlugins = make(map[string]bool)
|
||||
}
|
||||
settings.EnabledPlugins[name] = enabled
|
||||
l.saveSettings(settings)
|
||||
}
|
||||
|
||||
func (l *Loader) saveSettings(s *Settings) {
|
||||
settingsPath := l.settingsPath()
|
||||
data, err := json.MarshalIndent(s, "", " ")
|
||||
if err != nil {
|
||||
slog.Debug("plugin: failed to marshal settings", "error", err)
|
||||
return
|
||||
}
|
||||
_ = os.MkdirAll(filepath.Dir(settingsPath), 0o700)
|
||||
_ = os.WriteFile(settingsPath, data, 0o600)
|
||||
}
|
||||
|
||||
// GetPluginConfig returns the value of a config key for a plugin.
|
||||
// It checks pluginConfigs in settings.json first, then falls back to
|
||||
// the userConfig default in the plugin's manifest.
|
||||
func (l *Loader) GetPluginConfig(pluginName, key string) (string, bool) {
|
||||
settings := l.loadSettings()
|
||||
if settings.PluginConfigs != nil {
|
||||
if pluginCfg, ok := settings.PluginConfigs[pluginName]; ok {
|
||||
if val, ok := pluginCfg[key]; ok {
|
||||
if s, ok := val.(string); ok {
|
||||
return s, true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// SetPluginConfig persists a config key-value pair for a plugin.
|
||||
func (l *Loader) SetPluginConfig(pluginName, key, value string) {
|
||||
settings := l.loadSettings()
|
||||
if settings.PluginConfigs == nil {
|
||||
settings.PluginConfigs = make(map[string]map[string]any)
|
||||
}
|
||||
if settings.PluginConfigs[pluginName] == nil {
|
||||
settings.PluginConfigs[pluginName] = make(map[string]any)
|
||||
}
|
||||
settings.PluginConfigs[pluginName][key] = value
|
||||
l.saveSettings(settings)
|
||||
}
|
||||
|
||||
// UnsetPluginConfig removes a config key for a plugin.
|
||||
func (l *Loader) UnsetPluginConfig(pluginName, key string) bool {
|
||||
settings := l.loadSettings()
|
||||
if settings.PluginConfigs == nil {
|
||||
return false
|
||||
}
|
||||
pluginCfg, ok := settings.PluginConfigs[pluginName]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if _, exists := pluginCfg[key]; !exists {
|
||||
return false
|
||||
}
|
||||
delete(pluginCfg, key)
|
||||
if len(pluginCfg) == 0 {
|
||||
delete(settings.PluginConfigs, pluginName)
|
||||
}
|
||||
l.saveSettings(settings)
|
||||
return true
|
||||
}
|
||||
|
||||
// ListPluginConfig returns all config key-value pairs for a plugin.
|
||||
func (l *Loader) ListPluginConfig(pluginName string) map[string]string {
|
||||
settings := l.loadSettings()
|
||||
result := make(map[string]string)
|
||||
if settings.PluginConfigs == nil {
|
||||
return result
|
||||
}
|
||||
pluginCfg, ok := settings.PluginConfigs[pluginName]
|
||||
if !ok {
|
||||
return result
|
||||
}
|
||||
for k, v := range pluginCfg {
|
||||
if s, ok := v.(string); ok {
|
||||
result[k] = s
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// InjectPluginConfigEnv reads pluginConfigs from settings.json and sets
|
||||
// environment variables for each configured key. This allows
|
||||
// expandPluginVars (which calls os.Expand) to resolve ${KEY} references
|
||||
// in plugin.json headers, endpoints, etc.
|
||||
//
|
||||
// Environment variables already set by the user take precedence — only
|
||||
// keys not already present in the environment are injected.
|
||||
// dangerousEnvVars contains environment variable names that must never be
|
||||
// set from plugin config because they can alter process behavior in
|
||||
// security-critical ways (library injection, executable search path, etc.).
|
||||
var dangerousEnvVars = map[string]bool{
|
||||
"PATH": true, "HOME": true, "USER": true, "SHELL": true,
|
||||
"LD_PRELOAD": true, "LD_LIBRARY_PATH": true,
|
||||
"DYLD_INSERT_LIBRARIES": true, "DYLD_LIBRARY_PATH": true, "DYLD_FRAMEWORK_PATH": true,
|
||||
"NODE_OPTIONS": true, "PYTHONPATH": true, "RUBYLIB": true,
|
||||
"GOPATH": true, "GOROOT": true,
|
||||
"HTTP_PROXY": true, "HTTPS_PROXY": true, "ALL_PROXY": true, "NO_PROXY": true,
|
||||
"http_proxy": true, "https_proxy": true, "all_proxy": true, "no_proxy": true,
|
||||
}
|
||||
|
||||
func (l *Loader) InjectPluginConfigEnv() {
|
||||
settings := l.loadSettings()
|
||||
if len(settings.PluginConfigs) == 0 {
|
||||
return
|
||||
}
|
||||
for _, pluginCfg := range settings.PluginConfigs {
|
||||
for key, val := range pluginCfg {
|
||||
strVal, ok := val.(string)
|
||||
if !ok || strVal == "" {
|
||||
continue
|
||||
}
|
||||
// Block dangerous environment variable names.
|
||||
if dangerousEnvVars[key] {
|
||||
slog.Warn("plugin: blocked dangerous env var from config",
|
||||
"key", key)
|
||||
continue
|
||||
}
|
||||
// Do not override existing environment variables.
|
||||
if _, exists := os.LookupEnv(key); exists {
|
||||
continue
|
||||
}
|
||||
_ = os.Setenv(key, strVal)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// LoadDev loads dev plugins registered via `dws plugin dev`.
|
||||
// Dev plugins are loaded from their source directories without copying.
|
||||
func (l *Loader) LoadDev() []*Plugin {
|
||||
settings := l.loadSettings()
|
||||
if len(settings.DevPlugins) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
var plugins []*Plugin
|
||||
for name, dir := range settings.DevPlugins {
|
||||
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err != nil {
|
||||
slog.Debug("plugin: dev plugin directory missing, skipping",
|
||||
"name", name, "dir", dir)
|
||||
continue
|
||||
}
|
||||
p := l.loadPlugin(dir)
|
||||
if p != nil {
|
||||
plugins = append(plugins, p)
|
||||
slog.Debug("plugin: loaded dev plugin", "name", name, "dir", dir)
|
||||
}
|
||||
}
|
||||
return plugins
|
||||
}
|
||||
|
||||
// RegisterDevPlugin registers a source directory as a dev plugin.
|
||||
func (l *Loader) RegisterDevPlugin(name, absDir string) error {
|
||||
settings := l.loadSettings()
|
||||
if settings.DevPlugins == nil {
|
||||
settings.DevPlugins = make(map[string]string)
|
||||
}
|
||||
settings.DevPlugins[name] = absDir
|
||||
l.saveSettings(settings)
|
||||
return nil
|
||||
}
|
||||
|
||||
// UnregisterDevPlugin removes a dev plugin registration.
|
||||
func (l *Loader) UnregisterDevPlugin(name string) error {
|
||||
settings := l.loadSettings()
|
||||
if settings.DevPlugins == nil || settings.DevPlugins[name] == "" {
|
||||
return fmt.Errorf("dev plugin %q is not registered", name)
|
||||
}
|
||||
delete(settings.DevPlugins, name)
|
||||
l.saveSettings(settings)
|
||||
return nil
|
||||
}
|
||||
|
||||
// SyncSkills copies plugin SKILL.md files into all detected agent
|
||||
// skill directories (e.g. ~/.claude/skills/dws/, ~/.cursor/skills/dws/).
|
||||
// This makes plugin skills available to AI agents without CLI releases.
|
||||
func SyncSkills(plugins []*Plugin) {
|
||||
if len(plugins) == 0 {
|
||||
return
|
||||
}
|
||||
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
slog.Debug("plugin: cannot get home dir for skill sync", "error", err)
|
||||
return
|
||||
}
|
||||
|
||||
// Known agent skill directories (subset of upgrade/paths.go knownSkillDirs).
|
||||
agentDirs := []string{
|
||||
".agents/skills",
|
||||
".claude/skills",
|
||||
".cursor/skills",
|
||||
".qoder/skills",
|
||||
".codex/skills",
|
||||
}
|
||||
|
||||
for _, p := range plugins {
|
||||
skillsDir := p.SkillsDir()
|
||||
if _, err := os.Stat(skillsDir); err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Walk the plugin's skills directory and copy files to each agent dir.
|
||||
entries, err := os.ReadDir(skillsDir)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
|
||||
for _, agentDir := range agentDirs {
|
||||
agentBase := filepath.Join(homeDir, agentDir)
|
||||
// Only sync to agents that are actually installed (parent dir exists).
|
||||
parentGate := filepath.Dir(agentBase)
|
||||
if _, err := os.Stat(parentGate); os.IsNotExist(err) {
|
||||
continue
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
src := filepath.Join(skillsDir, entry.Name())
|
||||
// Place plugin skills under dws/plugins/{plugin-name}/
|
||||
dest := filepath.Join(agentBase, "dws", "plugins", p.Manifest.Name, entry.Name())
|
||||
if entry.IsDir() {
|
||||
_ = copyDir(src, dest)
|
||||
} else {
|
||||
_ = os.MkdirAll(filepath.Dir(dest), 0o755)
|
||||
data, readErr := os.ReadFile(src)
|
||||
if readErr == nil {
|
||||
_ = os.WriteFile(dest, data, 0o644)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
slog.Debug("plugin: skill sync completed", "plugins", len(plugins))
|
||||
}
|
||||
|
||||
// BuildPlugin runs the build command declared in plugin.json.
|
||||
// It compiles the plugin's stdio server into a native binary so that
|
||||
// users don't need language runtimes. Returns nil if no build is configured.
|
||||
func BuildPlugin(pluginDir string) error {
|
||||
manifest, err := ParseManifest(filepath.Join(pluginDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return fmt.Errorf("parse manifest: %w", err)
|
||||
}
|
||||
if manifest.Build == nil {
|
||||
return nil // no build configured
|
||||
}
|
||||
return runBuild(pluginDir, manifest.Build)
|
||||
}
|
||||
|
||||
// runBuild executes the build command and verifies the output exists.
|
||||
func runBuild(pluginDir string, build *BuildConfig) error {
|
||||
if build.Command == "" {
|
||||
return fmt.Errorf("build.command is empty")
|
||||
}
|
||||
|
||||
// Validate build.output is a relative path within the plugin directory.
|
||||
if build.Output != "" {
|
||||
if filepath.IsAbs(build.Output) {
|
||||
return fmt.Errorf("build.output must be a relative path, got %q", build.Output)
|
||||
}
|
||||
cleanOut := filepath.Clean(build.Output)
|
||||
if strings.HasPrefix(cleanOut, "..") {
|
||||
return fmt.Errorf("build.output must not escape plugin directory: %q", build.Output)
|
||||
}
|
||||
}
|
||||
|
||||
slog.Info("plugin: building", "dir", pluginDir, "command", build.Command)
|
||||
|
||||
var cmd *exec.Cmd
|
||||
if runtime.GOOS == "windows" {
|
||||
cmd = exec.Command("cmd", "/C", build.Command)
|
||||
} else {
|
||||
cmd = exec.Command("sh", "-c", build.Command)
|
||||
}
|
||||
cmd.Dir = pluginDir
|
||||
cmd.Stdout = os.Stdout
|
||||
cmd.Stderr = os.Stderr
|
||||
|
||||
// Pass through environment + plugin root
|
||||
cmd.Env = append(os.Environ(), "DWS_PLUGIN_ROOT="+pluginDir)
|
||||
|
||||
if err := cmd.Run(); err != nil {
|
||||
return fmt.Errorf("build failed: %w", err)
|
||||
}
|
||||
|
||||
// Verify output binary exists
|
||||
if build.Output != "" {
|
||||
outPath := filepath.Join(pluginDir, build.Output)
|
||||
info, err := os.Stat(outPath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("build output not found at %s: %w", build.Output, err)
|
||||
}
|
||||
// Ensure the output is executable
|
||||
if info.Mode()&0o111 == 0 {
|
||||
_ = os.Chmod(outPath, info.Mode()|0o755)
|
||||
}
|
||||
}
|
||||
|
||||
slog.Info("plugin: build succeeded", "output", build.Output)
|
||||
return nil
|
||||
}
|
||||
|
||||
// copyDir recursively copies src to dst, skipping files whose content
|
||||
// is identical to the destination. This avoids overwriting locked
|
||||
// executables (e.g. a running stdio plugin on Windows).
|
||||
// Symlinks are skipped for security (prevents path traversal attacks).
|
||||
func copyDir(src, dst string) error {
|
||||
cleanDst := filepath.Clean(dst) + string(os.PathSeparator)
|
||||
if err := os.MkdirAll(dst, 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
return filepath.Walk(src, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Skip symlinks to prevent path traversal.
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
return nil
|
||||
}
|
||||
rel, err := filepath.Rel(src, path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
target := filepath.Join(dst, rel)
|
||||
// Guard against path traversal via crafted relative paths.
|
||||
if target != cleanDst[:len(cleanDst)-1] && !strings.HasPrefix(target, cleanDst) {
|
||||
return fmt.Errorf("path traversal detected: %s", rel)
|
||||
}
|
||||
if info.IsDir() {
|
||||
return os.MkdirAll(target, info.Mode())
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// Skip if destination already has identical content (cheap size check first).
|
||||
if targetInfo, statErr := os.Stat(target); statErr == nil && targetInfo.Size() == int64(len(data)) {
|
||||
if existing, readErr := os.ReadFile(target); readErr == nil && bytes.Equal(existing, data) {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return os.WriteFile(target, data, info.Mode())
|
||||
})
|
||||
}
|
||||
|
||||
// removeStaleFiles deletes files under dst that do not exist in src.
|
||||
// Best-effort: errors are logged but do not fail the install.
|
||||
func removeStaleFiles(src, dst string) {
|
||||
srcSet := make(map[string]struct{})
|
||||
_ = filepath.Walk(src, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
rel, relErr := filepath.Rel(src, path)
|
||||
if relErr != nil {
|
||||
return nil
|
||||
}
|
||||
srcSet[rel] = struct{}{}
|
||||
return nil
|
||||
})
|
||||
|
||||
_ = filepath.Walk(dst, func(path string, info os.FileInfo, err error) error {
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
rel, relErr := filepath.Rel(dst, path)
|
||||
if relErr != nil {
|
||||
return nil
|
||||
}
|
||||
if rel == "." {
|
||||
return nil
|
||||
}
|
||||
if _, exists := srcSet[rel]; !exists {
|
||||
if info.IsDir() {
|
||||
_ = os.RemoveAll(path)
|
||||
return filepath.SkipDir
|
||||
}
|
||||
if removeErr := os.Remove(path); removeErr != nil {
|
||||
slog.Debug("plugin: failed to remove stale file", "path", path, "error", removeErr)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// 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 plugin
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSetAndGetPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Initially empty.
|
||||
val, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if ok {
|
||||
t.Errorf("expected not found, got %q", val)
|
||||
}
|
||||
|
||||
// Set a value.
|
||||
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
|
||||
|
||||
// Read it back.
|
||||
val, ok = loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if !ok {
|
||||
t.Fatal("expected to find config after set")
|
||||
}
|
||||
if val != "sk-test-12345" {
|
||||
t.Errorf("got %q, want sk-test-12345", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetPluginConfigMultipleKeys(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("my-plugin", "API_KEY", "key-1")
|
||||
loader.SetPluginConfig("my-plugin", "API_ENDPOINT", "https://example.com")
|
||||
loader.SetPluginConfig("other-plugin", "TOKEN", "tok-abc")
|
||||
|
||||
val, ok := loader.GetPluginConfig("my-plugin", "API_KEY")
|
||||
if !ok || val != "key-1" {
|
||||
t.Errorf("API_KEY = %q (ok=%v), want key-1", val, ok)
|
||||
}
|
||||
|
||||
val, ok = loader.GetPluginConfig("my-plugin", "API_ENDPOINT")
|
||||
if !ok || val != "https://example.com" {
|
||||
t.Errorf("API_ENDPOINT = %q (ok=%v), want https://example.com", val, ok)
|
||||
}
|
||||
|
||||
val, ok = loader.GetPluginConfig("other-plugin", "TOKEN")
|
||||
if !ok || val != "tok-abc" {
|
||||
t.Errorf("TOKEN = %q (ok=%v), want tok-abc", val, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsetPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Unset on empty returns false.
|
||||
if loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
|
||||
t.Error("expected false for unset on empty config")
|
||||
}
|
||||
|
||||
// Set then unset.
|
||||
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
|
||||
if !loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
|
||||
t.Error("expected true for unset of existing key")
|
||||
}
|
||||
|
||||
// Verify it's gone.
|
||||
_, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if ok {
|
||||
t.Error("expected not found after unset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsetPluginConfigCleansEmptyMap(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", "KEY1", "val1")
|
||||
loader.UnsetPluginConfig("demo-devtool", "KEY1")
|
||||
|
||||
// After removing the last key, the plugin entry should be cleaned up.
|
||||
configs := loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 0 {
|
||||
t.Errorf("expected empty config map after removing last key, got %v", configs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Empty list.
|
||||
configs := loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 0 {
|
||||
t.Errorf("expected empty, got %v", configs)
|
||||
}
|
||||
|
||||
// Set some values.
|
||||
loader.SetPluginConfig("demo-devtool", "KEY_A", "val-a")
|
||||
loader.SetPluginConfig("demo-devtool", "KEY_B", "val-b")
|
||||
|
||||
configs = loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 2 {
|
||||
t.Fatalf("expected 2 configs, got %d", len(configs))
|
||||
}
|
||||
if configs["KEY_A"] != "val-a" {
|
||||
t.Errorf("KEY_A = %q, want val-a", configs["KEY_A"])
|
||||
}
|
||||
if configs["KEY_B"] != "val-b" {
|
||||
t.Errorf("KEY_B = %q, want val-b", configs["KEY_B"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectPluginConfigEnv(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Use a unique env var name to avoid test pollution.
|
||||
envKey := "DWS_TEST_INJECT_CONFIG_" + t.Name()
|
||||
t.Cleanup(func() { os.Unsetenv(envKey) })
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", envKey, "injected-value")
|
||||
|
||||
// Ensure it's not already set.
|
||||
os.Unsetenv(envKey)
|
||||
|
||||
loader.InjectPluginConfigEnv()
|
||||
|
||||
got := os.Getenv(envKey)
|
||||
if got != "injected-value" {
|
||||
t.Errorf("env %s = %q, want injected-value", envKey, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectPluginConfigEnvDoesNotOverride(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
envKey := "DWS_TEST_INJECT_NOOVERRIDE_" + t.Name()
|
||||
t.Cleanup(func() { os.Unsetenv(envKey) })
|
||||
|
||||
// Pre-set the env var.
|
||||
os.Setenv(envKey, "user-value")
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", envKey, "config-value")
|
||||
loader.InjectPluginConfigEnv()
|
||||
|
||||
got := os.Getenv(envKey)
|
||||
if got != "user-value" {
|
||||
t.Errorf("env %s = %q, want user-value (should not be overridden)", envKey, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetPluginConfigOverwritesExisting(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", "KEY", "old-value")
|
||||
loader.SetPluginConfig("demo-devtool", "KEY", "new-value")
|
||||
|
||||
val, ok := loader.GetPluginConfig("demo-devtool", "KEY")
|
||||
if !ok || val != "new-value" {
|
||||
t.Errorf("got %q (ok=%v), want new-value", val, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPluginConfigWrongPlugin(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("plugin-a", "KEY", "value")
|
||||
|
||||
_, ok := loader.GetPluginConfig("plugin-b", "KEY")
|
||||
if ok {
|
||||
t.Error("expected not found for different plugin name")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,257 @@
|
||||
// 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 plugin implements the DWS CLI plugin system. It loads,
|
||||
// validates, and injects plugin capabilities (MCP servers, skills,
|
||||
// pipeline hooks) into the existing CLI infrastructure.
|
||||
package plugin
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// namePattern validates plugin names: lowercase kebab-case, 3–50 chars.
|
||||
var namePattern = regexp.MustCompile(`^[a-z][a-z0-9-]{2,49}$`)
|
||||
|
||||
// Manifest represents the parsed contents of a plugin.json file.
|
||||
type Manifest struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Type string `json:"type,omitempty"` // "managed" or "user"
|
||||
MinCLIVersion string `json:"minCLIVersion,omitempty"`
|
||||
MCPServers map[string]*MCPServer `json:"mcpServers,omitempty"`
|
||||
Skills string `json:"skills,omitempty"`
|
||||
Hooks string `json:"hooks,omitempty"`
|
||||
Permissions []string `json:"permissions,omitempty"`
|
||||
UserConfig map[string]ConfigItem `json:"userConfig,omitempty"`
|
||||
Build *BuildConfig `json:"build,omitempty"`
|
||||
}
|
||||
|
||||
// BuildConfig declares how to compile the plugin's stdio server into
|
||||
// a native binary. DWS runs this automatically during install so that
|
||||
// plugin users never need language runtimes or dependency managers.
|
||||
type BuildConfig struct {
|
||||
// Command is the shell command to compile the server.
|
||||
// Executed via "sh -c" in the plugin root directory.
|
||||
// Examples: "bun build --compile src/server.ts --outfile bin/server"
|
||||
// "go build -o bin/server ./cmd/server"
|
||||
// "pip install pyinstaller && pyinstaller --onefile src/server.py -n server --distpath bin/"
|
||||
Command string `json:"command"`
|
||||
|
||||
// Output is the path to the compiled binary, relative to the plugin root.
|
||||
// Used to verify the build succeeded. Example: "bin/server"
|
||||
Output string `json:"output"`
|
||||
}
|
||||
|
||||
// MCPServer describes a single MCP server declared by a plugin.
|
||||
type MCPServer struct {
|
||||
Type string `json:"type"` // "streamable-http" or "stdio"
|
||||
Endpoint string `json:"endpoint,omitempty"` // required for streamable-http
|
||||
Command string `json:"command,omitempty"` // required for stdio
|
||||
Args []string `json:"args,omitempty"`
|
||||
Env map[string]string `json:"env,omitempty"`
|
||||
Headers map[string]string `json:"headers,omitempty"` // custom HTTP headers (e.g. Authorization for third-party APIs)
|
||||
CLI json.RawMessage `json:"cli,omitempty"` // CLIOverlay, passed through
|
||||
}
|
||||
|
||||
// ConfigItem describes a user-configurable setting for a plugin.
|
||||
type ConfigItem struct {
|
||||
Description string `json:"description,omitempty"`
|
||||
Default string `json:"default,omitempty"`
|
||||
Sensitive bool `json:"sensitive,omitempty"`
|
||||
}
|
||||
|
||||
// HooksConfig describes pipeline hooks declared in a hooks.json file.
|
||||
type HooksConfig struct {
|
||||
Hooks []HookEntry `json:"hooks"`
|
||||
}
|
||||
|
||||
// HookEntry describes a single pipeline hook.
|
||||
type HookEntry struct {
|
||||
Phase string `json:"phase"` // "pre-request", "post-response", etc.
|
||||
Matcher string `json:"matcher,omitempty"` // glob pattern, e.g. "conference.*"
|
||||
Command string `json:"command"` // shell command to execute
|
||||
Timeout int `json:"timeout,omitempty"` // seconds, default 30
|
||||
}
|
||||
|
||||
// Plugin is a loaded, validated plugin ready for injection.
|
||||
type Plugin struct {
|
||||
Manifest Manifest
|
||||
Root string // absolute path to plugin directory
|
||||
}
|
||||
|
||||
// ParseManifest reads and parses a plugin.json file.
|
||||
func ParseManifest(path string) (*Manifest, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read plugin.json: %w", err)
|
||||
}
|
||||
var m Manifest
|
||||
if err := json.Unmarshal(data, &m); err != nil {
|
||||
return nil, fmt.Errorf("parse plugin.json: %w", err)
|
||||
}
|
||||
return &m, nil
|
||||
}
|
||||
|
||||
// Validate checks that a manifest is well-formed. It returns an error
|
||||
// describing the first problem found, or nil if the manifest is valid.
|
||||
// cliVersion is the current CLI version string for compatibility checks.
|
||||
func (m *Manifest) Validate(cliVersion string) error {
|
||||
if !namePattern.MatchString(m.Name) {
|
||||
return fmt.Errorf("invalid plugin name %q: must be lowercase kebab-case, 3–50 chars", m.Name)
|
||||
}
|
||||
if !isValidSemver(m.Version) {
|
||||
return fmt.Errorf("invalid plugin version %q: must be valid semver (e.g. 1.0.0)", m.Version)
|
||||
}
|
||||
if m.Type != "" && m.Type != "managed" && m.Type != "user" {
|
||||
return fmt.Errorf("invalid plugin type %q: must be \"managed\" or \"user\"", m.Type)
|
||||
}
|
||||
if m.MinCLIVersion != "" && cliVersion != "" && cliVersion != "dev" {
|
||||
if compareSemver(cliVersion, m.MinCLIVersion) < 0 {
|
||||
return fmt.Errorf("plugin requires CLI >= %s, current is %s", m.MinCLIVersion, cliVersion)
|
||||
}
|
||||
}
|
||||
for key, srv := range m.MCPServers {
|
||||
if err := validateMCPServer(key, srv); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if m.Skills != "" {
|
||||
if err := validateSafePath(m.Skills); err != nil {
|
||||
return fmt.Errorf("skills path: %w", err)
|
||||
}
|
||||
}
|
||||
if m.Hooks != "" {
|
||||
if err := validateSafePath(m.Hooks); err != nil {
|
||||
return fmt.Errorf("hooks path: %w", err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateMCPServer(key string, srv *MCPServer) error {
|
||||
switch srv.Type {
|
||||
case "streamable-http":
|
||||
if strings.TrimSpace(srv.Endpoint) == "" {
|
||||
return fmt.Errorf("mcpServers[%q]: streamable-http requires endpoint", key)
|
||||
}
|
||||
case "stdio":
|
||||
if strings.TrimSpace(srv.Command) == "" {
|
||||
return fmt.Errorf("mcpServers[%q]: stdio requires command", key)
|
||||
}
|
||||
// Reject absolute paths in command to encourage relative paths within plugin root.
|
||||
if filepath.IsAbs(srv.Command) {
|
||||
return fmt.Errorf("mcpServers[%q]: command must be a relative path, got %q", key, srv.Command)
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("mcpServers[%q]: unsupported type %q (must be streamable-http or stdio)", key, srv.Type)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateSafePath rejects paths containing ".." traversal.
|
||||
func validateSafePath(p string) error {
|
||||
cleaned := filepath.Clean(p)
|
||||
if strings.Contains(cleaned, "..") {
|
||||
return fmt.Errorf("unsafe path %q: must not contain \"..\"", p)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadHooks reads the hooks.json file referenced by the manifest.
|
||||
func (p *Plugin) LoadHooks() (*HooksConfig, error) {
|
||||
if p.Manifest.Hooks == "" {
|
||||
return nil, nil
|
||||
}
|
||||
hooksPath := filepath.Join(p.Root, p.Manifest.Hooks)
|
||||
data, err := os.ReadFile(hooksPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("read hooks: %w", err)
|
||||
}
|
||||
var cfg HooksConfig
|
||||
if err := json.Unmarshal(data, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("parse hooks: %w", err)
|
||||
}
|
||||
return &cfg, nil
|
||||
}
|
||||
|
||||
// SkillsDir returns the absolute path to the plugin's skills directory.
|
||||
func (p *Plugin) SkillsDir() string {
|
||||
dir := p.Manifest.Skills
|
||||
if dir == "" {
|
||||
dir = "./skills/"
|
||||
}
|
||||
return filepath.Join(p.Root, dir)
|
||||
}
|
||||
|
||||
// isValidSemver checks if a string is a valid semantic version (major.minor.patch).
|
||||
func isValidSemver(v string) bool {
|
||||
parts := strings.SplitN(strings.TrimPrefix(v, "v"), "-", 2)
|
||||
nums := strings.Split(parts[0], ".")
|
||||
if len(nums) != 3 {
|
||||
return false
|
||||
}
|
||||
for _, n := range nums {
|
||||
if _, err := strconv.Atoi(n); err != nil {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// parseSemver extracts major, minor, patch from a version string.
|
||||
func parseSemver(v string) (int, int, int) {
|
||||
v = strings.TrimPrefix(v, "v")
|
||||
parts := strings.SplitN(v, "-", 2) // strip pre-release
|
||||
nums := strings.Split(parts[0], ".")
|
||||
if len(nums) != 3 {
|
||||
return 0, 0, 0
|
||||
}
|
||||
major, _ := strconv.Atoi(nums[0])
|
||||
minor, _ := strconv.Atoi(nums[1])
|
||||
patch, _ := strconv.Atoi(nums[2])
|
||||
return major, minor, patch
|
||||
}
|
||||
|
||||
// compareSemver compares two semver strings. Returns -1, 0, or 1.
|
||||
func compareSemver(a, b string) int {
|
||||
aMaj, aMin, aPat := parseSemver(a)
|
||||
bMaj, bMin, bPat := parseSemver(b)
|
||||
if aMaj != bMaj {
|
||||
return cmpInt(aMaj, bMaj)
|
||||
}
|
||||
if aMin != bMin {
|
||||
return cmpInt(aMin, bMin)
|
||||
}
|
||||
return cmpInt(aPat, bPat)
|
||||
}
|
||||
|
||||
func cmpInt(a, b int) int {
|
||||
if a < b {
|
||||
return -1
|
||||
}
|
||||
if a > b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,587 @@
|
||||
// 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 plugin
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestParseManifest(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
manifestPath := filepath.Join(dir, "plugin.json")
|
||||
content := `{
|
||||
"name": "conference",
|
||||
"version": "1.0.0",
|
||||
"description": "音视频会议",
|
||||
"type": "managed",
|
||||
"minCLIVersion": "0.9.0",
|
||||
"mcpServers": {
|
||||
"conference": {
|
||||
"type": "streamable-http",
|
||||
"endpoint": "https://mcp.conference.dingtalk.com"
|
||||
},
|
||||
"conference-local": {
|
||||
"type": "stdio",
|
||||
"command": "${DWS_PLUGIN_ROOT}/bin/conference-local",
|
||||
"args": ["--mode", "cli"]
|
||||
}
|
||||
},
|
||||
"skills": "./skills/"
|
||||
}`
|
||||
if err := os.WriteFile(manifestPath, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
m, err := ParseManifest(manifestPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseManifest: %v", err)
|
||||
}
|
||||
|
||||
if m.Name != "conference" {
|
||||
t.Errorf("name = %q, want conference", m.Name)
|
||||
}
|
||||
if m.Version != "1.0.0" {
|
||||
t.Errorf("version = %q, want 1.0.0", m.Version)
|
||||
}
|
||||
if m.Type != "managed" {
|
||||
t.Errorf("type = %q, want managed", m.Type)
|
||||
}
|
||||
if len(m.MCPServers) != 2 {
|
||||
t.Errorf("mcpServers count = %d, want 2", len(m.MCPServers))
|
||||
}
|
||||
if m.MCPServers["conference"].Type != "streamable-http" {
|
||||
t.Errorf("conference server type = %q, want streamable-http", m.MCPServers["conference"].Type)
|
||||
}
|
||||
if m.MCPServers["conference-local"].Type != "stdio" {
|
||||
t.Errorf("conference-local server type = %q, want stdio", m.MCPServers["conference-local"].Type)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestValidate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
manifest Manifest
|
||||
cliVersion string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid manifest",
|
||||
manifest: Manifest{
|
||||
Name: "conference",
|
||||
Version: "1.0.0",
|
||||
Type: "managed",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"conference": {Type: "streamable-http", Endpoint: "https://example.com"},
|
||||
},
|
||||
},
|
||||
cliVersion: "1.0.0",
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid name - too short",
|
||||
manifest: Manifest{
|
||||
Name: "ab",
|
||||
Version: "1.0.0",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid name - uppercase",
|
||||
manifest: Manifest{
|
||||
Name: "MyPlugin",
|
||||
Version: "1.0.0",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid version",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "not-semver",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid type",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "1.0.0",
|
||||
Type: "invalid",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "cli version too low",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "1.0.0",
|
||||
MinCLIVersion: "2.0.0",
|
||||
},
|
||||
cliVersion: "1.0.0",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "streamable-http without endpoint",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "1.0.0",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"srv": {Type: "streamable-http"},
|
||||
},
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "stdio without command",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "1.0.0",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"srv": {Type: "stdio"},
|
||||
},
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "unsafe skills path",
|
||||
manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Version: "1.0.0",
|
||||
Skills: "../../../etc/passwd",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := tt.manifest.Validate(tt.cliVersion)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginToServerDescriptors(t *testing.T) {
|
||||
cliOverlay, _ := json.Marshal(map[string]any{
|
||||
"id": "conference",
|
||||
"command": "conference",
|
||||
})
|
||||
|
||||
p := &Plugin{
|
||||
Manifest: Manifest{
|
||||
Name: "conference",
|
||||
Description: "音视频会议",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"conference": {
|
||||
Type: "streamable-http",
|
||||
Endpoint: "https://mcp.conference.dingtalk.com",
|
||||
CLI: cliOverlay,
|
||||
},
|
||||
"conference-local": {
|
||||
Type: "stdio",
|
||||
Command: "/usr/local/bin/conference-local",
|
||||
},
|
||||
},
|
||||
},
|
||||
Root: "/tmp/plugins/conference",
|
||||
}
|
||||
|
||||
descriptors := p.ToServerDescriptors()
|
||||
|
||||
// Only streamable-http should be converted
|
||||
if len(descriptors) != 1 {
|
||||
t.Fatalf("got %d descriptors, want 1 (stdio should be skipped)", len(descriptors))
|
||||
}
|
||||
|
||||
d := descriptors[0]
|
||||
if d.Key != "conference" {
|
||||
t.Errorf("key = %q, want conference", d.Key)
|
||||
}
|
||||
if d.Endpoint != "https://mcp.conference.dingtalk.com" {
|
||||
t.Errorf("endpoint = %q", d.Endpoint)
|
||||
}
|
||||
if d.Source != "plugin" {
|
||||
t.Errorf("source = %q, want plugin", d.Source)
|
||||
}
|
||||
if d.CLI.ID != "conference" {
|
||||
t.Errorf("cli.id = %q, want conference", d.CLI.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginToServerDescriptorsWithHeaders(t *testing.T) {
|
||||
cliOverlay, _ := json.Marshal(map[string]any{
|
||||
"id": "web-search",
|
||||
"command": "web-search",
|
||||
})
|
||||
|
||||
// Set an environment variable to test expansion
|
||||
t.Setenv("TEST_API_KEY", "sk-test-12345")
|
||||
|
||||
p := &Plugin{
|
||||
Manifest: Manifest{
|
||||
Name: "my-plugin",
|
||||
Description: "Test plugin with headers",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"web-search": {
|
||||
Type: "streamable-http",
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
CLI: cliOverlay,
|
||||
Headers: map[string]string{
|
||||
"Authorization": "Bearer ${TEST_API_KEY}",
|
||||
"X-Custom": "static-value",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
Root: "/tmp/plugins/my-plugin",
|
||||
}
|
||||
|
||||
descriptors := p.ToServerDescriptors()
|
||||
if len(descriptors) != 1 {
|
||||
t.Fatalf("got %d descriptors, want 1", len(descriptors))
|
||||
}
|
||||
|
||||
d := descriptors[0]
|
||||
if d.Key != "web-search" {
|
||||
t.Errorf("key = %q, want web-search", d.Key)
|
||||
}
|
||||
if len(d.AuthHeaders) != 2 {
|
||||
t.Fatalf("AuthHeaders len = %d, want 2", len(d.AuthHeaders))
|
||||
}
|
||||
// Environment variable should be expanded
|
||||
if d.AuthHeaders["Authorization"] != "Bearer sk-test-12345" {
|
||||
t.Errorf("AuthHeaders[Authorization] = %q, want 'Bearer sk-test-12345'", d.AuthHeaders["Authorization"])
|
||||
}
|
||||
if d.AuthHeaders["X-Custom"] != "static-value" {
|
||||
t.Errorf("AuthHeaders[X-Custom] = %q, want static-value", d.AuthHeaders["X-Custom"])
|
||||
}
|
||||
if d.Source != "plugin" {
|
||||
t.Errorf("source = %q, want plugin", d.Source)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginToServerDescriptorsNoHeaders(t *testing.T) {
|
||||
cliOverlay, _ := json.Marshal(map[string]any{
|
||||
"id": "conference",
|
||||
"command": "conference",
|
||||
})
|
||||
|
||||
p := &Plugin{
|
||||
Manifest: Manifest{
|
||||
Name: "conference",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"conference": {
|
||||
Type: "streamable-http",
|
||||
Endpoint: "https://mcp.conference.dingtalk.com",
|
||||
CLI: cliOverlay,
|
||||
},
|
||||
},
|
||||
},
|
||||
Root: "/tmp/plugins/conference",
|
||||
}
|
||||
|
||||
descriptors := p.ToServerDescriptors()
|
||||
if len(descriptors) != 1 {
|
||||
t.Fatalf("got %d descriptors, want 1", len(descriptors))
|
||||
}
|
||||
if descriptors[0].AuthHeaders != nil {
|
||||
t.Errorf("AuthHeaders = %v, want nil for server without headers", descriptors[0].AuthHeaders)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManifestWithHeaders(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
manifestPath := filepath.Join(dir, "plugin.json")
|
||||
content := `{
|
||||
"name": "api-plugin",
|
||||
"version": "1.0.0",
|
||||
"mcpServers": {
|
||||
"api-server": {
|
||||
"type": "streamable-http",
|
||||
"endpoint": "https://api.example.com/mcp",
|
||||
"headers": {
|
||||
"Authorization": "Bearer ${MY_API_KEY}",
|
||||
"X-Custom-Header": "custom-value"
|
||||
}
|
||||
}
|
||||
}
|
||||
}`
|
||||
if err := os.WriteFile(manifestPath, []byte(content), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
m, err := ParseManifest(manifestPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseManifest: %v", err)
|
||||
}
|
||||
|
||||
srv := m.MCPServers["api-server"]
|
||||
if srv == nil {
|
||||
t.Fatal("api-server not found in MCPServers")
|
||||
}
|
||||
if len(srv.Headers) != 2 {
|
||||
t.Fatalf("Headers len = %d, want 2", len(srv.Headers))
|
||||
}
|
||||
if srv.Headers["Authorization"] != "Bearer ${MY_API_KEY}" {
|
||||
t.Errorf("Headers[Authorization] = %q, want raw template", srv.Headers["Authorization"])
|
||||
}
|
||||
if srv.Headers["X-Custom-Header"] != "custom-value" {
|
||||
t.Errorf("Headers[X-Custom-Header] = %q, want custom-value", srv.Headers["X-Custom-Header"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoaderScanEmpty(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{
|
||||
PluginsDir: dir,
|
||||
CLIVersion: "1.0.0",
|
||||
}
|
||||
|
||||
user := loader.LoadUser()
|
||||
if len(user) != 0 {
|
||||
t.Errorf("expected 0 user plugins, got %d", len(user))
|
||||
}
|
||||
}
|
||||
|
||||
// TestRemovePluginPurgesSettings verifies RemovePlugin fully purges the
|
||||
// plugin's settings — both its enabled flag and any pluginConfigs entry —
|
||||
// so settings.json does not retain dangling state for a plugin that no
|
||||
// longer exists on disk.
|
||||
func TestRemovePluginPurgesSettings(t *testing.T) {
|
||||
const pkgName = "my-plugin"
|
||||
dir := t.TempDir()
|
||||
pluginDir := filepath.Join(dir, "user", pkgName)
|
||||
if err := os.MkdirAll(pluginDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(pluginDir, "plugin.json"),
|
||||
[]byte(`{"name":"`+pkgName+`","version":"1.0.0"}`), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Seed settings.json with an explicit enabled flag and a
|
||||
// pluginConfigs entry to verify both get purged.
|
||||
settings := &Settings{
|
||||
EnabledPlugins: map[string]bool{pkgName: true, "other-plugin": true},
|
||||
PluginConfigs: map[string]map[string]any{
|
||||
pkgName: {"API_KEY": "secret"},
|
||||
"other-plugin": {"TOKEN": "keep-me"},
|
||||
},
|
||||
}
|
||||
loader.saveSettings(settings)
|
||||
|
||||
if err := loader.RemovePlugin(pkgName, false); err != nil {
|
||||
t.Fatalf("RemovePlugin: %v", err)
|
||||
}
|
||||
|
||||
reloaded := loader.loadSettings()
|
||||
if _, exists := reloaded.EnabledPlugins[pkgName]; exists {
|
||||
t.Errorf("EnabledPlugins should not retain removed plugin %q", pkgName)
|
||||
}
|
||||
if _, exists := reloaded.PluginConfigs[pkgName]; exists {
|
||||
t.Errorf("PluginConfigs should not retain removed plugin %q", pkgName)
|
||||
}
|
||||
if !reloaded.EnabledPlugins["other-plugin"] {
|
||||
t.Error("unrelated EnabledPlugins entry should be preserved")
|
||||
}
|
||||
if reloaded.PluginConfigs["other-plugin"]["TOKEN"] != "keep-me" {
|
||||
t.Error("unrelated PluginConfigs entry should be preserved")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPluginEnabled(t *testing.T) {
|
||||
s := &Settings{
|
||||
EnabledPlugins: map[string]bool{
|
||||
"my-plugin": true,
|
||||
"disabled": false,
|
||||
},
|
||||
}
|
||||
|
||||
if !isPluginEnabled(s, "my-plugin") {
|
||||
t.Error("my-plugin should be enabled")
|
||||
}
|
||||
if isPluginEnabled(s, "disabled") {
|
||||
t.Error("disabled should not be enabled")
|
||||
}
|
||||
if !isPluginEnabled(s, "not-in-list") {
|
||||
t.Error("unlisted plugin should default to enabled")
|
||||
}
|
||||
if !isPluginEnabled(nil, "anything") {
|
||||
t.Error("nil settings should default to enabled")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseGitURL(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
url string
|
||||
wantWS string
|
||||
wantRepo string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "https with .git",
|
||||
url: "https://github.com/PeterGuy326/hello-plugin.git",
|
||||
wantWS: "PeterGuy326",
|
||||
wantRepo: "hello-plugin",
|
||||
},
|
||||
{
|
||||
name: "https without .git",
|
||||
url: "https://github.com/DingTalk-Real-AI/conference",
|
||||
wantWS: "DingTalk-Real-AI",
|
||||
wantRepo: "conference",
|
||||
},
|
||||
{
|
||||
name: "ssh format",
|
||||
url: "git@github.com:DingTalk-Real-AI/conference.git",
|
||||
wantWS: "DingTalk-Real-AI",
|
||||
wantRepo: "conference",
|
||||
},
|
||||
{
|
||||
name: "invalid - no repo",
|
||||
url: "https://github.com/onlyone",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
ws, repo, err := parseGitURL(tt.url)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("parseGitURL() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if !tt.wantErr {
|
||||
if ws != tt.wantWS {
|
||||
t.Errorf("workspace = %q, want %q", ws, tt.wantWS)
|
||||
}
|
||||
if repo != tt.wantRepo {
|
||||
t.Errorf("repo = %q, want %q", repo, tt.wantRepo)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestDevPluginRegistration(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Create a dev plugin directory
|
||||
devDir := filepath.Join(t.TempDir(), "my-dev-plugin")
|
||||
if err := os.MkdirAll(devDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
manifest := `{"name":"my-dev-plugin","version":"0.1.0","type":"user"}`
|
||||
if err := os.WriteFile(filepath.Join(devDir, "plugin.json"), []byte(manifest), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// Register dev plugin
|
||||
if err := loader.RegisterDevPlugin("my-dev-plugin", devDir); err != nil {
|
||||
t.Fatalf("RegisterDevPlugin: %v", err)
|
||||
}
|
||||
|
||||
// Load dev plugins
|
||||
plugins := loader.LoadDev()
|
||||
if len(plugins) != 1 {
|
||||
t.Fatalf("expected 1 dev plugin, got %d", len(plugins))
|
||||
}
|
||||
if plugins[0].Manifest.Name != "my-dev-plugin" {
|
||||
t.Errorf("name = %q, want my-dev-plugin", plugins[0].Manifest.Name)
|
||||
}
|
||||
if plugins[0].Root != devDir {
|
||||
t.Errorf("root = %q, want %q (should load from source dir, not copy)", plugins[0].Root, devDir)
|
||||
}
|
||||
|
||||
// Unregister
|
||||
if err := loader.UnregisterDevPlugin("my-dev-plugin"); err != nil {
|
||||
t.Fatalf("UnregisterDevPlugin: %v", err)
|
||||
}
|
||||
|
||||
// Should be empty now
|
||||
plugins = loader.LoadDev()
|
||||
if len(plugins) != 0 {
|
||||
t.Errorf("expected 0 dev plugins after unregister, got %d", len(plugins))
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnregisterDevPluginNotFound(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
err := loader.UnregisterDevPlugin("nonexistent")
|
||||
if err == nil {
|
||||
t.Error("expected error when unregistering nonexistent dev plugin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSyncSkills(t *testing.T) {
|
||||
// Create a plugin with skills
|
||||
pluginDir := t.TempDir()
|
||||
skillsDir := filepath.Join(pluginDir, "skills", "test-plugin")
|
||||
if err := os.MkdirAll(skillsDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
skillContent := "# Test Plugin Skill"
|
||||
if err := os.WriteFile(filepath.Join(skillsDir, "SKILL.md"), []byte(skillContent), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
p := &Plugin{
|
||||
Manifest: Manifest{
|
||||
Name: "test-plugin",
|
||||
Skills: "./skills/test-plugin",
|
||||
},
|
||||
Root: pluginDir,
|
||||
}
|
||||
|
||||
// Create a mock agent directory
|
||||
home, _ := os.UserHomeDir()
|
||||
agentDir := filepath.Join(home, ".agents", "skills")
|
||||
// Only run if .agents exists (don't create in CI)
|
||||
if _, err := os.Stat(filepath.Dir(agentDir)); err == nil {
|
||||
SyncSkills([]*Plugin{p})
|
||||
|
||||
synced := filepath.Join(agentDir, "dws", "plugins", "test-plugin", "SKILL.md")
|
||||
if _, err := os.Stat(synced); err == nil {
|
||||
data, _ := os.ReadFile(synced)
|
||||
if string(data) != skillContent {
|
||||
t.Errorf("synced content = %q, want %q", string(data), skillContent)
|
||||
}
|
||||
// Cleanup
|
||||
os.RemoveAll(filepath.Join(agentDir, "dws", "plugins", "test-plugin"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func contains(s, substr string) bool {
|
||||
return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsSubstring(s, substr))
|
||||
}
|
||||
|
||||
func containsSubstring(s, sub string) bool {
|
||||
for i := 0; i+len(sub) <= len(s); i++ {
|
||||
if s[i:i+len(sub)] == sub {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -34,9 +34,27 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/validate"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_ALLOW_HTTP_ENDPOINTS",
|
||||
Category: configmeta.CategorySecurity,
|
||||
Description: "允许非 HTTPS 的 MCP 端点 (仅限 loopback)",
|
||||
DefaultValue: "(禁用)",
|
||||
Example: "1",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_TRUSTED_DOMAINS",
|
||||
Category: configmeta.CategoryNetwork,
|
||||
Description: "信任的 HTTPS 域名白名单 (逗号分隔,* 信任所有)",
|
||||
DefaultValue: "*.dingtalk.com",
|
||||
Example: "*.dingtalk.com,custom.example.com",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
trustedDomainsEnv = "DWS_TRUSTED_DOMAINS"
|
||||
defaultTrustedDomains = "*.dingtalk.com"
|
||||
@@ -194,7 +212,7 @@ func (r *ToolCallResult) UnmarshalJSON(data []byte) error {
|
||||
func defaultTransport() *http.Transport {
|
||||
return &http.Transport{
|
||||
DialContext: (&net.Dialer{
|
||||
Timeout: 10 * time.Second,
|
||||
Timeout: 3 * time.Second,
|
||||
KeepAlive: 30 * time.Second,
|
||||
}).DialContext,
|
||||
TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12},
|
||||
@@ -280,6 +298,7 @@ func SupportedProtocolVersions() []string {
|
||||
}
|
||||
|
||||
func (c *Client) Initialize(ctx context.Context, endpoint string) (InitializeResult, error) {
|
||||
var lastErr error
|
||||
for _, version := range SupportedProtocolVersions() {
|
||||
params := map[string]any{
|
||||
"capabilities": map[string]any{},
|
||||
@@ -291,18 +310,33 @@ func (c *Client) Initialize(ctx context.Context, endpoint string) (InitializeRes
|
||||
}
|
||||
|
||||
var payload InitializeResult
|
||||
if err := c.callJSONRPC(ctx, endpoint, requestEnvelope{
|
||||
err := c.callJSONRPC(ctx, endpoint, requestEnvelope{
|
||||
JSONRPC: "2.0",
|
||||
ID: 1,
|
||||
Method: "initialize",
|
||||
Params: params,
|
||||
}, true, &payload); err == nil {
|
||||
}, true, &payload)
|
||||
if err == nil {
|
||||
if payload.ProtocolVersion == "" {
|
||||
payload.ProtocolVersion = version
|
||||
}
|
||||
payload.RequestedProtocolVersion = version
|
||||
return payload, nil
|
||||
}
|
||||
lastErr = err
|
||||
// Only protocol-level JSON-RPC errors justify trying another version.
|
||||
// Transport/HTTP failures (dial timeout, connection refused, HTTP 5xx,
|
||||
// etc.) fail identically regardless of protocol version, so looping
|
||||
// over three versions only multiplies the dial cost — e.g. on an
|
||||
// unreachable endpoint, three 3-second dials amount to 9 seconds
|
||||
// before the outer context surrenders.
|
||||
var callErr *CallError
|
||||
if !errors.As(err, &callErr) || callErr.Stage != CallStageJSONRPC {
|
||||
return InitializeResult{}, err
|
||||
}
|
||||
}
|
||||
if lastErr != nil {
|
||||
return InitializeResult{}, lastErr
|
||||
}
|
||||
return InitializeResult{}, apperrors.NewDiscovery(fmt.Sprintf("initialize failed for all supported protocol versions at %s", RedactURL(endpoint)))
|
||||
}
|
||||
@@ -491,6 +525,15 @@ func (c *Client) doWithRetry(ctx context.Context, endpoint string, body []byte)
|
||||
}
|
||||
}
|
||||
|
||||
// Diagnostic: log identity-related headers on first attempt.
|
||||
if attempt == 0 && c.FileLogger != nil {
|
||||
c.FileLogger.LogAttrs(context.Background(), slog.LevelDebug, "http_request_headers",
|
||||
slog.String("endpoint", endpoint),
|
||||
slog.String("x-user-access-token-present", fmt.Sprintf("%t", req.Header.Get("x-user-access-token") != "")),
|
||||
slog.Int("extra_headers_count", len(c.ExtraHeaders)),
|
||||
)
|
||||
}
|
||||
|
||||
resp, err := c.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -71,6 +72,72 @@ func TestInitializeNegotiatesProtocolVersion(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitializeShortCircuitsOnHTTPError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Track which protocol versions actually get sent. With the short-circuit
|
||||
// in place, a transport-layer (HTTP) failure should fail Initialize on the
|
||||
// FIRST version without iterating through every supported version. Without
|
||||
// the short-circuit, three round-trips would happen — needlessly tripling
|
||||
// every CLI startup when a plugin endpoint is broken (issue #119).
|
||||
var seenVersions []string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
t.Fatalf("decode request: %v", err)
|
||||
}
|
||||
params := req["params"].(map[string]any)
|
||||
seenVersions = append(seenVersions, params["protocolVersion"].(string))
|
||||
http.Error(w, "boom", http.StatusInternalServerError)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(server.Client())
|
||||
client.MaxRetries = 0 // skip the HTTP retry loop — we only care about version iteration
|
||||
|
||||
if _, err := client.Initialize(context.Background(), server.URL); err == nil {
|
||||
t.Fatal("Initialize() error = nil, want HTTP error")
|
||||
}
|
||||
|
||||
if len(seenVersions) != 1 {
|
||||
t.Fatalf("Initialize() attempted %d protocol versions (%v), want 1 — HTTP failures must short-circuit",
|
||||
len(seenVersions), seenVersions)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitializeShortCircuitsOnDialFailure(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Bind to an ephemeral port, then close the listener. Subsequent connects
|
||||
// to that address fail with "connection refused" almost instantly. With
|
||||
// the short-circuit, three protocol versions would otherwise stack three
|
||||
// dial-error returns; we only want one — and we want Initialize to return
|
||||
// well under the per-dial budget.
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
endpoint := "http://" + ln.Addr().String()
|
||||
_ = ln.Close()
|
||||
|
||||
client := NewClient(nil)
|
||||
client.MaxRetries = 0
|
||||
|
||||
start := time.Now()
|
||||
if _, err := client.Initialize(context.Background(), endpoint); err == nil {
|
||||
t.Fatal("Initialize() error = nil, want dial failure")
|
||||
}
|
||||
elapsed := time.Since(start)
|
||||
|
||||
// Three dial attempts (one per supported version) on a refused-connection
|
||||
// path is still fast on loopback, so this assertion is a sanity bound, not
|
||||
// the primary signal — but if the short-circuit regresses, on a real
|
||||
// unreachable address this jumps from one dial timeout to three.
|
||||
if elapsed > 2*time.Second {
|
||||
t.Fatalf("Initialize() took %v, want <2s — dial failure should short-circuit", elapsed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListToolsRetriesOnServerError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
// 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 transport
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// mockMCPHandler is a minimal MCP JSON-RPC handler for testing.
|
||||
func mockMCPHandler(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
var req struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &req); err != nil {
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0", "error": map[string]any{"code": -32700, "message": "parse error"},
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
switch req.Method {
|
||||
case "initialize":
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"result": map[string]any{
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": map[string]any{"tools": map[string]any{}},
|
||||
"serverInfo": map[string]any{"name": "mock-server", "version": "0.0.1"},
|
||||
},
|
||||
})
|
||||
|
||||
case "notifications/initialized":
|
||||
w.WriteHeader(http.StatusOK)
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"result": map[string]any{},
|
||||
})
|
||||
|
||||
case "tools/list":
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"result": map[string]any{
|
||||
"tools": []map[string]any{
|
||||
{
|
||||
"name": "mock_hello",
|
||||
"description": "Say hello",
|
||||
"inputSchema": map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{"name": map[string]any{"type": "string"}},
|
||||
"required": []string{"name"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
case "tools/call":
|
||||
var params struct {
|
||||
Name string `json:"name"`
|
||||
Arguments map[string]any `json:"arguments"`
|
||||
}
|
||||
_ = json.Unmarshal(req.Params, ¶ms)
|
||||
|
||||
if params.Name == "mock_hello" {
|
||||
name, _ := params.Arguments["name"].(string)
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"result": map[string]any{
|
||||
"content": []map[string]any{
|
||||
{"type": "text", "text": "Hello, " + name + "!"},
|
||||
},
|
||||
},
|
||||
})
|
||||
} else {
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"error": map[string]any{"code": -32601, "message": "unknown tool"},
|
||||
})
|
||||
}
|
||||
|
||||
default:
|
||||
writeJSONResp(w, map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req.ID,
|
||||
"error": map[string]any{"code": -32601, "message": "method not found"},
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func writeJSONResp(w http.ResponseWriter, resp any) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(resp)
|
||||
}
|
||||
|
||||
func TestHTTPClientEndToEnd(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(mockMCPHandler))
|
||||
defer server.Close()
|
||||
|
||||
client := NewClient(nil)
|
||||
endpoint := server.URL
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Initialize
|
||||
initResult, err := client.Initialize(ctx, endpoint)
|
||||
if err != nil {
|
||||
t.Fatalf("Initialize: %v", err)
|
||||
}
|
||||
if initResult.ProtocolVersion != "2025-03-26" {
|
||||
t.Errorf("protocolVersion = %q, want 2025-03-26", initResult.ProtocolVersion)
|
||||
}
|
||||
|
||||
// ListTools
|
||||
toolsResult, err := client.ListTools(ctx, endpoint)
|
||||
if err != nil {
|
||||
t.Fatalf("ListTools: %v", err)
|
||||
}
|
||||
if len(toolsResult.Tools) != 1 {
|
||||
t.Fatalf("ListTools: got %d tools, want 1", len(toolsResult.Tools))
|
||||
}
|
||||
if toolsResult.Tools[0].Name != "mock_hello" {
|
||||
t.Errorf("tool name = %q, want mock_hello", toolsResult.Tools[0].Name)
|
||||
}
|
||||
|
||||
// CallTool
|
||||
callResult, err := client.CallTool(ctx, endpoint, "mock_hello", map[string]any{
|
||||
"name": "DWS",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CallTool: %v", err)
|
||||
}
|
||||
if callResult.IsError {
|
||||
t.Fatal("CallTool returned isError=true")
|
||||
}
|
||||
if len(callResult.Blocks) == 0 {
|
||||
t.Fatal("CallTool: no content blocks")
|
||||
}
|
||||
if callResult.Blocks[0].Text != "Hello, DWS!" {
|
||||
t.Errorf("CallTool text = %q, want %q", callResult.Blocks[0].Text, "Hello, DWS!")
|
||||
}
|
||||
|
||||
// CallTool with unknown tool
|
||||
_, err = client.CallTool(ctx, endpoint, "nonexistent", nil)
|
||||
if err == nil {
|
||||
t.Error("CallTool with unknown tool should return error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestHTTPClientInitializeFailsWithBadEndpoint(t *testing.T) {
|
||||
client := NewClient(nil)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err := client.Initialize(ctx, "http://127.0.0.1:0/nonexistent")
|
||||
if err == nil {
|
||||
t.Error("Initialize with bad endpoint should fail")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
// 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 transport
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
// StdioClient manages a local MCP server subprocess, communicating via
|
||||
// stdin/stdout using JSON-RPC 2.0 (newline-delimited).
|
||||
type StdioClient struct {
|
||||
command string
|
||||
args []string
|
||||
env map[string]string
|
||||
|
||||
cmd *exec.Cmd
|
||||
stdin io.WriteCloser
|
||||
stdout *bufio.Reader
|
||||
stderr io.ReadCloser
|
||||
|
||||
mu sync.Mutex // serializes JSON-RPC requests
|
||||
nextID int64
|
||||
started bool
|
||||
}
|
||||
|
||||
// NewStdioClient creates a StdioClient for the given command.
|
||||
// The subprocess is not started until Start() is called.
|
||||
func NewStdioClient(command string, args []string, env map[string]string) *StdioClient {
|
||||
return &StdioClient{
|
||||
command: command,
|
||||
args: args,
|
||||
env: env,
|
||||
}
|
||||
}
|
||||
|
||||
// Start launches the subprocess.
|
||||
func (s *StdioClient) Start(ctx context.Context) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if s.started {
|
||||
return nil
|
||||
}
|
||||
|
||||
cmd := exec.CommandContext(ctx, s.command, s.args...)
|
||||
|
||||
// Build environment: inherit current env + merge plugin-specific vars.
|
||||
cmd.Env = os.Environ()
|
||||
for k, v := range s.env {
|
||||
cmd.Env = append(cmd.Env, k+"="+v)
|
||||
}
|
||||
|
||||
stdin, err := cmd.StdinPipe()
|
||||
if err != nil {
|
||||
return fmt.Errorf("stdio: create stdin pipe: %w", err)
|
||||
}
|
||||
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
stdin.Close()
|
||||
return fmt.Errorf("stdio: create stdout pipe: %w", err)
|
||||
}
|
||||
|
||||
stderr, err := cmd.StderrPipe()
|
||||
if err != nil {
|
||||
stdin.Close()
|
||||
stdout.Close()
|
||||
return fmt.Errorf("stdio: create stderr pipe: %w", err)
|
||||
}
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
return fmt.Errorf("stdio: start process %q: %w", s.command, err)
|
||||
}
|
||||
|
||||
s.cmd = cmd
|
||||
s.stdin = stdin
|
||||
s.stdout = bufio.NewReaderSize(stdout, 64*1024)
|
||||
s.stderr = stderr
|
||||
s.started = true
|
||||
|
||||
// Drain stderr in background for debug logging.
|
||||
go s.drainStderr()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop kills the subprocess and waits for it to exit.
|
||||
func (s *StdioClient) Stop() error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if !s.started || s.cmd == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
s.stdin.Close()
|
||||
|
||||
if s.cmd.Process != nil {
|
||||
_ = s.cmd.Process.Kill()
|
||||
}
|
||||
err := s.cmd.Wait()
|
||||
s.started = false
|
||||
return err
|
||||
}
|
||||
|
||||
// Initialize sends the JSON-RPC initialize request.
|
||||
func (s *StdioClient) Initialize(ctx context.Context) (InitializeResult, error) {
|
||||
params := map[string]any{
|
||||
"protocolVersion": supportedProtocolVersions[0],
|
||||
"capabilities": map[string]any{},
|
||||
"clientInfo": map[string]any{
|
||||
"name": "dws-cli",
|
||||
"version": "1.0.0",
|
||||
},
|
||||
}
|
||||
|
||||
var result InitializeResult
|
||||
if err := s.call(ctx, "initialize", params, &result); err != nil {
|
||||
return InitializeResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// ListTools sends the tools/list JSON-RPC request.
|
||||
func (s *StdioClient) ListTools(ctx context.Context) (ToolsListResult, error) {
|
||||
var result ToolsListResult
|
||||
if err := s.call(ctx, "tools/list", nil, &result); err != nil {
|
||||
return ToolsListResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// CallTool sends the tools/call JSON-RPC request.
|
||||
func (s *StdioClient) CallTool(ctx context.Context, tool string, arguments map[string]any) (ToolCallResult, error) {
|
||||
params := map[string]any{
|
||||
"name": tool,
|
||||
"arguments": arguments,
|
||||
}
|
||||
|
||||
var result ToolCallResult
|
||||
if err := s.call(ctx, "tools/call", params, &result); err != nil {
|
||||
return ToolCallResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// call sends a JSON-RPC request and reads the response. It is serialized
|
||||
// by the mutex to ensure one request at a time over the stdio pipe.
|
||||
func (s *StdioClient) call(ctx context.Context, method string, params any, result any) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
|
||||
if !s.started {
|
||||
return fmt.Errorf("stdio: process not started")
|
||||
}
|
||||
|
||||
id := atomic.AddInt64(&s.nextID, 1)
|
||||
|
||||
req := requestEnvelope{
|
||||
JSONRPC: "2.0",
|
||||
ID: int(id),
|
||||
Method: method,
|
||||
Params: params,
|
||||
}
|
||||
|
||||
reqData, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("stdio: marshal request: %w", err)
|
||||
}
|
||||
|
||||
// Write request line.
|
||||
reqData = append(reqData, '\n')
|
||||
if _, err := s.stdin.Write(reqData); err != nil {
|
||||
return fmt.Errorf("stdio: write request: %w", err)
|
||||
}
|
||||
|
||||
// Read response line (respects context cancellation).
|
||||
type readResult struct {
|
||||
line []byte
|
||||
err error
|
||||
}
|
||||
ch := make(chan readResult, 1)
|
||||
go func() {
|
||||
line, err := s.stdout.ReadBytes('\n')
|
||||
ch <- readResult{line, err}
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return fmt.Errorf("stdio: %w", ctx.Err())
|
||||
case rr := <-ch:
|
||||
if rr.err != nil {
|
||||
return fmt.Errorf("stdio: read response: %w", rr.err)
|
||||
}
|
||||
|
||||
var resp responseEnvelope
|
||||
if err := json.Unmarshal(rr.line, &resp); err != nil {
|
||||
return fmt.Errorf("stdio: unmarshal response: %w", err)
|
||||
}
|
||||
|
||||
if resp.Error != nil {
|
||||
return fmt.Errorf("stdio: RPC error %d: %s", resp.Error.Code, resp.Error.Message)
|
||||
}
|
||||
|
||||
if result != nil {
|
||||
if err := json.Unmarshal(resp.Result, result); err != nil {
|
||||
return fmt.Errorf("stdio: unmarshal result: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// drainStderr reads stderr in the background and logs lines at debug level.
|
||||
func (s *StdioClient) drainStderr() {
|
||||
scanner := bufio.NewScanner(s.stderr)
|
||||
for scanner.Scan() {
|
||||
line := strings.TrimSpace(scanner.Text())
|
||||
if line != "" {
|
||||
slog.Debug("stdio: subprocess stderr", "command", s.command, "line", line)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
// 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 transport
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// TestStdioClientEndToEnd tests the full stdio MCP lifecycle:
|
||||
// Start → Initialize → ListTools → CallTool → Stop.
|
||||
//
|
||||
// It compiles a minimal MCP server helper from testdata and runs it as
|
||||
// a subprocess, exercising the real JSON-RPC protocol over stdin/stdout.
|
||||
func TestStdioClientEndToEnd(t *testing.T) {
|
||||
// Build the test helper server.
|
||||
helperBin := buildTestHelper(t)
|
||||
|
||||
client := NewStdioClient(helperBin, nil, nil)
|
||||
|
||||
// Use background context for Start so subprocess lives for the test duration.
|
||||
if err := client.Start(context.Background()); err != nil {
|
||||
t.Fatalf("Start: %v", err)
|
||||
}
|
||||
defer client.Stop()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// Initialize
|
||||
_, err := client.Initialize(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Initialize: %v", err)
|
||||
}
|
||||
|
||||
// ListTools
|
||||
toolsResult, err := client.ListTools(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("ListTools: %v", err)
|
||||
}
|
||||
if len(toolsResult.Tools) == 0 {
|
||||
t.Fatal("ListTools: no tools returned")
|
||||
}
|
||||
|
||||
// Find the test_echo tool
|
||||
found := false
|
||||
for _, tool := range toolsResult.Tools {
|
||||
if tool.Name == "test_echo" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatalf("ListTools: test_echo tool not found, got tools: %v", toolNames(toolsResult.Tools))
|
||||
}
|
||||
|
||||
// CallTool
|
||||
callResult, err := client.CallTool(ctx, "test_echo", map[string]any{
|
||||
"message": "hello world",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("CallTool: %v", err)
|
||||
}
|
||||
if callResult.IsError {
|
||||
t.Fatalf("CallTool returned isError=true")
|
||||
}
|
||||
|
||||
// Verify response content
|
||||
if len(callResult.Blocks) == 0 {
|
||||
t.Fatal("CallTool: no content blocks")
|
||||
}
|
||||
if callResult.Blocks[0].Text != "Echo: hello world" {
|
||||
t.Errorf("CallTool text = %q, want %q", callResult.Blocks[0].Text, "Echo: hello world")
|
||||
}
|
||||
|
||||
// CallTool with unknown tool should return RPC error
|
||||
_, err = client.CallTool(ctx, "nonexistent", nil)
|
||||
if err == nil {
|
||||
t.Error("CallTool with unknown tool should return error")
|
||||
}
|
||||
|
||||
// Stop
|
||||
if err := client.Stop(); err != nil {
|
||||
// Process killed, expected to return an error
|
||||
_ = err
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioClientStartFailsWithBadCommand(t *testing.T) {
|
||||
client := NewStdioClient("/nonexistent/binary", nil, nil)
|
||||
err := client.Start(context.Background())
|
||||
if err == nil {
|
||||
t.Error("expected error when starting with nonexistent binary")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioClientCallBeforeStart(t *testing.T) {
|
||||
client := NewStdioClient("echo", nil, nil)
|
||||
_, err := client.CallTool(context.Background(), "test", nil)
|
||||
if err == nil {
|
||||
t.Error("expected error when calling before Start")
|
||||
}
|
||||
}
|
||||
|
||||
// buildTestHelper compiles testdata/stdio_test_server.go into a temporary binary.
|
||||
func buildTestHelper(t *testing.T) string {
|
||||
t.Helper()
|
||||
serverSrc := filepath.Join("testdata", "stdio_test_server.go")
|
||||
if _, err := os.Stat(serverSrc); err != nil {
|
||||
t.Skipf("testdata/stdio_test_server.go not found: %v", err)
|
||||
}
|
||||
|
||||
binPath := filepath.Join(t.TempDir(), "stdio-test-server")
|
||||
cmd := exec.Command("go", "build", "-o", binPath, serverSrc)
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to build test helper: %v\n%s", err, out)
|
||||
}
|
||||
return binPath
|
||||
}
|
||||
|
||||
func toolNames(tools []ToolDescriptor) []string {
|
||||
names := make([]string, len(tools))
|
||||
for i, t := range tools {
|
||||
names[i] = t.Name
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// TestStdioProtocolNewlineDelimited verifies that the protocol is correctly
|
||||
// newline-delimited (one JSON object per line).
|
||||
func TestStdioProtocolNewlineDelimited(t *testing.T) {
|
||||
helperBin := buildTestHelper(t)
|
||||
|
||||
cmd := exec.Command(helperBin)
|
||||
stdin, _ := cmd.StdinPipe()
|
||||
stdout, _ := cmd.StdoutPipe()
|
||||
|
||||
if err := cmd.Start(); err != nil {
|
||||
t.Fatalf("start: %v", err)
|
||||
}
|
||||
defer cmd.Process.Kill()
|
||||
|
||||
scanner := bufio.NewScanner(stdout)
|
||||
|
||||
// Send initialize
|
||||
req := `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}` + "\n"
|
||||
fmt.Fprint(stdin, req)
|
||||
|
||||
if !scanner.Scan() {
|
||||
t.Fatal("no response from server")
|
||||
}
|
||||
var resp struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID int `json:"id"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
}
|
||||
if err := json.Unmarshal(scanner.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("unmarshal response: %v", err)
|
||||
}
|
||||
if resp.JSONRPC != "2.0" {
|
||||
t.Errorf("jsonrpc = %q, want 2.0", resp.JSONRPC)
|
||||
}
|
||||
if resp.ID != 1 {
|
||||
t.Errorf("id = %d, want 1", resp.ID)
|
||||
}
|
||||
|
||||
stdin.Close()
|
||||
cmd.Wait()
|
||||
}
|
||||
+126
@@ -0,0 +1,126 @@
|
||||
// Minimal MCP stdio server for integration tests.
|
||||
// Implements initialize, tools/list, tools/call over newline-delimited JSON-RPC.
|
||||
package main
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
)
|
||||
|
||||
type request struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
type response struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
func main() {
|
||||
scanner := bufio.NewScanner(os.Stdin)
|
||||
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
|
||||
|
||||
for scanner.Scan() {
|
||||
line := scanner.Bytes()
|
||||
if len(line) == 0 {
|
||||
continue
|
||||
}
|
||||
var req request
|
||||
if err := json.Unmarshal(line, &req); err != nil {
|
||||
writeResp(response{JSONRPC: "2.0", Error: &rpcError{Code: -32700, Message: "parse error"}})
|
||||
continue
|
||||
}
|
||||
|
||||
switch req.Method {
|
||||
case "initialize":
|
||||
writeResp(response{JSONRPC: "2.0", ID: req.ID, Result: map[string]any{
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": map[string]any{"tools": map[string]any{}},
|
||||
"serverInfo": map[string]any{"name": "test-server", "version": "0.0.1"},
|
||||
}})
|
||||
|
||||
case "tools/list":
|
||||
writeResp(response{JSONRPC: "2.0", ID: req.ID, Result: map[string]any{
|
||||
"tools": []map[string]any{
|
||||
{
|
||||
"name": "test_echo",
|
||||
"description": "Echo the input message",
|
||||
"inputSchema": map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"message": map[string]any{"type": "string", "description": "Message to echo"},
|
||||
},
|
||||
"required": []string{"message"},
|
||||
},
|
||||
},
|
||||
{
|
||||
"name": "test_add",
|
||||
"description": "Add two numbers",
|
||||
"inputSchema": map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"a": map[string]any{"type": "integer"},
|
||||
"b": map[string]any{"type": "integer"},
|
||||
},
|
||||
"required": []string{"a", "b"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}})
|
||||
|
||||
case "tools/call":
|
||||
handleCall(req.ID, req.Params)
|
||||
|
||||
case "notifications/initialized":
|
||||
// no response
|
||||
continue
|
||||
|
||||
default:
|
||||
writeResp(response{JSONRPC: "2.0", ID: req.ID, Error: &rpcError{Code: -32601, Message: "method not found"}})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func handleCall(id, params json.RawMessage) {
|
||||
var p struct {
|
||||
Name string `json:"name"`
|
||||
Arguments map[string]any `json:"arguments"`
|
||||
}
|
||||
if err := json.Unmarshal(params, &p); err != nil {
|
||||
writeResp(response{JSONRPC: "2.0", ID: id, Error: &rpcError{Code: -32602, Message: "invalid params"}})
|
||||
return
|
||||
}
|
||||
|
||||
switch p.Name {
|
||||
case "test_echo":
|
||||
msg, _ := p.Arguments["message"].(string)
|
||||
writeResp(response{JSONRPC: "2.0", ID: id, Result: map[string]any{
|
||||
"content": []map[string]any{{"type": "text", "text": fmt.Sprintf("Echo: %s", msg)}},
|
||||
}})
|
||||
case "test_add":
|
||||
a, _ := p.Arguments["a"].(float64)
|
||||
b, _ := p.Arguments["b"].(float64)
|
||||
writeResp(response{JSONRPC: "2.0", ID: id, Result: map[string]any{
|
||||
"content": []map[string]any{{"type": "text", "text": fmt.Sprintf("%.0f", a+b)}},
|
||||
}})
|
||||
default:
|
||||
writeResp(response{JSONRPC: "2.0", ID: id, Error: &rpcError{Code: -32601, Message: "unknown tool: " + p.Name}})
|
||||
}
|
||||
}
|
||||
|
||||
func writeResp(resp response) {
|
||||
data, _ := json.Marshal(resp)
|
||||
fmt.Fprintf(os.Stdout, "%s\n", data)
|
||||
}
|
||||
@@ -14,8 +14,32 @@ import (
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_UPGRADE_URL",
|
||||
Category: configmeta.CategoryNetwork,
|
||||
Description: "覆盖 GitHub API 地址 (镜像/测试)",
|
||||
DefaultValue: "https://api.github.com",
|
||||
Example: "https://mirror.example.com/api",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "GITHUB_TOKEN",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "GitHub API Token (提升 API 限额)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "GH_TOKEN",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "GitHub API Token 备选 (GITHUB_TOKEN 为空时使用)",
|
||||
Sensitive: true,
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
gitHubAPIBase = "https://api.github.com"
|
||||
defaultOwner = "DingTalk-Real-AI"
|
||||
|
||||
@@ -20,7 +20,7 @@ func VerifySHA256(filePath, expectedHash string) error {
|
||||
|
||||
expectedHash = strings.ToLower(strings.TrimSpace(expectedHash))
|
||||
if actual != expectedHash {
|
||||
return fmt.Errorf("SHA256 校验失败: 期望 %s..., 实际 %s...", expectedHash[:16], actual[:16])
|
||||
return fmt.Errorf("SHA256 mismatch: want %s, got %s", expectedHash[:16], actual[:16])
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
// 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 cli
|
||||
|
||||
import "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app"
|
||||
|
||||
// MCPIdentityHeaders returns HTTP headers aligned with MCP tool calls
|
||||
// (identity + edition merge). Overlays may pass this to auxiliary clients.
|
||||
func MCPIdentityHeaders() map[string]string {
|
||||
return app.MCPIdentityHeaders()
|
||||
}
|
||||
@@ -18,6 +18,8 @@ package config
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -106,3 +108,64 @@ const (
|
||||
// MaxUploadFileSize is the maximum file size for attachment uploads.
|
||||
MaxUploadFileSize int64 = 100 * 1024 * 1024 // 100 MB
|
||||
)
|
||||
|
||||
// ── Plugin system ──────────────────────────────────────────────────────
|
||||
|
||||
const (
|
||||
// PluginUserDir is the subdirectory under ~/.dws/plugins/ where all
|
||||
// third-party plugins are installed. Every plugin — whether authored
|
||||
// by the DingTalk team or anyone else — lives here with equal status.
|
||||
PluginUserDir = "user"
|
||||
|
||||
// PluginDataDir is the subdirectory under ~/.dws/plugins/ for
|
||||
// plugin persistent data that survives across version updates.
|
||||
PluginDataDir = "data"
|
||||
|
||||
// PluginHookTimeout is the default timeout for plugin hook commands.
|
||||
PluginHookTimeout = 30 * time.Second
|
||||
)
|
||||
|
||||
// ── Platform URLs ────────────────────────────────────────────────────────────
|
||||
// Shared across auth, errors, and device-flow packages.
|
||||
|
||||
const (
|
||||
// DefaultTerminalBaseURL is the DingTalk developer platform base URL.
|
||||
// Override at runtime via ~/.dws/terminal_url file.
|
||||
DefaultTerminalBaseURL = "https://open-dev.dingtalk.com"
|
||||
|
||||
// DeveloperSettingsPath is the path to the organization developer
|
||||
// settings page (CLI access management).
|
||||
DeveloperSettingsPath = "/fe/old#/developerSettings"
|
||||
)
|
||||
|
||||
// DefaultConfigDir returns the default DWS configuration directory.
|
||||
// Priority: DWS_CONFIG_DIR env var > ~/.dws
|
||||
func DefaultConfigDir() string {
|
||||
if envDir := os.Getenv("DWS_CONFIG_DIR"); envDir != "" {
|
||||
return envDir
|
||||
}
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return ".dws"
|
||||
}
|
||||
return filepath.Join(homeDir, ".dws")
|
||||
}
|
||||
|
||||
// GetTerminalBaseURL returns the terminal base URL with priority:
|
||||
// 1. ~/.dws/terminal_url file content (for pre-release environment)
|
||||
// 2. Default value (https://open-dev.dingtalk.com)
|
||||
func GetTerminalBaseURL() string {
|
||||
terminalURLPath := filepath.Join(DefaultConfigDir(), "terminal_url")
|
||||
if data, err := os.ReadFile(terminalURLPath); err == nil {
|
||||
if u := strings.TrimSpace(string(data)); u != "" {
|
||||
return u
|
||||
}
|
||||
}
|
||||
return DefaultTerminalBaseURL
|
||||
}
|
||||
|
||||
// GetDeveloperSettingsURL returns the full URL to the organization developer
|
||||
// settings page, derived from the terminal base URL.
|
||||
func GetDeveloperSettingsURL() string {
|
||||
return GetTerminalBaseURL() + DeveloperSettingsPath
|
||||
}
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
// 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 configmeta provides a central registry of all environment-variable
|
||||
// based configuration items used by the DWS CLI. Each package registers its
|
||||
// own items via init(), and the "dws config list" command reads the registry
|
||||
// to present a unified view to the developer.
|
||||
package configmeta
|
||||
|
||||
import (
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// Category groups related configuration items for display purposes.
|
||||
type Category string
|
||||
|
||||
const (
|
||||
CategoryCore Category = "core"
|
||||
CategoryAuth Category = "auth"
|
||||
CategoryNetwork Category = "network"
|
||||
CategorySecurity Category = "security"
|
||||
CategoryRuntime Category = "runtime"
|
||||
CategoryDebug Category = "debug"
|
||||
CategoryExternal Category = "external"
|
||||
)
|
||||
|
||||
// categoryOrder defines the display order for categories.
|
||||
var categoryOrder = map[Category]int{
|
||||
CategoryCore: 0,
|
||||
CategoryAuth: 1,
|
||||
CategoryNetwork: 2,
|
||||
CategorySecurity: 3,
|
||||
CategoryRuntime: 4,
|
||||
CategoryDebug: 5,
|
||||
CategoryExternal: 6,
|
||||
}
|
||||
|
||||
// ConfigItem describes a single environment-variable configuration item.
|
||||
type ConfigItem struct {
|
||||
Name string // Environment variable name, e.g. "DWS_CONFIG_DIR"
|
||||
Category Category // Logical grouping
|
||||
Description string // Short human-readable description
|
||||
DefaultValue string // Description of the default value
|
||||
Example string // Example value for documentation
|
||||
Sensitive bool // If true, actual value is masked in output
|
||||
Hidden bool // If true, omitted from default list output
|
||||
}
|
||||
|
||||
var (
|
||||
mu sync.RWMutex
|
||||
items []ConfigItem
|
||||
)
|
||||
|
||||
// Register adds a configuration item to the global registry.
|
||||
// Duplicate names are silently ignored (first registration wins).
|
||||
func Register(item ConfigItem) {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
for _, existing := range items {
|
||||
if existing.Name == item.Name {
|
||||
return
|
||||
}
|
||||
}
|
||||
items = append(items, item)
|
||||
}
|
||||
|
||||
// All returns every registered configuration item sorted by category
|
||||
// (display order) then by name.
|
||||
func All() []ConfigItem {
|
||||
mu.RLock()
|
||||
defer mu.RUnlock()
|
||||
out := make([]ConfigItem, len(items))
|
||||
copy(out, items)
|
||||
sort.Slice(out, func(i, j int) bool {
|
||||
ci, cj := categoryOrder[out[i].Category], categoryOrder[out[j].Category]
|
||||
if ci != cj {
|
||||
return ci < cj
|
||||
}
|
||||
return out[i].Name < out[j].Name
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
// ByCategory returns registered items that match the given category.
|
||||
func ByCategory(cat Category) []ConfigItem {
|
||||
all := All()
|
||||
var out []ConfigItem
|
||||
for _, item := range all {
|
||||
if item.Category == cat {
|
||||
out = append(out, item)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// Resolve returns the current value of the named environment variable.
|
||||
// For sensitive items the value is masked. Returns ("", false) when the
|
||||
// variable is not set.
|
||||
func Resolve(name string) (string, bool) {
|
||||
val, ok := os.LookupEnv(name)
|
||||
if !ok {
|
||||
return "", false
|
||||
}
|
||||
|
||||
mu.RLock()
|
||||
defer mu.RUnlock()
|
||||
for _, item := range items {
|
||||
if item.Name == name && item.Sensitive {
|
||||
return maskValue(val), true
|
||||
}
|
||||
}
|
||||
return val, true
|
||||
}
|
||||
|
||||
// Categories returns all known category values in display order.
|
||||
func Categories() []Category {
|
||||
cats := make([]Category, 0, len(categoryOrder))
|
||||
for c := range categoryOrder {
|
||||
cats = append(cats, c)
|
||||
}
|
||||
sort.Slice(cats, func(i, j int) bool {
|
||||
return categoryOrder[cats[i]] < categoryOrder[cats[j]]
|
||||
})
|
||||
return cats
|
||||
}
|
||||
|
||||
func maskValue(v string) string {
|
||||
if len(v) == 0 {
|
||||
return ""
|
||||
}
|
||||
if len(v) <= 4 {
|
||||
return strings.Repeat("*", len(v))
|
||||
}
|
||||
return v[:2] + strings.Repeat("*", len(v)-4) + v[len(v)-2:]
|
||||
}
|
||||
|
||||
// Reset clears the registry. Intended for testing only.
|
||||
func Reset() {
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
items = nil
|
||||
}
|
||||
@@ -0,0 +1,164 @@
|
||||
// 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 configmeta
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRegisterAndAll(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "ZZZ_LAST", Category: CategoryDebug, Description: "last"})
|
||||
Register(ConfigItem{Name: "AAA_FIRST", Category: CategoryCore, Description: "first"})
|
||||
Register(ConfigItem{Name: "MMM_MID", Category: CategoryAuth, Description: "mid"})
|
||||
|
||||
all := All()
|
||||
if len(all) != 3 {
|
||||
t.Fatalf("expected 3 items, got %d", len(all))
|
||||
}
|
||||
// core < auth < debug
|
||||
if all[0].Name != "AAA_FIRST" {
|
||||
t.Errorf("expected AAA_FIRST first, got %s", all[0].Name)
|
||||
}
|
||||
if all[1].Name != "MMM_MID" {
|
||||
t.Errorf("expected MMM_MID second, got %s", all[1].Name)
|
||||
}
|
||||
if all[2].Name != "ZZZ_LAST" {
|
||||
t.Errorf("expected ZZZ_LAST third, got %s", all[2].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterDuplicateIgnored(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "DUP", Category: CategoryCore, Description: "original"})
|
||||
Register(ConfigItem{Name: "DUP", Category: CategoryCore, Description: "duplicate"})
|
||||
|
||||
all := All()
|
||||
if len(all) != 1 {
|
||||
t.Fatalf("expected 1 item, got %d", len(all))
|
||||
}
|
||||
if all[0].Description != "original" {
|
||||
t.Errorf("expected original description, got %q", all[0].Description)
|
||||
}
|
||||
}
|
||||
|
||||
func TestByCategory(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "A", Category: CategoryCore})
|
||||
Register(ConfigItem{Name: "B", Category: CategoryAuth})
|
||||
Register(ConfigItem{Name: "C", Category: CategoryCore})
|
||||
|
||||
core := ByCategory(CategoryCore)
|
||||
if len(core) != 2 {
|
||||
t.Fatalf("expected 2 core items, got %d", len(core))
|
||||
}
|
||||
|
||||
empty := ByCategory(CategoryDebug)
|
||||
if len(empty) != 0 {
|
||||
t.Fatalf("expected 0 debug items, got %d", len(empty))
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveNonSensitive(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "TEST_VAR_PLAIN", Category: CategoryCore})
|
||||
|
||||
t.Setenv("TEST_VAR_PLAIN", "hello")
|
||||
val, ok := Resolve("TEST_VAR_PLAIN")
|
||||
if !ok || val != "hello" {
|
||||
t.Errorf("expected (hello, true), got (%q, %v)", val, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSensitiveMasked(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "TEST_SECRET", Category: CategoryAuth, Sensitive: true})
|
||||
|
||||
t.Setenv("TEST_SECRET", "abcdefgh")
|
||||
val, ok := Resolve("TEST_SECRET")
|
||||
if !ok {
|
||||
t.Fatal("expected ok=true")
|
||||
}
|
||||
if val == "abcdefgh" {
|
||||
t.Error("sensitive value should be masked")
|
||||
}
|
||||
// ab****gh
|
||||
if val != "ab****gh" {
|
||||
t.Errorf("unexpected masked value: %q", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveUnset(t *testing.T) {
|
||||
Reset()
|
||||
defer Reset()
|
||||
|
||||
Register(ConfigItem{Name: "UNSET_VAR", Category: CategoryCore})
|
||||
os.Unsetenv("UNSET_VAR")
|
||||
|
||||
_, ok := Resolve("UNSET_VAR")
|
||||
if ok {
|
||||
t.Error("expected ok=false for unset variable")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaskValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
in, want string
|
||||
}{
|
||||
{"", ""},
|
||||
{"ab", "**"},
|
||||
{"abcd", "****"},
|
||||
{"abcde", "ab*de"},
|
||||
{"abcdefghij", "ab******ij"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
got := maskValue(tc.in)
|
||||
if got != tc.want {
|
||||
t.Errorf("maskValue(%q) = %q, want %q", tc.in, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCategories(t *testing.T) {
|
||||
cats := Categories()
|
||||
if len(cats) != 7 {
|
||||
t.Fatalf("expected 7 categories, got %d", len(cats))
|
||||
}
|
||||
if cats[0] != CategoryCore {
|
||||
t.Errorf("expected core first, got %s", cats[0])
|
||||
}
|
||||
if cats[len(cats)-1] != CategoryExternal {
|
||||
t.Errorf("expected external last, got %s", cats[len(cats)-1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestReset(t *testing.T) {
|
||||
Reset()
|
||||
Register(ConfigItem{Name: "X", Category: CategoryCore})
|
||||
Reset()
|
||||
if len(All()) != 0 {
|
||||
t.Error("expected empty registry after Reset")
|
||||
}
|
||||
}
|
||||
@@ -19,5 +19,12 @@ package edition
|
||||
func defaultHooks() *Hooks {
|
||||
return &Hooks{
|
||||
Name: "open",
|
||||
MergeHeaders: func(base map[string]string) map[string]string {
|
||||
if base == nil {
|
||||
base = make(map[string]string)
|
||||
}
|
||||
base["claw-type"] = "openClaw"
|
||||
return base
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
+21
-3
@@ -73,14 +73,32 @@ type Hooks struct {
|
||||
// --- HTTP headers ---
|
||||
MergeHeaders func(base map[string]string) map[string]string
|
||||
|
||||
// --- auth hooks ---
|
||||
OnAuthError func(configDir string, err error) error
|
||||
TokenProvider func(ctx context.Context, fallback func() (string, error)) (string, error)
|
||||
// --- auth ---
|
||||
AuthClientID string // OAuth client ID for device-flow authorisation
|
||||
AuthClientFromMCP bool // true → fetch client ID from MCP at runtime
|
||||
OnAuthError func(configDir string, err error) error
|
||||
TokenProvider func(ctx context.Context, fallback func() (string, error)) (string, error)
|
||||
|
||||
// --- token persistence (overlay-managed keychain / encrypted storage) ---
|
||||
SaveToken func(configDir string, data []byte) error // persist token blob
|
||||
LoadToken func(configDir string) ([]byte, error) // retrieve token blob
|
||||
DeleteToken func(configDir string) error // remove persisted token
|
||||
|
||||
// --- MCP result classification ---
|
||||
// ClassifyToolResult inspects raw MCP tool-call content and returns a typed
|
||||
// error (e.g. PATError, CLIError) when the response contains a known
|
||||
// gateway-auth or PAT-permission failure. nil → no special handling.
|
||||
ClassifyToolResult func(content map[string]any) error
|
||||
|
||||
// --- product & endpoint ---
|
||||
StaticServers func() []ServerInfo // non-nil → skip Market discovery
|
||||
VisibleProducts func() []string // non-nil → override help visibility
|
||||
RegisterExtraCommands func(root *cobra.Command, caller ToolCaller) // register overlay-only commands
|
||||
|
||||
// AfterPersistentPreRun runs at the end of the root PersistentPreRunE after
|
||||
// global setup (OAuth flag overrides, log level, output sink). Overlays use
|
||||
// this for clients that bypass the MCP runner (e.g. A2A gateway).
|
||||
AfterPersistentPreRun func(cmd *cobra.Command, args []string) error
|
||||
}
|
||||
|
||||
var (
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
// 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 runtimetoken resolves API bearer tokens for features that bypass
|
||||
// the MCP runner (e.g. A2A gateway) but should behave like tool calls.
|
||||
package runtimetoken
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app"
|
||||
)
|
||||
|
||||
// ResolveAccessToken returns a non-empty bearer token using the same sources
|
||||
// and caching rules as MCP when configDir matches the active edition directory;
|
||||
// see app.ResolveAuxiliaryAccessToken.
|
||||
func ResolveAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
return app.ResolveAuxiliaryAccessToken(ctx, configDir, explicitToken)
|
||||
}
|
||||
+9
-3
@@ -1,7 +1,7 @@
|
||||
---
|
||||
name: dws
|
||||
description: 管理钉钉产品能力(AI表格/日历/通讯录/群聊与机器人/待办/审批/考勤/日志/DING消息/工作台/开放平台文档等)。当用户需要操作表格数据、管理日程会议、查询通讯录、管理群聊、机器人发消息、创建待办、提交审批、查看考勤、提交日报周报(钉钉日志模版)时使用。
|
||||
cli_version: ">=1.1.0"
|
||||
description: 管理钉钉产品能力(AI表格/日历/通讯录/群聊与机器人/待办/审批/考勤/日志/DING消息/工作台/开放平台文档/钉钉文档/AI听记等)。当用户需要操作表格数据、管理日程会议、查询通讯录、管理群聊、机器人发消息、创建待办、提交审批、查看考勤、提交日报周报(钉钉日志模版)、读写钉钉文档、查询听记纪要时使用。
|
||||
cli_version: ">=1.0.6"
|
||||
---
|
||||
|
||||
# 钉钉全产品 Skill
|
||||
@@ -24,7 +24,7 @@ cli_version: ">=1.1.0"
|
||||
|
||||
| 产品 | 用途 | 参考文件 |
|
||||
|-------------------|------------------------------------------------------|----------------------------------------------------------------|
|
||||
| `aitable` | AI表格:表格/数据表/字段/记录增删改查/模板搜索 | [aitable.md](./references/products/aitable.md) |
|
||||
| `aitable` | AI表格:Base/数据表/字段/记录/附件/模板搜索 | [aitable.md](./references/products/aitable.md) |
|
||||
| `approval` | 审批:审批表单/发起实例/审批/撤销 | [simple.md](./references/products/simple.md) |
|
||||
| `attendance` | 考勤:打卡记录/排班查询 | [attendance.md](./references/products/attendance.md) |
|
||||
| `calendar` | 日历:日程/参与者/会议室/闲忙查询 | [calendar.md](./references/products/calendar.md) |
|
||||
@@ -32,6 +32,8 @@ cli_version: ">=1.1.0"
|
||||
| `contact` | 通讯录:用户查询(当前用户/搜索/详情)/部门查询(搜索/子部门/成员列表) | [contact.md](./references/products/contact.md) |
|
||||
| `devdoc` | 开放平台文档:搜索开发文档 | [simple.md](./references/products/simple.md) |
|
||||
| `ding` | DING消息:发送/撤回(应用内/短信/电话) | [ding.md](./references/products/ding.md) |
|
||||
| `doc` | 钉钉文档:搜索/浏览/读写/块级编辑/评论 | [doc.md](./references/products/doc.md) |
|
||||
| `minutes` | AI听记:听记列表/摘要/关键词/转写/待办/思维导图/发言人/热词 | [minutes.md](./references/products/minutes.md) |
|
||||
| `report` | 日志:按模版创建/收件箱/已发送/模版查看/详情/已读统计 | [report.md](./references/products/report.md) |
|
||||
| `todo` | 待办:创建(含优先级/截止时间)/查询/修改/标记完成/删除 | [todo.md](./references/products/todo.md) |
|
||||
| `workbench` | 工作台:应用管理 | [workbench.md](./references/products/workbench.md) |
|
||||
@@ -46,6 +48,8 @@ cli_version: ">=1.1.0"
|
||||
用户提到"通讯录/同事/部门/组织架构" → `contact`
|
||||
用户提到"开发/API/调用错误 文档" → `devdoc`
|
||||
用户提到"DING/紧急消息/电话提醒" → `ding`
|
||||
用户提到"钉钉文档/云文档/知识库/读写文档/块级编辑/文档评论" → `doc`
|
||||
用户提到"听记/AI听记/会议纪要/转写/摘要/思维导图/发言人/热词" → `minutes`
|
||||
用户提到"日志/日报/周报/日志统计/写日报/提交周报/发日志/填日志" → `report`
|
||||
用户提到"待办/TODO/任务提醒" → `todo`
|
||||
用户提到"工作台/应用管理" → `workbench`
|
||||
@@ -69,6 +73,8 @@ cli_version: ">=1.1.0"
|
||||
| `calendar` | `participant delete` | 移除日程参与者 |
|
||||
| `calendar` | `room delete` | 取消会议室预定 |
|
||||
| `chat` | `group members remove` | 移除群成员 |
|
||||
| `doc` | `delete` | 删除钉钉文档(不可恢复) |
|
||||
| `doc` | `block delete` | 删除文档块 |
|
||||
| `todo` | `task delete` | 删除待办 |
|
||||
|
||||
### 确认流程
|
||||
|
||||
@@ -31,9 +31,9 @@ Flags:
|
||||
|
||||
> 📎 **操作后返回文档链接**:遍历返回的每个 base,拼接 `https://alidocs.dingtalk.com/i/nodes/{baseId}` 返回给用户。
|
||||
|
||||
> ⚠️ **注意**:`base list` 仅返回**最近访问过**的 Base,不是全部 Base。
|
||||
> 新创建的 Base 如果尚未在钉钉前端打开过,可能不会出现在此列表中。
|
||||
> 如需查找特定 Base,请使用 `base search`;如果刚创建完,直接使用 `create` 返回的 `baseId` 即可。
|
||||
> ⚠️ **重要**:`base list` 仅返回**最近访问过**的 Base,**不是全部 Base**。
|
||||
> 如需查找表格,**请优先使用 `base search`**;`base list` 仅作为浏览最近表格的辅助手段。
|
||||
> 如果刚创建完,直接使用 `create` 返回的 `baseId` 即可。
|
||||
|
||||
#### 搜索 AI 表格
|
||||
```
|
||||
@@ -56,6 +56,8 @@ Flags:
|
||||
--base-id string Base 唯一标识 (必填)
|
||||
```
|
||||
|
||||
> 💡 **用户提供 URL 时**:如果用户给出了链接如 `https://alidocs.dingtalk.com/i/nodes/ABC123`,请提取末尾的 `ABC123` 作为 `--base-id` 传入。详见下方「URL → baseId 提取」章节。
|
||||
|
||||
返回 baseName、tables、dashboards 的 summary 信息(不含字段与记录详情)。
|
||||
后续如需 tableId,优先从这里读取。
|
||||
|
||||
@@ -186,15 +188,24 @@ Flags:
|
||||
Usage:
|
||||
dws aitable field create [flags]
|
||||
Example:
|
||||
# 单字段模式
|
||||
dws aitable field create --base-id <BASE_ID> --table-id <TABLE_ID> \
|
||||
--fields '[{"fieldName":"状态","type":"singleSelect","config":{"options":[{"name":"待办"},{"name":"进行中"},{"name":"已完成"}]}}]'
|
||||
--name "状态" --type "singleSelect" --config '{"options":[{"name":"待办"},{"name":"进行中"},{"name":"已完成"}]}'
|
||||
|
||||
# 批量模式
|
||||
dws aitable field create --base-id <BASE_ID> --table-id <TABLE_ID> \
|
||||
--fields '[{"fieldName":"状态","type":"singleSelect","config":{"options":[{"name":"待办"}]}}]'
|
||||
Flags:
|
||||
--base-id string Base ID (必填)
|
||||
--fields string 待新增字段 JSON 数组,至少 1 个,单次最多 15 个 (必填)
|
||||
--name string 单字段名称(与 --type 配合使用,替代 --fields)
|
||||
--type string 单字段类型(参考 table create 字段类型)
|
||||
--config string 单字段配置 JSON(可选,如 options)
|
||||
--fields string 批量新增字段 JSON 数组,单次最多 15 个(与 --name/--type 二选一)
|
||||
--table-id string Table ID (必填)
|
||||
```
|
||||
|
||||
允许部分成功,返回结果逐项标明成功/失败状态。
|
||||
`--name/--type/--config` 为单字段模式;`--fields` 为批量模式;两种模式二选一。
|
||||
|
||||
#### 更新字段
|
||||
```
|
||||
@@ -236,13 +247,13 @@ Usage:
|
||||
Example:
|
||||
dws aitable record query --base-id <BASE_ID> --table-id <TABLE_ID>
|
||||
dws aitable record query --base-id <BASE_ID> --table-id <TABLE_ID> --record-ids rec1,rec2
|
||||
dws aitable record query --base-id <BASE_ID> --table-id <TABLE_ID> --keyword "关键词" --limit 50
|
||||
dws aitable record query --base-id <BASE_ID> --table-id <TABLE_ID> --query "关键词" --limit 50
|
||||
Flags:
|
||||
--base-id string Base ID (必填)
|
||||
--cursor string 分页游标,首次不传
|
||||
--field-ids string 返回字段 ID 列表,逗号分隔,单次最多 100 个
|
||||
--filters string 结构化过滤条件 JSON
|
||||
--keyword string 全文关键词搜索
|
||||
--query string 全文关键词搜索
|
||||
--limit int 单次最大记录数,默认 100,最大 100
|
||||
--record-ids string 指定记录 ID 列表,逗号分隔,单次最多 100 个
|
||||
--sort string 排序条件 JSON 数组
|
||||
@@ -251,6 +262,10 @@ Flags:
|
||||
|
||||
两种模式: 按 ID 取(传 record-ids,忽略 filters/sort)或条件查(filters+sort+cursor 分页)。
|
||||
|
||||
> ⚠️ **排序参数规范(关键)**:`--sort` 需要传 JSON 数组,排序方向字段必须是 `direction`(`asc` 或 `desc`),**不要使用 `order`**。
|
||||
>
|
||||
> 正确示例:`--sort '[{"fieldId":"wm8ns9bw2vmucb45xj3ix","direction":"desc"}]'`
|
||||
|
||||
filters 结构:`{"operator":"and|or","operands":[{"operator":"<op>","operands":["<fieldId>","<value>"]}]}`
|
||||
|
||||
> 💡 **singleSelect/multipleSelect 过滤**:filters 中可传 option id 或 option name,但建议优先用 **option id**(通过 `field get` 获取),更可靠。
|
||||
@@ -365,10 +380,48 @@ Flags:
|
||||
|
||||
> 📎 **模板预览地址**:`https://docs.dingtalk.com/table/template/{templateId}`
|
||||
|
||||
## 复杂操作
|
||||
|
||||
### 仪表盘 / 图表(建议顺序)
|
||||
|
||||
```bash
|
||||
# 1) 先看配置模板(JSONC)
|
||||
dws aitable dashboard config-example --format json
|
||||
dws aitable chart widgets-example --format json
|
||||
|
||||
# 2) 先拿 dashboard,再拿 chart 详情
|
||||
dws aitable dashboard get --base-id <BASE_ID> --dashboard-id <DASHBOARD_ID> --format json
|
||||
dws aitable chart get --base-id <BASE_ID> --dashboard-id <DASHBOARD_ID> --chart-id <CHART_ID> --format json
|
||||
```
|
||||
|
||||
要点:
|
||||
|
||||
- `dashboard get` 返回的 `charts[].chartId` 可直接给 `chart get` 使用。
|
||||
- `dashboard share get` 可能返回 `404`(资源不存在或未开通),需按可重试错误处理,不要误判为参数拼错。
|
||||
- `chart share get` 可正常返回 `enabled/shareUrl`,用于分享状态判断。
|
||||
|
||||
### 导出数据(两阶段轮询)
|
||||
|
||||
`export data` 常见为异步任务:首次调用可能只返回 `taskId`,需要继续轮询。
|
||||
|
||||
```bash
|
||||
# 第一步:创建任务(按 scope 传必要参数)
|
||||
dws aitable export data --base-id <BASE_ID> --scope table --table-id <TABLE_ID> --format excel --timeout-ms 1000
|
||||
|
||||
# 第二步:拿 taskId 继续轮询,直到返回 downloadUrl
|
||||
dws aitable export data --base-id <BASE_ID> --task-id <TASK_ID> --timeout-ms 3000
|
||||
```
|
||||
|
||||
参数约束
|
||||
|
||||
- `scope=all`:只需 `base-id`
|
||||
- `scope=table`:必须 `table-id`
|
||||
- `scope=view`:必须同时 `table-id + view-id`
|
||||
|
||||
## 意图判断
|
||||
|
||||
用户说"表格/多维表/AI表格":
|
||||
- 查看/列表 → `base list`
|
||||
- 查看/查找/列表 → `base search`(优先)或 `base list`(仅浏览最近访问)
|
||||
- 搜索 → `base search`
|
||||
- 详情 → `base get`
|
||||
- 创建 → `base create`
|
||||
@@ -433,6 +486,63 @@ dws aitable record create --base-id <BASE_ID> --table-id <TABLE_ID> \
|
||||
- 所有操作使用 ID(baseId/tableId/fieldId/recordId),不使用名称
|
||||
- records 的 cells key 是 fieldId,不是字段名称
|
||||
|
||||
## `--filters` 筛选语法排错与使用规范(极易出错)
|
||||
|
||||
调用 `record query` 时,如果条件筛选**完全失效(查询返回了所有记录)**,通常是因为 `--filters` JSON 语法错误,API 默默丢弃了不合规的 filter。
|
||||
|
||||
**强制规则:**
|
||||
1. **根节点必须是逻辑操作符**:`"operator"` 必须是 `"and"` 或 `"or"`,不能是 `"eq"` 等比较操作符。
|
||||
2. 比较操作必须放在根节点的 `"operands"` 数组内的对象中。
|
||||
3. `singleSelect` 和 `multipleSelect` 字段,推荐使用 **选项的 exact String 名称 (name)** 作为比较值,而不是 ID。
|
||||
4. **内层比较操作符语义**:支持 `eq`(等于)、`not_eq`(不等于)、`contain`(包含/模糊搜索)、`not_contain`(不包含)、`gt/gte`(大于/大于等于)、`lt/lte`(小于/小于等于)、`is_empty/is_not_empty`(为空/不为空,对应 operands 内只传单个 fieldId)。
|
||||
|
||||
**精简防呆模板与 4 种衍生情况**
|
||||
```json
|
||||
{
|
||||
"operator": "and", // 情况 4 (OR 查询): 这里改为 "or"
|
||||
"operands": [
|
||||
{
|
||||
"operator": "eq", // 情况 3 (文本包含): 这里改为 "contain"
|
||||
"operands": ["fld_state", "进行中"] // 情况 1 (基础等于)
|
||||
}
|
||||
// 情况 2 (多条件 AND): 在此行增加类似 {"operator":"eq","operands":["fld_priority","高"]}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
**错误示例 1:缺失根节点 and/or**(API 将忽略该 filter,返回全表)
|
||||
```json
|
||||
{"operator":"eq","operands":["fldXXX","本科"]}
|
||||
```
|
||||
|
||||
**错误示例 2:传入选项 ID 而非名称**(可能导致匹配不到 0 记录)
|
||||
```json
|
||||
{"operator":"and","operands":[{"operator":"eq","operands":["fldXXX","CXzrOHK9JI"]}]}
|
||||
```
|
||||
|
||||
## URL → baseId 提取
|
||||
|
||||
用户经常通过钉钉链接指定表格,链接格式为:
|
||||
|
||||
```
|
||||
https://alidocs.dingtalk.com/i/nodes/{baseId}
|
||||
https://alidocs.dingtalk.com/i/nodes/{baseId}?xxx=yyy
|
||||
```
|
||||
|
||||
**处理规则**(必须严格遵守):
|
||||
1. 当用户提供了包含 `alidocs.dingtalk.com/i/nodes/` 的 URL 时,提取 `/nodes/` 后的路径段作为 `baseId`
|
||||
2. 去掉尾部的查询参数(`?` 及其后内容)和尾部斜杠
|
||||
3. 将提取得到的 ID 传入 `--base-id` 参数
|
||||
|
||||
**示例**:
|
||||
```
|
||||
用户输入:帮我查看 https://alidocs.dingtalk.com/i/nodes/ABC123XYZ 这个表格
|
||||
→ 提取 baseId = ABC123XYZ
|
||||
→ 执行: dws aitable base get --base-id ABC123XYZ --format json
|
||||
```
|
||||
|
||||
> 💡 **注意**:URL 中的 nodeId 在 AI 表格场景下等同于 baseId,可以直接作为 `--base-id` 使用。
|
||||
|
||||
### cells 写入/读取格式速查
|
||||
|
||||
| 字段类型 | 写入格式 | 读取返回格式 |
|
||||
@@ -450,3 +560,7 @@ dws aitable record create --base-id <BASE_ID> --table-id <TABLE_ID> \
|
||||
| group | `[{"cid":"xxx"}]` (注意: key 是 cid,不是 openConversationId) | 同写入 |
|
||||
|
||||
- 详见 [field-rules.md](../field-rules.md) 和 [error-codes.md](../error-codes.md)
|
||||
|
||||
## 相关产品
|
||||
|
||||
- [doc](./doc.md) — 富文本文档编辑,不是结构化数据表格
|
||||
|
||||
+652
-127
@@ -1,77 +1,113 @@
|
||||
# 群聊与机器人 (chat) 命令参考
|
||||
# 会话与群聊 (chat) 命令参考
|
||||
|
||||
> 命令别名: `dws im` 等价于 `dws chat`
|
||||
|
||||
## 命令总览
|
||||
|
||||
### group (群组管理)
|
||||
|
||||
| 子命令 | 用途 |
|
||||
|-------|------|
|
||||
| `search` | 搜索群聊 |
|
||||
| `group create` | 创建群 |
|
||||
| `group create` | 创建内部群 |
|
||||
| `group create-org` | 创建企业全员群 |
|
||||
| `group members list` | 查看群成员列表 |
|
||||
| `group members add` | 添加群成员 |
|
||||
| `group members remove` | 移除群成员(⚠️ 危险操作) |
|
||||
| `group members add-bot` | 添加机器人到群 |
|
||||
| `group rename` | 修改群名称 |
|
||||
| `bot search` | 搜索我的机器人 |
|
||||
| `search` | 搜索群会话 |
|
||||
| `search-common` | 搜索共同群 |
|
||||
|
||||
### message (会话消息管理)
|
||||
|
||||
| 子命令 | 用途 |
|
||||
|-------|------|
|
||||
| `message send` | 以当前用户身份发群消息或单聊消息 |
|
||||
| `message send-personal` | 发送个人消息(⚠️ 敏感操作) |
|
||||
| `message list` | 拉取群聊或单聊会话消息 |
|
||||
| `message list-all` | 按时间范围拉取当前用户所有会话消息 |
|
||||
| `message list-topic-replies` | 拉取群话题回复消息列表 |
|
||||
| `message list-by-sender` | 搜索指定发送者的消息 |
|
||||
| `message list-mentions` | 拉取 @我 的消息 |
|
||||
| `message list-focused` | 拉取特别关注人的消息 |
|
||||
| `message list-unread-conversations` | 获取未读会话列表 |
|
||||
| `message search` | 按关键词搜索消息 |
|
||||
| `message info` | 获取会话信息 |
|
||||
| `message send-by-bot` | 机器人发消息(群聊或批量单聊) |
|
||||
| `message recall-by-bot` | 机器人撤回消息(群聊或批量单聊) |
|
||||
| `message recall-by-bot` | 机器人撤回消息 |
|
||||
| `message send-by-webhook` | 自定义机器人 Webhook 发消息 |
|
||||
| `list-top-conversations` | 拉取置顶会话列表 |
|
||||
|
||||
### bot (机器人管理)
|
||||
|
||||
| 子命令 | 用途 |
|
||||
|-------|------|
|
||||
| `bot search` | 搜索我的机器人 |
|
||||
| `bot create` | 创建企业机器人 |
|
||||
| `bot search-groups` | 搜索机器人所在群 |
|
||||
|
||||
---
|
||||
|
||||
## search — 搜索群聊
|
||||
## group create — 创建内部群
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat search [flags]
|
||||
Example:
|
||||
dws chat search --query "项目冲刺" --format json
|
||||
Flags:
|
||||
--query string 搜索关键词 (必填)
|
||||
--cursor string 分页游标(首页留空)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## group create — 创建群
|
||||
当前登录用户自动成为群主。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat group create [flags]
|
||||
Example:
|
||||
dws chat group create --name "Q1 项目冲刺群" --users userId1,userId2,userId3 --format json
|
||||
dws chat group create --name "Q1 项目冲刺群" --users userId1,userId2,userId3
|
||||
Flags:
|
||||
--name string 群名称 (必填)
|
||||
--users string 群成员 userId 列表,逗号分隔 (必填)
|
||||
--users string 成员 userId 列表,用户本身会自动加入,无需包含,逗号分隔,不超过20个 (必填)
|
||||
--name string 群名称 (必填)
|
||||
```
|
||||
|
||||
> 当前用户自动作为群主加入,无需在 --users 中重复传入。
|
||||
---
|
||||
|
||||
## group create-org — 创建企业全员群
|
||||
|
||||
创建面向企业组织的群,成员通过 userId 列表指定。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat group create-org [flags]
|
||||
Example:
|
||||
dws chat group create-org --name "全员通知群" --users userId1,userId2
|
||||
Flags:
|
||||
--name string 群名称 (必填)
|
||||
--users string 成员 userId 列表,逗号分隔 (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## group members list — 查看群成员列表
|
||||
|
||||
分页查询指定群聊的成员。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat group members list [flags]
|
||||
Example:
|
||||
dws chat group members list --id <openConversationId> --format json
|
||||
dws chat group members list --id <openconversation_id>
|
||||
Flags:
|
||||
--id string 群会话 ID (必填)
|
||||
--cursor string 分页游标
|
||||
--cursor string 分页游标,首次从 0 开始
|
||||
--id string 群 ID / openconversation_id (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## group members add — 添加群成员
|
||||
|
||||
向指定群聊添加成员,需传入群 ID 与用户 ID 列表。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat group members add [flags]
|
||||
Example:
|
||||
dws chat group members add --id <openConversationId> --users userId1,userId2 --format json
|
||||
dws chat group members add --id <openconversation_id> --users userId1,userId2
|
||||
Flags:
|
||||
--id string 群会话 ID (必填)
|
||||
--users string 要添加的 userId 列表,逗号分隔 (必填)
|
||||
--id string 群 ID / openconversation_id (必填)
|
||||
--users string 要添加的用户 userId 列表,逗号分隔 (必填)
|
||||
```
|
||||
|
||||
---
|
||||
@@ -84,24 +120,26 @@ Flags:
|
||||
Usage:
|
||||
dws chat group members remove [flags]
|
||||
Example:
|
||||
dws chat group members remove --id <openConversationId> --users userId1,userId2 --format json
|
||||
dws chat group members remove --id <openconversation_id> --users userId1,userId2
|
||||
Flags:
|
||||
--id string 群会话 ID (必填)
|
||||
--users string 要移除的 userId 列表,逗号分隔 (必填)
|
||||
--id string 群 ID / openconversation_id (必填)
|
||||
--users string 要移除的用户 userId 列表,逗号分隔 (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## group members add-bot — 添加机器人到群
|
||||
|
||||
将自定义机器人添加到当前用户有管理权限的群聊中,如果没有权限则会报错。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat group members add-bot [flags]
|
||||
Example:
|
||||
dws chat group members add-bot --id <openConversationId> --robot-code <robotCode> --format json
|
||||
dws chat group members add-bot --robot-code <robot-code> --id <openconversation_id>
|
||||
Flags:
|
||||
--id string 群会话 ID (必填)
|
||||
--robot-code string 机器人 code (必填)
|
||||
--id string 群聊 openConversationId (必填)
|
||||
--robot-code string 机器人 Code (必填)
|
||||
```
|
||||
|
||||
---
|
||||
@@ -112,10 +150,351 @@ Flags:
|
||||
Usage:
|
||||
dws chat group rename [flags]
|
||||
Example:
|
||||
dws chat group rename --id <openConversationId> --name "新群名" --format json
|
||||
dws chat group rename --id <openconversation_id> --name "新群名"
|
||||
Flags:
|
||||
--id string 群会话 ID (必填)
|
||||
--name string 新群名称 (必填)
|
||||
--id string 群 ID / openconversation_id (必填)
|
||||
--name string 修改后的群名称 (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## search — 搜索群会话
|
||||
|
||||
根据名称搜索会话列表。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat search [flags]
|
||||
Example:
|
||||
dws chat search --query "项目冲刺"
|
||||
Flags:
|
||||
--cursor string 分页游标 (首页留空)
|
||||
--query string 搜索关键词 (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## search-common — 搜索共同群
|
||||
|
||||
根据昵称列表搜索共同群聊。--nicks 指定要搜索的人员昵称(逗号分隔,必填)。--match-mode 控制匹配模式:AND 表示所有人都在群里,OR 表示任一人在群里(默认 AND)。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat search-common [flags]
|
||||
Example:
|
||||
dws chat search-common --nicks "风雷,山乔" --limit 20 --cursor 0
|
||||
dws chat search-common --nicks "天鸡,乐函" --match-mode OR --limit 20 --cursor 0
|
||||
dws chat search-common --nicks "风雷,山乔,天鸡" --limit 10 --cursor <nextCursor>
|
||||
Flags:
|
||||
--nicks string 要搜索的昵称列表,逗号分隔 (必填)
|
||||
--match-mode string 匹配模式:AND=所有人都在群里,OR=任一人在群里(默认 AND)
|
||||
--limit int 每页返回数量(默认 20)
|
||||
--cursor string 分页游标(默认 "0",翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- --nicks 传人员昵称(花名),逗号分隔,如 "风雷,山乔"
|
||||
- --match-mode AND 表示群里必须包含所有指定的人;OR 表示包含任意一人即可
|
||||
- 翻页:hasMore=true 时,用返回的 nextCursor 作为下次 --cursor
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message send — 以当前用户身份发消息
|
||||
|
||||
--group 指定群聊 ID 发群消息;--user 指定用户 userId 发单聊;--open-dingtalk-id 指定用户 openDingTalkId 发单聊。三者只能选其一,不能同时指定。消息内容为位置参数(恰好 1 个),支持 Markdown。可选 --title 作为消息标题。
|
||||
--群聊时可选 --at-all @所有人,或 --at-users 指定成员(仅群聊时生效)。
|
||||
--发送图片消息:指定 --media-id(通过 dt_media_upload 工具上传获得),自动设置 msgType=image,此时不需要传文本内容。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message send [flags] [<text>]
|
||||
Example:
|
||||
dws chat message send --group <openconversation_id> --text "hello"
|
||||
dws chat message send --user <userId> --text "请查收"
|
||||
dws chat message send --open-dingtalk-id <openDingTalkId> --text "请查收"
|
||||
dws chat message send --group <openconversation_id> "hello"
|
||||
dws chat message send --group <openconversation_id> --title "周报提醒" --text "请大家本周五前提交周报"
|
||||
dws chat message send --group <openconversation_id> --at-all "<@all> 请大家注意"
|
||||
dws chat message send --group <openconversation_id> --at-users userId1,userId2 "<@userId1> <@userId2> 请查收"
|
||||
dws chat message send --group <openconversation_id> --media-id <mediaId>
|
||||
dws chat message send --open-dingtalk-id <openDingTalkId> --media-id <mediaId>
|
||||
Flags:
|
||||
--text string 消息内容(推荐使用,也可用位置参数)
|
||||
--group string 群聊 openconversation_id(群聊时必填)
|
||||
--user string 接收人 userId(单聊时与 --open-dingtalk-id 二选一)
|
||||
--open-dingtalk-id string 接收人 openDingTalkId(单聊时与 --user 二选一,适用于三方应用等无法获取 userId 的场景)
|
||||
--title string 消息标题(可选,默认「消息」)
|
||||
--at-all @所有人(仅群聊时生效,可选,默认 false)
|
||||
--at-users string @指定成员的 userId 列表,逗号分隔(仅群聊时生效,可选)
|
||||
--media-id string 图片 mediaId(通过 dt_media_upload 工具上传获得,需从返回链接中去除 _宽_高.格式 后缀并加上 @ 前缀),指定后发送图片消息,不需要传文本内容
|
||||
|
||||
注意:
|
||||
- --text 和位置参数二选一,--text 优先
|
||||
- --group、--user、--open-dingtalk-id 三者互斥,只需指定其一:群聊用 --group,单聊用 --user 或 --open-dingtalk-id
|
||||
- --group 的别名: --id, --chat, --conversation-id (均可替代 --group)
|
||||
- --at-all 和 --at-users 仅在 --group 群聊时生效;当设置--at-all时,消息内容中一定要包含对应的占位符<@all>;当设置--at-users userId1,userId2时,消息内容中一定要包含对应格式的占位符<@userId1> <@userId2>
|
||||
- --media-id 指定图片 mediaId 时自动发送图片消息(msgType=image),不需要传 --text;图片单聊仅支持 --open-dingtalk-id,不支持 --user
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message send-personal — 发送个人消息
|
||||
|
||||
> ⚠️ 敏感操作:执行前必须向用户确认,同意后才加 `--yes`。
|
||||
|
||||
发送个人消息到指定会话或指定用户。支持 @指定人。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message send-personal [flags]
|
||||
Example:
|
||||
dws chat message send-personal --id <openConversationId> --content "你好" --type text
|
||||
dws chat message send-personal --open-id <openDingTalkId> --content "消息内容" --type text
|
||||
dws chat message send-personal --id <openConversationId> --content "内容" --at-all
|
||||
Flags:
|
||||
--content string 消息内容 (必填)
|
||||
--type string 消息类型,如 text、markdown (必填)
|
||||
--id string 群聊会话 ID openConversationId(与 --open-id 二选一)
|
||||
--open-id string 接收人 openDingTalkId(与 --id 二选一)
|
||||
--at-all @所有人(可选)
|
||||
--at-users string @指定人的 openDingTalkId 列表,逗号分隔(可选)
|
||||
|
||||
注意:
|
||||
- --id(群聊会话)和 --open-id(指定用户 openDingTalkId)二选一
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list — 拉取会话消息内容
|
||||
|
||||
拉取指定群聊或单聊的会话消息内容。
|
||||
|
||||
--group 指定群聊,--user 指定单聊用户(通过 userId),--open-dingtalk-id 指定单聊用户(通过 openDingTalkId),三者互斥。默认拉取给定时间之后的消息,--forward=false 拉之前的。hasMore=true 时用结果中的边界 createTime 作为下次 --time 翻页。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list [flags]
|
||||
Example:
|
||||
dws chat message list --group <openconversation_id> --time "2025-03-01 00:00:00"
|
||||
dws chat message list --user <userId> --time "2025-03-01 00:00:00" --limit 50
|
||||
dws chat message list --open-dingtalk-id <openDingTalkId> --time "2025-03-01 00:00:00" --limit 50
|
||||
dws chat message list --group <openconversation_id> --time "2025-03-01 00:00:00" --forward=false
|
||||
Flags:
|
||||
--forward true=拉给定时间之后的消息,false=拉给定时间之前的消息 (default true)
|
||||
--group string 群聊 openconversation_id(群聊时必填)
|
||||
--limit int 返回数量,不传则不限制
|
||||
--time string 开始时间,格式: yyyy-MM-dd HH:mm:ss (必填)
|
||||
--user string 单聊用户 userId(单聊时与 --open-dingtalk-id 二选一)
|
||||
--open-dingtalk-id string 单聊用户 openDingTalkId(单聊时与 --user 二选一,适用于三方应用等无法获取 userId 的场景)
|
||||
|
||||
注意:
|
||||
- --group、--user、--open-dingtalk-id 三者互斥,只需指定其一:群聊用 --group,单聊用 --user 或 --open-dingtalk-id
|
||||
- --group 的别名: --id, --chat, --conversation-id (均可替代 --group)
|
||||
- 翻页:hasMore=true 时,用结果中的边界 createTime 作为下次 --time
|
||||
- 如果返回的会话消息中包含 openConvThreadId 字段,说明是话题类消息,需要调用 dws chat message list-topic-replies 拉取话题的回复内容列表,openConvThreadId 作为 --topic-id 参数
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-all — 拉取指定时间范围内当前用户的所有会话消息
|
||||
|
||||
分页拉取当前登录用户在指定时间范围内的所有会话消息。
|
||||
|
||||
--start 和 --end 限定时间范围,--limit 指定每页数量,--cursor 传分页游标(首页传 "0",后续从响应中的 nextCursor 获取)。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-all [flags]
|
||||
Example:
|
||||
dws chat message list-all --start "2025-03-01 00:00:00" --end "2025-03-31 23:59:59" --limit 50
|
||||
dws chat message list-all --start "2025-03-01 00:00:00" --end "2025-03-31 23:59:59" --limit 50 --cursor "abc123token"
|
||||
Flags:
|
||||
--start string 起始时间,格式: yyyy-MM-dd HH:mm:ss (必填)
|
||||
--end string 结束时间,格式: yyyy-MM-dd HH:mm:ss (必填)
|
||||
--limit int 每页返回数量(默认 50)
|
||||
--cursor string 分页游标(首页传 "0",后续从响应中的 nextCursor 获取)
|
||||
|
||||
注意:
|
||||
- 四个参数每次请求都会传递给服务端,cursor 首页传 "0"
|
||||
- 与 chat message list 的区别:list 拉取指定单个会话(群聊或单聊)的消息,list-all 拉取当前用户所有会话的消息
|
||||
- 翻页:hasMore=true 时,用响应中的 nextCursor 值作为下次 --cursor 参数继续翻页
|
||||
- 时间格式统一为 yyyy-MM-dd HH:mm:ss
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-topic-replies — 拉取群话题回复消息列表
|
||||
|
||||
查询指定群聊中某条话题消息的全部回复。--group 指定群会话 ID,--topic-id 指定话题 ID(由 dws chat message list 返回)。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-topic-replies [flags]
|
||||
Example:
|
||||
dws chat message list-topic-replies --group <openconversation_id> --topic-id <topicId>
|
||||
dws chat message list-topic-replies --group <openconversation_id> --topic-id <topicId> --time "2025-03-01 00:00:00" --limit 20
|
||||
Flags:
|
||||
--group string 群会话 openconversationId (必填)
|
||||
--topic-id string 话题 ID,由 dws chat message list 返回 (必填)
|
||||
--time string 开始时间,格式: yyyy-MM-dd HH:mm:ss(可选)
|
||||
--limit int 返回数量(默认 50)
|
||||
--forward true=从老往新,false=从新往老(默认 false)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-by-sender — 拉取指定发送者的消息
|
||||
|
||||
搜索特定人发送给我的消息,返回结果包含单聊和群聊标识。--sender-user-id 指定发送者 userId,--sender-open-dingtalk-id 指定发送者 openDingTalkId,二者互斥。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-by-sender [flags]
|
||||
Example:
|
||||
dws chat message list-by-sender --sender-user-id <userId> --start "2026-03-10T00:00:00+08:00" --end "2026-03-11T00:00:00+08:00" --limit 50 --cursor 0
|
||||
dws chat message list-by-sender --sender-open-dingtalk-id <openDingTalkId> --start "2026-03-10T00:00:00+08:00" --end "2026-03-11T00:00:00+08:00" --limit 50 --cursor 0
|
||||
Flags:
|
||||
--sender-user-id string 发送者 userId(与 --sender-open-dingtalk-id 二选一)
|
||||
--sender-open-dingtalk-id string 发送者 openDingTalkId(与 --sender-user-id 二选一)
|
||||
--start string 开始时间,ISO-8601 格式 (必填)
|
||||
--end string 结束时间,ISO-8601 格式 (必填)
|
||||
--limit int 每页返回数量(默认 50)
|
||||
--cursor string 分页游标(默认 "0",翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- --sender-user-id 和 --sender-open-dingtalk-id 二者互斥,必须且只能指定其一
|
||||
- 不需要指定单聊/群聊,MCP 返回结果自带会话类型标识
|
||||
- 时间支持多种 ISO-8601 格式,如 "2026-03-10T00:00:00+08:00"、"2026-03-10 14:00:00" 等
|
||||
- 翻页:hasMore=true 时,用返回的 nextCursor 作为下次 --cursor
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-mentions — 拉取 @我 的消息
|
||||
|
||||
搜索时间范围内 @我 的消息,可选指定群聊。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-mentions [flags]
|
||||
Example:
|
||||
dws chat message list-mentions --start "2026-03-10T00:00:00+08:00" --end "2026-03-11T00:00:00+08:00" --limit 50 --cursor 0
|
||||
dws chat message list-mentions --group <openconversation_id> --start "2026-03-10T00:00:00+08:00" --end "2026-03-11T00:00:00+08:00" --limit 50 --cursor 0
|
||||
Flags:
|
||||
--group string 群聊 openconversation_id(可选,不传则查全部)
|
||||
--start string 开始时间,ISO-8601 格式 (必填)
|
||||
--end string 结束时间,ISO-8601 格式 (必填)
|
||||
--limit int 每页返回数量(默认 50)
|
||||
--cursor string 分页游标(默认 "0",翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- --group 可选,不传则查询所有会话中 @我 的消息;传入则只查指定群聊
|
||||
- --group 的别名: --id, --chat, --conversation-id (均可替代 --group)
|
||||
- 翻页:hasMore=true 时,用返回的 nextCursor 作为下次 --cursor
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-focused — 拉取特别关注人的消息
|
||||
|
||||
拉取当前用户特别关注人的消息。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-focused [flags]
|
||||
Example:
|
||||
dws chat message list-focused --limit 50
|
||||
dws chat message list-focused --limit 20 --cursor <nextCursor>
|
||||
Flags:
|
||||
--limit int 每页返回数量(默认 50)
|
||||
--cursor int64 分页游标(首次不传或传 0,翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- 首次调用不传 --cursor 或传 0,后续翻页传 nextCursor
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message list-unread-conversations — 获取未读会话列表
|
||||
|
||||
获取当前用户有未读消息的会话信息。可选通过 `--count` 限制返回条数。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message list-unread-conversations [flags]
|
||||
Example:
|
||||
dws chat message list-unread-conversations
|
||||
dws chat message list-unread-conversations --count 20
|
||||
Flags:
|
||||
--count int 返回未读会话条数(可选)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message search — 按关键词搜索消息
|
||||
|
||||
在当前用户的会话中按关键词搜索消息。--keyword 必填,可选 --group 限定搜索某个会话。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message search [flags]
|
||||
Example:
|
||||
dws chat message search --keyword "changefree" --start "2026-04-01T00:00:00+08:00" --end "2026-04-15T00:00:00+08:00" --limit 50 --cursor 0
|
||||
dws chat message search --keyword "codereview" --group <openconversation_id> --start "2026-04-01T00:00:00+08:00" --end "2026-04-15T00:00:00+08:00" --limit 100 --cursor 0
|
||||
Flags:
|
||||
--keyword string 搜索关键词 (必填)
|
||||
--group string 群聊 openconversation_id(可选,不传则搜索所有会话)
|
||||
--start string 开始时间,ISO-8601 格式 (必填)
|
||||
--end string 结束时间,ISO-8601 格式 (必填)
|
||||
--limit int 每页返回数量(默认 100)
|
||||
--cursor string 分页游标(默认 "0",翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- --group 可选,不传则搜索所有会话中的消息;传入则只搜索指定会话
|
||||
- --group 的别名: --id, --chat, --conversation-id (均可替代 --group)
|
||||
- 翻页:hasMore=true 时,用返回的 nextCursor 作为下次 --cursor
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message info — 获取会话信息
|
||||
|
||||
获取指定群聊或单聊会话的详情信息。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message info [flags]
|
||||
Example:
|
||||
dws chat message info --id <openConversationId>
|
||||
dws chat message info --open-id <openDingTalkId>
|
||||
Flags:
|
||||
--id string 群聊会话 ID openConversationId(与 --open-id 二选一)
|
||||
--open-id string 用户 openDingTalkId(单聊时与 --id 二选一)
|
||||
|
||||
注意:
|
||||
- --id(群聊)和 --open-id(单聊用户 openDingTalkId)二选一
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## list-top-conversations — 拉取置顶会话列表
|
||||
|
||||
拉取当前用户的置顶会话列表。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat list-top-conversations [flags]
|
||||
Example:
|
||||
dws chat list-top-conversations --limit 1000
|
||||
dws chat list-top-conversations --limit 1000 --cursor <nextCursor>
|
||||
Flags:
|
||||
--limit int 每页返回数量(默认 1000)
|
||||
--cursor int 分页游标(首次不传或传 0,翻页传 nextCursor)
|
||||
|
||||
注意:
|
||||
- 用户询问"置顶会话"时,直接调用此命令返回置顶会话列表即可
|
||||
- 用户询问"置顶消息"时,需两步:先调用此命令拉取置顶会话列表获取各会话的 openConversationId,再用 `chat message list --group <openConversationId>` 分别拉取每个会话内的消息
|
||||
- 翻页:hasMore=true 时,用返回的 nextCursor 作为下次 --cursor
|
||||
```
|
||||
|
||||
---
|
||||
@@ -126,161 +505,307 @@ Flags:
|
||||
Usage:
|
||||
dws chat bot search [flags]
|
||||
Example:
|
||||
dws chat bot search --name "考勤" --format json
|
||||
dws chat bot search --page 1
|
||||
dws chat bot search --page 1 --size 10 --name "日报"
|
||||
Flags:
|
||||
--name string 机器人名称(模糊搜索)
|
||||
--page int 页码(默认 1)
|
||||
--size int 每页数量(默认 50)
|
||||
--name string 按名称搜索
|
||||
--page int 页码,从1开始 (默认 1)
|
||||
--size int 每页条数 (默认 50),别名: --limit
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## bot create — 创建企业机器人
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat bot create [flags]
|
||||
Example:
|
||||
dws chat bot create --name "日报提醒机器人" --desc "负责每日日报提醒"
|
||||
Flags:
|
||||
--name string 机器人名称 (必填)
|
||||
--desc string 机器人描述(可选)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## bot search-groups — 搜索机器人所在群
|
||||
|
||||
搜索指定机器人已加入的群列表。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat bot search-groups [flags]
|
||||
Example:
|
||||
dws chat bot search-groups --keyword "项目"
|
||||
dws chat bot search-groups --keyword "冲刺" --cursor <nextCursor>
|
||||
Flags:
|
||||
--keyword string 搜索关键词 (必填)
|
||||
--cursor string 分页游标(首页留空,翻页传返回的 cursor)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## message send-by-bot — 机器人发消息
|
||||
|
||||
支持两种模式:群聊发送 和 批量单聊发送,通过 `--group` 和 `--users` 互斥区分。
|
||||
群聊:传 --group 指定群;单聊:传 --users 指定用户列表,二者只能选其一,不能同时指定。--text 支持 Markdown。
|
||||
|
||||
### 群聊发送
|
||||
```
|
||||
Usage:
|
||||
dws chat message send-by-bot [flags]
|
||||
Example:
|
||||
dws chat message send-by-bot --robot-code <code> --group <openConversationId> \
|
||||
--title "日报提醒" --text "请提交今日日报" --format json
|
||||
dws chat message send-by-bot --robot-code <robot-code> --group <openconversation_id> --title "日报" --text "## 今日完成..."
|
||||
dws chat message send-by-bot --robot-code <robot-code> --users userId1,userId2 --title "提醒" --text "请提交周报"
|
||||
Flags:
|
||||
--robot-code string 机器人 code (必填)
|
||||
--group string 群会话 ID (必填,与 --users 互斥)
|
||||
--group string 群聊 openConversationId(群聊时必填)
|
||||
--robot-code string 机器人 Code (必填)
|
||||
--text string 消息内容 Markdown (必填)
|
||||
--title string 消息标题 (必填)
|
||||
--text string 消息内容,支持 Markdown (必填)
|
||||
```
|
||||
--users string 用户 userId 列表,逗号分隔,最多20个(单聊时必填)
|
||||
|
||||
### 批量单聊发送
|
||||
注意:
|
||||
- --group 与 --users 互斥,必须且只能指定其一
|
||||
- --group 的别名: --id, --chat, --conversation-id (均可替代 --group)
|
||||
```
|
||||
Usage:
|
||||
dws chat message send-by-bot [flags]
|
||||
Example:
|
||||
dws chat message send-by-bot --robot-code <code> --users "user1,user2" \
|
||||
--title "通知" --text "会议已取消" --format json
|
||||
Flags:
|
||||
--robot-code string 机器人 code (必填)
|
||||
--users string 用户 ID 列表,逗号分隔,最多 20 个 (必填,与 --group 互斥)
|
||||
--title string 消息标题 (必填)
|
||||
--text string 消息内容,支持 Markdown (必填)
|
||||
```
|
||||
|
||||
> ⚠️ `--group` 和 `--users` 互斥:群聊用 `--group`,单聊用 `--users`,不能同时传。
|
||||
|
||||
---
|
||||
|
||||
## message recall-by-bot — 机器人撤回消息
|
||||
|
||||
支持两种模式:群聊撤回 和 批量单聊撤回。
|
||||
群聊:传 --group 与 --keys;单聊:仅传 --keys。--keys 为发送时返回的 processQueryKey 列表,逗号分隔。
|
||||
|
||||
### 群聊撤回
|
||||
```
|
||||
Usage:
|
||||
dws chat message recall-by-bot [flags]
|
||||
Example:
|
||||
dws chat message recall-by-bot --robot-code <code> --group <openConversationId> \
|
||||
--keys "key1,key2" --format json
|
||||
dws chat message recall-by-bot --robot-code <robot-code> --group <openconversation_id> --keys <process-query-key>
|
||||
dws chat message recall-by-bot --robot-code <robot-code> --keys key1,key2
|
||||
Flags:
|
||||
--robot-code string 机器人 code (必填)
|
||||
--group string 群会话 ID (必填,与批量单聊互斥)
|
||||
--keys string 消息 key 列表,逗号分隔 (必填)
|
||||
--group string 群聊 openConversationId(群聊撤回时必填)
|
||||
--keys string 消息 processQueryKey 列表,逗号分隔 (必填)
|
||||
--robot-code string 机器人 Code (必填)
|
||||
```
|
||||
|
||||
### 批量单聊撤回
|
||||
```
|
||||
Usage:
|
||||
dws chat message recall-by-bot [flags]
|
||||
Example:
|
||||
dws chat message recall-by-bot --robot-code <code> --keys "key1,key2" --format json
|
||||
Flags:
|
||||
--robot-code string 机器人 code (必填)
|
||||
--keys string 消息 key 列表,逗号分隔 (必填)
|
||||
```
|
||||
|
||||
> ⚠️ 消息 key 从 `send-by-bot` 返回结果中提取。
|
||||
|
||||
---
|
||||
|
||||
## message send-by-webhook — 自定义机器人 Webhook 发消息
|
||||
|
||||
@ 人时需在 --text 中包含 @userId 或 @手机号,否则 @ 不生效。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws chat message send-by-webhook [flags]
|
||||
Example:
|
||||
dws chat message send-by-webhook --token <robotToken> \
|
||||
--title "告警" --text "CPU 使用率超过 90%" --format json
|
||||
dws chat message send-by-webhook --token <webhook-token> --title "告警" --text "CPU 超 90%" --at-all
|
||||
dws chat message send-by-webhook --token <webhook-token> --title "test" --text "hi @118785" --at-users 118785
|
||||
Flags:
|
||||
--token string 自定义机器人 Webhook Token (必填)
|
||||
--title string 消息标题 (必填)
|
||||
--at-all @ 所有人
|
||||
--at-mobiles string @ 指定手机号,逗号分隔
|
||||
--at-users string @ 指定用户,逗号分隔(需在 text 中包含 @userId)
|
||||
--text string 消息内容 (必填)
|
||||
--at-all @所有人
|
||||
--at-mobiles string @指定手机号列表,逗号分隔
|
||||
--at-users string @指定用户 ID 列表,逗号分隔
|
||||
--title string 消息标题 (必填)
|
||||
--token string Webhook Token (必填)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 意图判断
|
||||
|
||||
- 用户说"搜索一个群" → `search`
|
||||
- 用户说"帮我建个群" → `group create`
|
||||
- 用户说"看看群里有谁" → `group members list`
|
||||
- 用户说"把张三拉进群" → 先 `contact user search` 获取 userId,再 `group members add`
|
||||
- 用户说"把张三移出群" → 先 `contact user search` 获取 userId,再 `group members remove`(⚠️ 需确认)
|
||||
- 用户说"改一下群名" → `group rename`
|
||||
- 用户说"让机器人在群里发通知" → `message send-by-bot --group`
|
||||
- 用户说"机器人给张三发消息" → 先 `contact user search` 获取 userId,再 `message send-by-bot --users`
|
||||
- 用户说"通过 Webhook 发告警" / 用户有 Webhook Token → `message send-by-webhook`
|
||||
- 用户说"撤回机器人消息" → `message recall-by-bot`
|
||||
- 用户说"查一下我的机器人" → `bot search`
|
||||
- 用户说"把机器人加到群里" → `group members add-bot`
|
||||
用户说"建群/创建群聊" → `chat group create`
|
||||
用户说"创建企业全员群/组织群" → `chat group create-org`
|
||||
用户说"搜索群/找群" → `chat search`
|
||||
用户说"群成员/看群里有谁" → `chat group members list`
|
||||
用户说"拉人进群/加群成员" → `chat group members add`
|
||||
用户说"踢人/移除群成员" → `chat group members remove`
|
||||
用户说"加机器人到群" → `chat group members add-bot`
|
||||
用户说"改群名" → `chat group rename`
|
||||
用户说"聊天记录/会话消息/拉取会话" → `chat message list`
|
||||
用户说"某人发给我的消息/指定发送者/某人的消息" → `chat message list-by-sender`(用户未明确说"单聊"时优先使用,跨单聊/群聊)
|
||||
用户说"拉取和某人的单聊记录/单聊消息" → `chat message list --user`(用户明确说"单聊"时使用)
|
||||
用户说"@我的消息/at我的/提及我的" → `chat message list-mentions`
|
||||
用户说"未读消息会话/未读会话列表/我的未读会话" → `chat message list-unread-conversations`
|
||||
用户说"发群消息(以个人身份)" → `chat message send --group`
|
||||
用户说"发单聊消息(以个人身份)" → `chat message send --user`(有 userId 时)或 `chat message send --open-dingtalk-id`(有 openDingTalkId 时)
|
||||
用户说"发个人消息/个人通知" → `chat message send-personal`(⚠️ 敏感操作,需确认)
|
||||
用户说"机器人发消息/机器人群发" → `chat message send-by-bot`
|
||||
用户说"机器人撤回消息" → `chat message recall-by-bot`
|
||||
用户说"Webhook 发消息/告警消息" → `chat message send-by-webhook`
|
||||
用户说"话题回复/群话题消息回复/拉取话题回复" → `chat message list-topic-replies`
|
||||
用户说"所有消息/全部会话消息/拉取全部消息/时间范围内消息/我的消息/我今天的消息/查我的钉钉消息/最近的消息" → `chat message list-all`
|
||||
用户说"特别关注人的消息/关注的人的消息/星标联系人的消息" → `chat message list-focused`
|
||||
用户说"查看我的机器人" → `chat bot search`
|
||||
用户说"创建机器人" → `chat bot create`
|
||||
用户说"搜索消息/查找关键词/搜一下消息里的XX" → `chat message search`
|
||||
用户说"我和XX的共同群/我们都在哪些群/查共同群" → `chat search-common`
|
||||
用户说"置顶会话/置顶消息/我的置顶/查看置顶" → `chat list-top-conversations`
|
||||
用户说"获取会话信息/会话详情" → `chat message info`
|
||||
用户说"机器人在哪些群/机器人的群" → `chat bot search-groups`
|
||||
|
||||
**关键区分**: `send-by-bot`(企业内部机器人,需 robotCode) vs `send-by-webhook`(自定义机器人 Webhook,需 token)
|
||||
关键区分:
|
||||
- `chat message list` — 拉取指定会话的消息(需指定 --group 或 --user),按时间点 + 方向翻页
|
||||
- `chat message list --user` — list 的单聊模式,拉取与指定用户的单聊记录(用户明确说"单聊""私聊"时使用)
|
||||
- `chat message list-by-sender` — 搜索指定发送者发给我的消息,跨所有会话(单聊+群聊均包含,用户只说"某人发的消息"时优先使用)
|
||||
- `chat message list-mentions` — 拉取 @我 的消息(跨单聊/群聊,可选指定群)
|
||||
- `chat message list-unread-conversations` — 拉取当前用户存在未读消息的会话列表(可选 `--count`)
|
||||
- `chat message list-all` — 拉取当前用户所有会话的消息,按时间范围 + cursor 分页。只要用户没有指定某个具体的会话(如某个群名、某个人名),即使提到"单聊消息""群聊消息"等笼统范围,也应路由到此命令
|
||||
- `chat message list-topic-replies` — 拉取群话题的回复消息列表
|
||||
- `chat message list-focused` — 拉取特别关注人的消息,cursor 分页
|
||||
- `chat list-top-conversations` — 拉取置顶会话列表(用户询问"置顶会话"或"置顶消息"时路由到此),cursor 分页
|
||||
- `chat message send` — 以**当前用户**身份发消息(群聊或单聊),text 为位置参数;支持 --media-id 发送图片消息
|
||||
- `chat message send-personal` — 发送个人消息,支持通过 openConversationId 或 openDingTalkId 指定目标(⚠️ 敏感操作)
|
||||
- `chat message search` — 按关键词搜索消息内容(跨所有会话,可选指定群)
|
||||
- `chat search-common` — 搜索共同群,查询指定人共同所在的群聊(AND=所有人都在,OR=任一人在)
|
||||
- `chat message send-by-bot` — 以**机器人**身份发消息(群聊或单聊),text 为 --text flag
|
||||
- `chat message send-by-webhook` — 通过**自定义机器人 Webhook** 发群消息
|
||||
- `chat message recall-by-bot` — 通过机器人撤回已发送的消息
|
||||
- `chat message info` — 获取指定会话的详情信息
|
||||
- `chat bot create` — 创建新的企业机器人
|
||||
- `chat bot search-groups` — 搜索机器人所在群列表
|
||||
|
||||
## 核心工作流
|
||||
|
||||
```bash
|
||||
# ── 工作流: 建群并添加机器人 ──
|
||||
# 1. 搜索群 — 提取 openconversation_id
|
||||
dws chat search --query "项目冲刺" --format json
|
||||
|
||||
# 1. 搜索同事 userId
|
||||
dws contact user search --keyword "张三" --format json
|
||||
# 2. 拉取群消息
|
||||
dws chat message list --group <openconversation_id> --time "2025-03-01 00:00:00" --format json
|
||||
|
||||
# 2. 创建群
|
||||
dws chat group create --name "项目群" --users <userId1>,<userId2> --format json
|
||||
# 2b. 拉取未读会话列表
|
||||
dws chat message list-unread-conversations --count 20 --format json
|
||||
|
||||
# 3. 搜索机器人
|
||||
dws chat bot search --format json
|
||||
# 3. 以个人身份发送群消息
|
||||
dws chat message send --group <openconversation_id> --title "周报提醒" "请大家本周五前提交周报" --format json
|
||||
|
||||
# 4. 添加机器人到群
|
||||
dws chat group members add-bot --id <openConversationId> --robot-code <code> --format json
|
||||
# 4. 以个人身份单聊(通过 userId)
|
||||
dws chat message send --user <userId> "你好" --format json
|
||||
|
||||
# 4b. 以个人身份单聊(通过 openDingTalkId,三方应用等无法获取 userId 时使用)
|
||||
dws chat message send --open-dingtalk-id <openDingTalkId> "你好" --format json
|
||||
|
||||
# 5. 机器人发群消息(Markdown)
|
||||
dws chat message send-by-bot --robot-code <robot-code> \
|
||||
--group <openconversation_id> --title "日报" --text "## 今日完成..." --format json
|
||||
|
||||
# 6. 机器人单聊发消息
|
||||
dws chat message send-by-bot --robot-code <robot-code> \
|
||||
--users userId1,userId2 --title "提醒" --text "请提交周报" --format json
|
||||
|
||||
# 7. Webhook 发告警
|
||||
dws chat message send-by-webhook --token <webhook-token> \
|
||||
--title "告警" --text "CPU 超 90%" --at-all --format json
|
||||
```
|
||||
|
||||
```bash
|
||||
# ── 工作流: 机器人群发消息 ──
|
||||
## 复合工作流
|
||||
|
||||
# 1. 搜索可用机器人
|
||||
### 机器人发消息后撤回(完整流程)
|
||||
|
||||
撤回只能用于 `send-by-bot` 发出的消息。个人身份 (`chat message send`) 发出的消息**无法通过 API 撤回**。
|
||||
|
||||
```bash
|
||||
# Step 1: 查我的机器人 — 提取 robot-code
|
||||
dws chat bot search --format json
|
||||
|
||||
# 2. 发送群消息
|
||||
dws chat message send-by-bot --robot-code <code> --group <groupId> \
|
||||
# Step 2: 用机器人发消息 — 提取返回中的 processQueryKey
|
||||
dws chat message send-by-bot --robot-code <robot-code> --group <openconversation_id> \
|
||||
--title "通知" --text "内容" --format json
|
||||
|
||||
# Step 3: 用同一个 robot-code + processQueryKey 撤回
|
||||
dws chat message recall-by-bot --robot-code <robot-code> --group <openconversation_id> \
|
||||
--keys <processQueryKey> --format json
|
||||
```
|
||||
|
||||
```bash
|
||||
# ── 工作流: Webhook 告警 ──
|
||||
### 创建并使用机器人(完整流程)
|
||||
|
||||
# 直接通过 Webhook Token 发送
|
||||
dws chat message send-by-webhook --token <token> \
|
||||
--title "告警" --text "服务异常" --at-all --format json
|
||||
```bash
|
||||
# Step 1: 创建机器人
|
||||
dws chat bot create --name "项目提醒机器人" --desc "项目状态提醒" --format json
|
||||
|
||||
# Step 2: 搜索群 — 提取 openConversationId
|
||||
dws chat search --query "项目群" --format json
|
||||
|
||||
# Step 3: 将机器人添加到群
|
||||
dws chat group members add-bot --id <openConversationId> --robot-code <robotCode> --format json
|
||||
|
||||
# Step 4: 机器人发消息
|
||||
dws chat message send-by-bot --robot-code <robotCode> --group <openConversationId> \
|
||||
--title "提醒" --text "请及时更新项目状态" --format json
|
||||
```
|
||||
|
||||
### 机器人 @指定人发群消息
|
||||
|
||||
`--text` 中**必须**包含 `<@userId>` 占位符,否则 @ 不生效。
|
||||
|
||||
```bash
|
||||
# Step 1: 搜人获取 userId
|
||||
dws aisearch person --keyword "张三" --dimension name --format json
|
||||
|
||||
# Step 2: 取 userId 发送(注意 text 中的占位符)
|
||||
dws chat message send-by-bot --robot-code <robot-code> --group <openconversation_id> \
|
||||
--title "提醒" --text "<@userId1> <@userId2> 请查收本周报告" --format json
|
||||
```
|
||||
|
||||
### 发送图片/文件消息(跨产品: drive → chat)
|
||||
|
||||
```bash
|
||||
# Step 1: 上传文件到钉盘 — 获取 uploadId 和凭证
|
||||
dws drive upload-info --file-name "截图.png" --file-size <字节数> --format json
|
||||
|
||||
# Step 2: HTTP PUT 上传文件到 OSS
|
||||
curl -X PUT -T "截图.png" "<upload-info 返回的上传 URL>"
|
||||
|
||||
# Step 3: 提交上传 — 获取 dentryUuid
|
||||
dws drive commit --file-name "截图.png" --file-size <字节数> --upload-id <uploadId> --format json
|
||||
|
||||
# Step 4: 获取下载链接
|
||||
dws drive download --file-id <dentryUuid> --format json
|
||||
|
||||
# Step 5: 用 Markdown 图片语法发送
|
||||
dws chat message send --group <openconversation_id> \
|
||||
--text "" --format json
|
||||
```
|
||||
|
||||
## 上下文传递表
|
||||
|
||||
| 操作 | 从返回中提取 | 用于 |
|
||||
|------|-------------|------|
|
||||
| `search` | openConversationId | `group members` / `group rename` / `group members add` / `send-by-bot --group` |
|
||||
| `group create` | openConversationId | 同上 |
|
||||
| `bot search` | robotCode | `send-by-bot` / `recall-by-bot` / `add-bot` |
|
||||
| `message send-by-bot` | processQueryKey | `recall-by-bot --keys` |
|
||||
| `chat search` | `openConversationId` | message send/list、group members 等的 --group |
|
||||
| `chat group create` | `openConversationId` | 同上 |
|
||||
| `chat message list-all` | `nextCursor` | 下次 list-all 的 --cursor |
|
||||
| `aisearch person` | `userId` | message send 的 --user、--at-users、send-by-bot 的 --users、list-by-sender 的 --sender-user-id |
|
||||
| `aisearch person` → `contact user get` | `openDingTalkId` | list-by-sender 的 --sender-open-dingtalk-id、message send/list 的 --open-dingtalk-id |
|
||||
| `chat bot search` | `robotCode` | send-by-bot / recall-by-bot 的 --robot-code、group members add-bot 的 --robot-code |
|
||||
| `chat bot create` | `robotCode` | send-by-bot / recall-by-bot 的 --robot-code |
|
||||
| `chat message send-by-bot` | `processQueryKey` | recall-by-bot 的 --keys |
|
||||
| `chat message search` | `nextCursor` | 下次 message search 的 --cursor |
|
||||
| `chat search-common` | `openConversationId` | message send/list 等的 --group |
|
||||
| `drive download` | 下载链接 | message send 的 Markdown 图片/链接语法 |
|
||||
|
||||
## 注意事项
|
||||
|
||||
- `--group` 为群聊会话 ID (openconversation_id),可从群搜索或群聊信息中获取
|
||||
- `chat message send` 的 text 是位置参数(恰好 1 个),非 flag;群聊用 `--group`,单聊用 `--user`(userId)或 `--open-dingtalk-id`(openDingTalkId),三者互斥;`--at-all`、`--at-users` 仅在 `--group` 群聊时生效;发送图片消息用 `--media-id`
|
||||
- `chat message send-personal` 为敏感操作(isSensitive),执行前需用户明确确认
|
||||
- `chat message list-all` 的四个参数(--start、--end、--limit、--cursor)每次请求都必须传递;翻页时用响应中的 nextCursor 值作为下次 --cursor
|
||||
- `chat message list` 的 `--group`、`--user`、`--open-dingtalk-id` 三者互斥,必须且只能指定其一
|
||||
- `chat message list-by-sender` 不需要指定单聊/群聊,返回结果自带会话类型标识
|
||||
- `chat message list-mentions` 可选 `--group` 指定群聊,不传则查全部
|
||||
- `chat message list-unread-conversations` 获取当前用户未读会话列表,可选 `--count` 指定返回条数
|
||||
- `chat message search` 按关键词搜索消息内容,`--keyword` 必填,可选 `--group` 限定搜索某个会话
|
||||
- `chat search-common` 搜索共同群,`--nicks` 传人员昵称(逗号分隔),`--match-mode` AND/OR 控制匹配逻辑
|
||||
- `chat list-top-conversations` 拉取置顶会话列表,分页用 `--limit`(默认 1000)/`--cursor`
|
||||
- `send-by-bot` 群聊传 `--group`,单聊传 `--users`,二者互斥且必选其一
|
||||
- `recall-by-bot` 群聊传 `--group` + `--keys`,单聊仅传 `--keys`(不传 `--group` 即为单聊撤回)
|
||||
- `send-by-webhook` 支持 `--at-all`、`--at-mobiles`、`--at-users` 进行 @ 操作,但需在 `--text` 中包含 `@userId` 或 `@手机号` 才能生效
|
||||
|
||||
## 自动化脚本
|
||||
|
||||
| 脚本 | 场景 | 用法 |
|
||||
|------|------|------|
|
||||
| [chat_export_messages.py](../../scripts/chat_export_messages.py) | 导出群聊消息到 JSON 文件 | `python chat_export_messages.py --query "项目冲刺" --time "2026-03-10 00:00:00"` |
|
||||
| [chat_history_with_user.py](../../scripts/chat_history_with_user.py) | 查询与某人的单聊聊天记录 | `python chat_history_with_user.py --name "张三" --time "2026-03-10 00:00:00"` |
|
||||
|
||||
## 相关产品
|
||||
|
||||
- [contact](./contact.md) — 搜索同事/好友,获取 userId 用于 --user、--at-users、send-by-bot --users、list-by-sender --sender-user-id;获取 openDingTalkId 用于 list-by-sender --sender-open-dingtalk-id、--open-dingtalk-id
|
||||
- [drive](./drive.md) — 上传文件获取下载链接,用于 Markdown 图片/文件消息
|
||||
|
||||
@@ -0,0 +1,462 @@
|
||||
# 文档 (doc) 命令参考
|
||||
|
||||
## 命令总览
|
||||
|
||||
### 搜索文档
|
||||
```
|
||||
Usage:
|
||||
dws doc search [flags]
|
||||
Example:
|
||||
dws doc search --query "会议纪要"
|
||||
dws doc search
|
||||
dws doc search --extensions pdf,docx
|
||||
dws doc search --query "方案" --created-from 1700000000000 --created-to 1710000000000
|
||||
dws doc search --creator-uids uid1,uid2
|
||||
dws doc search --workspace-ids wsId1,wsId2
|
||||
Flags:
|
||||
--query string 搜索关键词 (不传则返回最近访问)
|
||||
--extensions strings 按文件扩展名过滤,不含点号,逗号分隔 (如 pdf,docx,png)。支持的在线文档类型后缀名: adoc=文字, axls=表格, appt=演示文稿, awbd=白板, adraw=画板, amind=脑图, able=多维表格, aform=收集表
|
||||
--created-from int 创建时间起始 (毫秒时间戳,含)
|
||||
--created-to int 创建时间截止 (毫秒时间戳,含)
|
||||
--visited-from int 访问时间起始 (毫秒时间戳,含)
|
||||
--visited-to int 访问时间截止 (毫秒时间戳,含)
|
||||
--creator-uids strings 按创建者用户 ID 过滤,逗号分隔
|
||||
--editor-uids strings 按编辑者用户 ID 过滤,逗号分隔
|
||||
--mentioned-uids strings 按 @提及的用户 ID 过滤,逗号分隔
|
||||
--workspace-ids strings 按知识库 ID 过滤,支持知识库 URL,逗号分隔
|
||||
--page-size int 每页数量 (默认 10,最大 30)
|
||||
--page-token string 分页游标 (从上次结果的 nextPageToken 获取)
|
||||
```
|
||||
|
||||
### 遍历文件列表
|
||||
```
|
||||
Usage:
|
||||
dws doc list [flags]
|
||||
Example:
|
||||
dws doc list
|
||||
dws doc list --folder <FOLDER_ID>
|
||||
dws doc list --workspace <WS_ID> --page-size 20
|
||||
Flags:
|
||||
--folder string 文件夹 ID 或 URL
|
||||
--workspace string 知识库 ID
|
||||
--page-size int 每页数量 (默认 50,最大 50)
|
||||
--page-token string 分页游标 (从上次结果的 nextPageToken 获取)
|
||||
```
|
||||
|
||||
### 获取文档元信息
|
||||
```
|
||||
Usage:
|
||||
dws doc info [flags]
|
||||
Example:
|
||||
dws doc info --node <DOC_ID>
|
||||
dws doc info --node "https://alidocs.dingtalk.com/i/nodes/<DOC_UUID>"
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
```
|
||||
|
||||
### 读取文档内容
|
||||
```
|
||||
Usage:
|
||||
dws doc read [flags]
|
||||
Example:
|
||||
dws doc read --node <DOC_ID>
|
||||
dws doc read --node "https://alidocs.dingtalk.com/i/nodes/<DOC_UUID>"
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
```
|
||||
|
||||
### 创建文档
|
||||
```
|
||||
Usage:
|
||||
dws doc create [flags]
|
||||
Example:
|
||||
dws doc create --name "项目周报"
|
||||
dws doc create --name "Q1 总结" --markdown "# Q1 总结" --folder <FOLDER_ID>
|
||||
dws doc create --name "知识库文档" --workspace <WS_ID>
|
||||
Flags:
|
||||
--name string 文档名称 (必填)
|
||||
--folder string 目标文件夹 ID 或 URL
|
||||
--workspace string 目标知识库 ID
|
||||
--markdown string 文档初始 Markdown 内容
|
||||
```
|
||||
|
||||
### 更新文档内容
|
||||
```
|
||||
Usage:
|
||||
dws doc update [flags]
|
||||
Example:
|
||||
dws doc update --node <DOC_ID> --markdown "# 追加内容" --mode append
|
||||
dws doc update --node <DOC_ID> --markdown "# 完整替换" --mode overwrite
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
--markdown string Markdown 内容 (必填)
|
||||
--mode string 更新模式: overwrite=覆盖, append=追加 (默认 append)
|
||||
```
|
||||
|
||||
### 上传文件到钉钉文档或钉钉知识库
|
||||
```
|
||||
Usage:
|
||||
dws doc upload [flags]
|
||||
Example:
|
||||
dws doc upload --file ./report.pdf
|
||||
dws doc upload --file ./slides.pptx --name "Q1汇报.pptx" --folder <FOLDER_ID>
|
||||
dws doc upload --file ./data.xlsx --workspace <WS_ID> --convert
|
||||
Flags:
|
||||
--file string 本地文件路径 (必填)
|
||||
--name string 文件显示名称 (默认使用文件名)
|
||||
--folder string 目标文件夹 ID 或 URL
|
||||
--workspace string 目标知识库 ID
|
||||
--convert 是否转换为钉钉在线文档
|
||||
```
|
||||
|
||||
### 下载文件到本地
|
||||
```
|
||||
Usage:
|
||||
dws doc download [flags]
|
||||
Example:
|
||||
dws doc download --node <NODE_ID>
|
||||
dws doc download --node <NODE_ID> --output ./report.pdf
|
||||
dws doc download --node "https://alidocs.dingtalk.com/i/nodes/<DOC_UUID>" --output ~/downloads/
|
||||
Flags:
|
||||
--node string 文件节点 ID 或 URL (必填)
|
||||
--output string 本地保存路径 (文件路径或目录,必填)
|
||||
```
|
||||
|
||||
### 创建文件夹
|
||||
```
|
||||
Usage:
|
||||
dws doc folder create [flags]
|
||||
Example:
|
||||
dws doc folder create --name "项目资料"
|
||||
dws doc folder create --name "子文件夹" --folder <PARENT_FOLDER_ID>
|
||||
Flags:
|
||||
--name string 文件夹名称 (必填)
|
||||
--folder string 父文件夹 ID 或 URL
|
||||
--workspace string 目标知识库 ID
|
||||
```
|
||||
|
||||
### 查询块元素
|
||||
```
|
||||
Usage:
|
||||
dws doc block list [flags]
|
||||
Example:
|
||||
dws doc block list --node <DOC_ID>
|
||||
dws doc block list --node <DOC_ID> --start-index 0 --end-index 5
|
||||
dws doc block list --node <DOC_ID> --block-type heading
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
--start-index int 起始位置 (从 0 开始)
|
||||
--end-index int 终止位置 (含)
|
||||
--block-type string 按块类型过滤
|
||||
```
|
||||
|
||||
### 插入块元素
|
||||
```
|
||||
Usage:
|
||||
dws doc block insert [flags]
|
||||
Example:
|
||||
dws doc block insert --node <DOC_ID> --text "这是一段文字"
|
||||
dws doc block insert --node <DOC_ID> --heading "二级标题" --level 2
|
||||
dws doc block insert --node <DOC_ID> --element '{"blockType":"paragraph","paragraph":{"text":"内容"}}'
|
||||
dws doc block insert --node <DOC_ID> --text "在此处之前插入" --ref-block <BLOCK_ID> --where before
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
--text string 快捷: 段落文本内容
|
||||
--heading string 快捷: 标题文本
|
||||
--level int 标题级别 1-6 (配合 --heading,默认 1)
|
||||
--element string 块元素 JSON (高级)
|
||||
--index int 参照位置索引 (从 0 开始)
|
||||
--where string 插入方向: before / after (默认 after)
|
||||
--ref-block string 参照块 ID (优先级高于 --index)
|
||||
```
|
||||
|
||||
### 更新块元素
|
||||
```
|
||||
Usage:
|
||||
dws doc block update [flags]
|
||||
Example:
|
||||
dws doc block update --node <DOC_ID> --block-id <BLOCK_ID> --text "新内容"
|
||||
dws doc block update --node <DOC_ID> --block-id <BLOCK_ID> --element '{"blockType":"heading","heading":{"text":"新标题","level":1}}'
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
--block-id string 目标块 ID (必填)
|
||||
--text string 快捷: 段落文本内容
|
||||
--heading string 快捷: 标题文本
|
||||
--level int 标题级别 1-6 (配合 --heading,默认 1)
|
||||
--element string 块元素 JSON (高级)
|
||||
```
|
||||
|
||||
### 删除块元素
|
||||
|
||||
> **CAUTION:** 不可逆操作 — 执行前必须向用户确认。
|
||||
|
||||
```
|
||||
Usage:
|
||||
dws doc block delete [flags]
|
||||
Example:
|
||||
dws doc block delete --node <DOC_ID> --block-id <BLOCK_ID> --yes
|
||||
Flags:
|
||||
--node string 文档 ID 或 URL (必填)
|
||||
--block-id string 目标块 ID (必填)
|
||||
```
|
||||
|
||||
### 查询文档评论列表
|
||||
```
|
||||
Usage:
|
||||
dws doc comment list [flags]
|
||||
Example:
|
||||
dws doc comment list --node <DOC_ID>
|
||||
dws doc comment list --node <DOC_ID> --type inline --resolve-status unresolved
|
||||
dws doc comment list --node <DOC_ID> --page-size 20 --next-token <TOKEN>
|
||||
Flags:
|
||||
--node string 目标文档的标识,支持传入 URL 或 ID (必填)
|
||||
--page-size int 每页返回的评论数量,默认 50,最大 50
|
||||
--next-token string 分页游标,从上一次请求的返回结果中获取 (首次请求不传)
|
||||
--type string 按评论类型过滤: global (全文评论) / inline (划词评论)
|
||||
--resolve-status string 按解决状态过滤: resolved (已解决) / unresolved (未解决)
|
||||
```
|
||||
|
||||
### 创建文档评论
|
||||
```
|
||||
Usage:
|
||||
dws doc comment create [flags]
|
||||
Example:
|
||||
dws doc comment create --node <DOC_ID> --content "这里需要修改"
|
||||
dws doc comment create --node <DOC_ID> --content "请review" --mention uid1,uid2
|
||||
Flags:
|
||||
--node string 目标文档的标识,支持传入 URL 或 ID (必填)
|
||||
--content string 评论的文字内容,纯文本 (必填)
|
||||
--mention string 被 @ 的用户 uid 列表,逗号分隔
|
||||
```
|
||||
|
||||
### 回复文档评论
|
||||
```
|
||||
Usage:
|
||||
dws doc comment reply [flags]
|
||||
Example:
|
||||
dws doc comment reply --node <DOC_ID> --comment-key <COMMENT_KEY> --content "同意"
|
||||
dws doc comment reply --node <DOC_ID> --comment-key <COMMENT_KEY> --content "比心" --emoji
|
||||
dws doc comment reply --node <DOC_ID> --comment-key <COMMENT_KEY> --content "请确认" --mention uid1,uid2
|
||||
Flags:
|
||||
--node string 目标文档的标识,支持传入 URL 或 ID (必填)
|
||||
--content string 回复的文字内容,表情回复时填写表情名称 (必填)
|
||||
--comment-key string 被回复评论的 commentKey,格式: {13位毫秒时间戳}{32位UUID},可从 list/create 结果获取 (必填)
|
||||
--emoji 设为 true 时作为表情贴图回复 (默认 false)
|
||||
--mention string 被 @ 的用户 uid 列表,逗号分隔
|
||||
```
|
||||
|
||||
## URL 识别与 DOC_ID 提取
|
||||
|
||||
当用户输入包含钉钉文档 URL 时,**必须先识别并提取 DOC_ID**,再判断意图。
|
||||
|
||||
### 支持的 URL 格式
|
||||
|
||||
| 格式 | 示例 | DOC_ID 提取方式 |
|
||||
|------|------|----------------|
|
||||
| `alidocs.dingtalk.com/i/nodes/{id}` | `https://alidocs.dingtalk.com/i/nodes/9E05BDRVQePjzLkZt2p2vE7kV63zgkYA` | 取 URL 路径最后一段:`9E05BDRVQePjzLkZt2p2vE7kV63zgkYA` |
|
||||
| `alidocs.dingtalk.com/i/nodes/{id}?queryParams` | `https://alidocs.dingtalk.com/i/nodes/abc123?doc_type=wiki_doc` | 忽略 query 参数,取路径最后一段:`abc123` |
|
||||
|
||||
### 提取规则
|
||||
|
||||
1. 匹配 URL 中 `alidocs.dingtalk.com` 域名
|
||||
2. 取 URL path 的最后一段作为 DOC_ID(去掉 query string 和 fragment)
|
||||
3. 提取出的 DOC_ID 可直接用于所有 `--node` 参数,也可将完整 URL 传给 `--node`(CLI 会自动解析)
|
||||
|
||||
### 处理流程
|
||||
|
||||
```
|
||||
用户输入含 alidocs.dingtalk.com URL
|
||||
→ 提取 DOC_ID(URL 路径最后一段)
|
||||
→ 结合用户意图选择命令(默认 read)
|
||||
→ 将 DOC_ID 传给 --node 参数
|
||||
```
|
||||
|
||||
## 意图判断
|
||||
|
||||
用户说"找文档/搜文档/最近文档":
|
||||
- 搜索 → `search`
|
||||
- 浏览 → `list`
|
||||
|
||||
用户说"看文档/读内容/文档内容":
|
||||
- 读取 → `read` (需文档 ID 或 URL)
|
||||
- 元信息 → `info`
|
||||
|
||||
用户说"写文档/创建文档":
|
||||
- 新建 → `create`
|
||||
- 追加内容 → `update --mode append`
|
||||
- 覆盖替换 → `update --mode overwrite`
|
||||
|
||||
用户说"建文件夹/新建目录":
|
||||
- 创建 → `folder create`
|
||||
|
||||
用户说"上传文件/传文件/上传到文档/上传到知识库":
|
||||
- 上传 → `upload`(需本地文件路径)
|
||||
- 上传并转换 → `upload --convert`
|
||||
|
||||
用户说"下载文件/导出文件/下载到本地":
|
||||
- 下载 → `download`(需文件节点 ID 或 URL)
|
||||
|
||||
用户说"编辑块/改段落/插入标题/删除块":
|
||||
- 查看结构 → `block list`
|
||||
- 插入 → `block insert`
|
||||
- 修改 → `block update`
|
||||
- 删除 → `block delete`
|
||||
|
||||
**用户直接粘贴文档 URL(无其他指令)**:
|
||||
- 默认 → `read`(读取文档内容)
|
||||
- 如 URL 明显是文件夹 → `list`(列出文件夹内容)
|
||||
|
||||
**用户粘贴 URL + 附加指令**:
|
||||
- "帮我看看这个文档" → `read`
|
||||
- "这个文档的信息" → `info`
|
||||
- "往这个文档追加内容" → `update --mode append`
|
||||
- "编辑这个文档的标题" → `block update`
|
||||
|
||||
关键区分: doc(文档编辑/阅读) vs aitable(数据表格操作) vs drive(钉盘文件管理)
|
||||
|
||||
## 核心工作流
|
||||
|
||||
```bash
|
||||
# ── 工作流 1: 浏览并阅读文档 ──
|
||||
|
||||
# 1. 浏览我的文档根目录
|
||||
dws doc list --format json
|
||||
|
||||
# 2. 浏览子文件夹
|
||||
dws doc list --folder <FOLDER_ID> --format json
|
||||
|
||||
# 3. 获取文档元信息 (标题、类型、权限)
|
||||
dws doc info --node <DOC_ID> --format json
|
||||
|
||||
# 4. 读取文档内容 (Markdown 格式)
|
||||
dws doc read --node <DOC_ID> --format json
|
||||
|
||||
# ── 工作流 2: 创建文档并写入内容 ──
|
||||
|
||||
# 1. (可选) 创建文件夹 — 提取 nodeId
|
||||
dws doc folder create --name "项目资料" --format json
|
||||
|
||||
# 2. 创建文档 — 提取 nodeId
|
||||
dws doc create --name "项目周报" --folder <FOLDER_ID> --format json
|
||||
|
||||
# 3. 写入内容 (追加模式)
|
||||
dws doc update --node <DOC_ID> --markdown "# 本周总结\n\n- 完成了 A\n- 推进了 B" --mode append --format json
|
||||
|
||||
# ── 工作流 3: 一步创建带内容的文档 ──
|
||||
|
||||
dws doc create --name "会议纪要" --markdown "# 会议纪要\n\n## 议题\n\n1. ..." --format json
|
||||
|
||||
# ── 工作流 4: 上传本地文件到钉钉文档/知识库 ──
|
||||
|
||||
# 1. 上传到"我的文档"根目录
|
||||
dws doc upload --file ./report.pdf
|
||||
|
||||
# 2. 上传到指定文件夹
|
||||
dws doc upload --file ./slides.pptx --name "Q1汇报.pptx" --folder <FOLDER_ID>
|
||||
|
||||
# 3. 上传到知识库并转换为在线文档
|
||||
dws doc upload --file ./data.xlsx --workspace <WS_ID> --convert
|
||||
|
||||
# ── 工作流 5: 下载文件到本地 ──
|
||||
|
||||
# 1. 下载到当前目录 (自动推断文件名)
|
||||
dws doc download --node <NODE_ID>
|
||||
|
||||
# 2. 下载到指定路径
|
||||
dws doc download --node <NODE_ID> --output ./report.pdf
|
||||
|
||||
# 3. 下载到指定目录 (自动推断文件名)
|
||||
dws doc download --node <NODE_ID> --output ~/downloads/
|
||||
|
||||
# ── 工作流 6: 块级精细编辑 ──
|
||||
|
||||
# 1. 查看文档块结构 — 获取 blockId
|
||||
dws doc block list --node <DOC_ID> --format json
|
||||
|
||||
# 2. 在文档末尾插入段落
|
||||
dws doc block insert --node <DOC_ID> --text "新增内容"
|
||||
|
||||
# 3. 在指定块之前插入标题
|
||||
dws doc block insert --node <DOC_ID> --heading "新章节" --level 2 --ref-block <BLOCK_ID> --where before
|
||||
|
||||
# 4. 更新某个块的内容
|
||||
dws doc block update --node <DOC_ID> --block-id <BLOCK_ID> --text "修改后的内容"
|
||||
|
||||
# 5. 删除块
|
||||
dws doc block delete --node <DOC_ID> --block-id <BLOCK_ID> --yes
|
||||
|
||||
# ── 工作流 7: 文档评论管理 ──
|
||||
|
||||
# 1. 查看文档的所有评论
|
||||
dws doc comment list --node <DOC_ID> --format json
|
||||
|
||||
# 2. 在文档上创建评论
|
||||
dws doc comment create --node <DOC_ID> --content "这里需要补充数据来源" --format json
|
||||
|
||||
# 3. 创建评论并 @ 相关人
|
||||
# 先搜索用户: dws contact user search --query "张三" --format json → 提取 userId
|
||||
# 再将 userId 传入 --mention
|
||||
dws doc comment create --node <DOC_ID> --content "请确认这部分内容" --mention <userId1>,<userId2> --format json
|
||||
|
||||
# 4. 回复某条评论(commentKey 从 list 或 create 返回中获取)
|
||||
dws doc comment reply --node <DOC_ID> --comment-key <COMMENT_KEY> --content "已修改" --format json
|
||||
|
||||
# 5. 用表情回复评论
|
||||
dws doc comment reply --node <DOC_ID> --comment-key <COMMENT_KEY> --content "比心" --emoji --format json
|
||||
```
|
||||
|
||||
## 上下文传递表
|
||||
|
||||
| 操作 | 从返回中提取 | 用于 |
|
||||
|------|-------------|------|
|
||||
| `list` | `nodes[].nodeId` | read / info / update / block 操作的 --node |
|
||||
| `list` | folder 类型的 `nodeId` | list 的 --folder, create 的 --folder |
|
||||
| `search` | 文档 `nodeId` / URL / `createTime` / `creatorUid` | read / info / update 的 --node;创建时间与创建者信息 |
|
||||
| `create` | `nodeId` | update / block 操作的 --node |
|
||||
| `folder create` | `nodeId` | create / list / upload 的 --folder |
|
||||
| `block list` | `blockId` | block insert 的 --ref-block, block update/delete 的 --block-id |
|
||||
| `upload` | `nodeId` / URL | 上传后文件的访问链接 |
|
||||
| `download` | 本地文件路径 | 下载后的文件保存位置 |
|
||||
| `comment list` | `commentList[].commentKey` | comment reply 的 --comment-key |
|
||||
| `comment create` | `commentKey` | comment reply 的 --comment-key |
|
||||
| `contact user search` | `userId` | comment create/reply 的 --mention |
|
||||
|
||||
## nodeId 双格式说明
|
||||
|
||||
所有 `--node` 参数同时支持两种格式,系统自动识别:
|
||||
- **文档 ID**: 字母数字字符串,如 `9E05BDRVQePjzLkZt2p2vE7kV63zgkYA`
|
||||
- **文档 URL**: `https://alidocs.dingtalk.com/i/nodes/{dentryUuid}`,如 `https://alidocs.dingtalk.com/i/nodes/9E05BDRVQePjzLkZt2p2vE7kV63zgkYA`
|
||||
|
||||
两种方式等价,以下命令效果相同:
|
||||
```bash
|
||||
dws doc read --node 9E05BDRVQePjzLkZt2p2vE7kV63zgkYA
|
||||
dws doc read --node "https://alidocs.dingtalk.com/i/nodes/9E05BDRVQePjzLkZt2p2vE7kV63zgkYA"
|
||||
```
|
||||
|
||||
`--folder` 参数同样支持文件夹 URL 或 ID。
|
||||
|
||||
## 注意事项
|
||||
|
||||
- `update --mode overwrite` 会**清空原内容后重写**,⚠️ 谨慎使用;默认 `--mode append` (追加) 更安全
|
||||
- `read` 返回 Markdown 格式的文档内容,仅限有"下载"权限的文档
|
||||
- `create` 不传 `--folder` 和 `--workspace` 时,默认创建在"我的文档"根目录
|
||||
- `block list/insert/update/delete` 是块级精细编辑,适合结构化修改;简单内容追加建议用 `update --mode append`
|
||||
- `block insert` 优先使用 `--text` 或 `--heading` 快捷方式;复杂块类型 (table, callout 等) 使用 `--element` JSON
|
||||
- `markdown` 参数中的换行必须使用**真实换行符**(即实际的换行字符,Unicode `U+000A`),而不是字面量字符串 `\n`(反斜杠加字母 n)。在通过程序或大模型构造此参数时,请确保字符串在发送前已正确反转义。如果传入的是两个字符的字面量 `\n`,所有内容将渲染在同一行,导致标题、段落和表格格式全部错乱。
|
||||
- 块类型包括: paragraph, heading, blockquote, callout, columns, orderedList, unorderedList, table, sheet, attachment, slot
|
||||
- 关键区分: doc(文档编辑/阅读) vs aitable(数据表格操作) vs drive(钉盘文件管理)
|
||||
- `upload` 支持上传任意类型文件 (PDF、Office、图片等) 到钉钉文档空间或知识库;`--convert` 可将 Office 文件转换为钉钉在线文档
|
||||
- `upload` 是三步自动完成的流程 (获取凭证 → OSS 上传 → 提交入库),无需手动分步操作
|
||||
- `download` 是两步自动完成的流程 (获取下载链接 → HTTP GET 下载),支持自动推断文件名;`--output` 可指定文件路径或目录
|
||||
|
||||
## 自动化脚本
|
||||
|
||||
| 脚本 | 场景 | 用法 |
|
||||
|------|------|------|
|
||||
| [doc_create_and_write.py](../../scripts/doc_create_and_write.py) | 创建文档并写入 Markdown 内容 | `python doc_create_and_write.py --name "周报" --content "# 本周总结"` |
|
||||
|
||||
## 相关产品
|
||||
|
||||
- [aitable](./aitable.md) — 结构化数据表格(行列/字段/记录),不是富文本文档
|
||||
- [drive](./drive.md) — 钉盘文件存储/上传/下载,不是文档内容编辑
|
||||
- [report](./report.md) — 钉钉日志系统(日报/周报模版),不是在线文档
|
||||
@@ -0,0 +1,274 @@
|
||||
# AI听记 (minutes) 命令参考
|
||||
|
||||
## 命令总览
|
||||
|
||||
### 查询我创建的听记列表
|
||||
```
|
||||
Usage:
|
||||
dws minutes list mine [flags]
|
||||
Example:
|
||||
dws minutes list mine
|
||||
dws minutes list mine --max 10
|
||||
dws minutes list mine --max 10 --next-token <nextToken>
|
||||
dws minutes list mine --query "周会"
|
||||
Flags:
|
||||
--max float 查询的听记篇数 (默认 10)
|
||||
--next-token string 分页 token (首页留空,后续填写前次返回的 nextToken)
|
||||
--query string 关键字筛选 (可选)
|
||||
--start string 开始时间 ISO-8601 (可选)
|
||||
--end string 结束时间 ISO-8601 (可选)
|
||||
```
|
||||
|
||||
查询我创建的听记列表,支持 `--max` 和 `--next-token` 分页,支持按关键字和时间范围筛选。
|
||||
|
||||
### 查询他人共享给我的听记列表
|
||||
```
|
||||
Usage:
|
||||
dws minutes list shared [flags]
|
||||
Example:
|
||||
dws minutes list shared
|
||||
dws minutes list shared --max 20
|
||||
dws minutes list shared --max 5 --next-token <nextToken>
|
||||
Flags:
|
||||
--max float 查询的听记篇数 (默认 10)
|
||||
--next-token string 分页 token (首页留空,后续填写前次返回的 nextToken)
|
||||
--query string 关键字筛选 (可选)
|
||||
--start string 开始时间 ISO-8601 (可选)
|
||||
--end string 结束时间 ISO-8601 (可选)
|
||||
```
|
||||
|
||||
查询他人共享给我的听记列表,支持 `--max` 和 `--next-token` 分页,支持按关键字和时间范围筛选。
|
||||
|
||||
### 查询我有权限访问的所有听记列表
|
||||
```
|
||||
Usage:
|
||||
dws minutes list all [flags]
|
||||
Example:
|
||||
dws minutes list all
|
||||
dws minutes list all --max 20
|
||||
dws minutes list all --query "周会" --max 20
|
||||
dws minutes list all --start "2026-03-01T00:00:00+08:00" --end "2026-03-20T23:59:59+08:00"
|
||||
dws minutes list all --max 10 --next-token <nextToken>
|
||||
Flags:
|
||||
--end string 结束时间 ISO-8601 (可选)
|
||||
--query string 关键字筛选 (可选)
|
||||
--max float 查询的听记篇数 (默认 10)
|
||||
--next-token string 分页 token (首页留空,后续填写前次返回的 nextToken)
|
||||
--start string 开始时间 ISO-8601 (可选)
|
||||
```
|
||||
|
||||
查询我有权限访问的所有听记列表(包括我创建的、他人共享给我的等所有有权限的听记)。支持按关键字和时间范围筛选。时间范围和关键字为可选参数,不传则返回所有有权限的听记。支持使用 `--max` 和 `--next-token` 进行分页查询。
|
||||
|
||||
### 获取听记基础信息
|
||||
```
|
||||
Usage:
|
||||
dws minutes get info [flags]
|
||||
Example:
|
||||
dws minutes get info --id <taskUuid>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
```
|
||||
|
||||
返回字段: 创建人、开始时间、截止时间、听记标题、听记访问链接URL
|
||||
|
||||
### 获取听记 AI 摘要
|
||||
```
|
||||
Usage:
|
||||
dws minutes get summary [flags]
|
||||
Example:
|
||||
dws minutes get summary --id <taskUuid>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
```
|
||||
|
||||
返回 Markdown 格式摘要,涵盖会议主题、核心结论、关键讨论点等
|
||||
|
||||
### 获取听记关键字列表
|
||||
```
|
||||
Usage:
|
||||
dws minutes get keywords [flags]
|
||||
Example:
|
||||
dws minutes get keywords --id <taskUuid>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
```
|
||||
|
||||
### 获取听记语音转写原文
|
||||
```
|
||||
Usage:
|
||||
dws minutes get transcription [flags]
|
||||
Example:
|
||||
dws minutes get transcription --id <taskUuid>
|
||||
dws minutes get transcription --id <taskUuid> --direction 1
|
||||
Flags:
|
||||
--direction string 排序方向: 0=正序, 1=倒序 (默认 0)
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
--next-token string 下一页的token 首次查询可空 后续查询需填写前次请求返回的nextToken
|
||||
```
|
||||
|
||||
每条记录包含: 发言人信息、转写文本、对应时间戳
|
||||
|
||||
### 获取听记中提取的待办事项
|
||||
```
|
||||
Usage:
|
||||
dws minutes get todos [flags]
|
||||
Example:
|
||||
dws minutes get todos --id <taskUuid>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
```
|
||||
|
||||
每条记录包含: 待办内容、待办唯一ID、参与人信息、待办时间
|
||||
|
||||
### 批量查询听记详情
|
||||
```
|
||||
Usage:
|
||||
dws minutes get batch [flags]
|
||||
Example:
|
||||
dws minutes get batch --ids uuid1,uuid2,uuid3
|
||||
Flags:
|
||||
--ids string 听记 taskUuid 列表,逗号分隔 (必填)
|
||||
```
|
||||
|
||||
返回字段: 听记标题、时长、参与人列表、创建时间、taskUuid、听记状态
|
||||
|
||||
### 修改听记标题
|
||||
```
|
||||
Usage:
|
||||
dws minutes update title [flags]
|
||||
Example:
|
||||
dws minutes update title --id <taskUuid> --title "Q2 复盘会议"
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填),取值逻辑参考 ## 注意事项
|
||||
--title string 新标题 (必填)
|
||||
```
|
||||
|
||||
### 发起听记(开始录音)
|
||||
```
|
||||
Usage:
|
||||
dws minutes record start [flags]
|
||||
Example:
|
||||
dws minutes record start
|
||||
dws minutes record start --session-id <sessionId>
|
||||
Flags:
|
||||
--session-id string AI 助理会话 ID (可选)
|
||||
```
|
||||
|
||||
### 暂停听记录音
|
||||
```
|
||||
Usage:
|
||||
dws minutes record pause [flags]
|
||||
Example:
|
||||
dws minutes record pause --id <taskUuid>
|
||||
dws minutes record pause --id <taskUuid> --session-id <sessionId>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填)
|
||||
--session-id string AI 助理会话 ID (可选)
|
||||
```
|
||||
|
||||
### 恢复听记录音
|
||||
```
|
||||
Usage:
|
||||
dws minutes record resume [flags]
|
||||
Example:
|
||||
dws minutes record resume --id <taskUuid>
|
||||
dws minutes record resume --id <taskUuid> --session-id <sessionId>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填)
|
||||
--session-id string AI 助理会话 ID (可选)
|
||||
```
|
||||
|
||||
### 结束听记录音
|
||||
```
|
||||
Usage:
|
||||
dws minutes record stop [flags]
|
||||
Example:
|
||||
dws minutes record stop --id <taskUuid>
|
||||
dws minutes record stop --id <taskUuid> --session-id <sessionId>
|
||||
Flags:
|
||||
--id string 听记 taskUuid (必填)
|
||||
--session-id string AI 助理会话 ID (可选)
|
||||
```
|
||||
|
||||
## 意图判断
|
||||
|
||||
用户说"我的听记/我创建的听记" → `list mine`(可附加 `--query`、`--start`、`--end` 筛选)
|
||||
用户说"别人给我的听记/共享听记" → `list shared`(可附加 `--query`、`--start`、`--end` 筛选)
|
||||
用户说"有权限的听记/我能访问的听记/所有听记" → `list all`(可附加 `--query`、`--start`、`--end` 筛选)
|
||||
用户说"某时间段内的听记/按时间查听记/按关键词查听记" → 根据所属范围选择 `list mine`/`list shared`/`list all`,附加 `--start`、`--end`、`--query` 参数
|
||||
用户说"听记详情/听记信息" → `get info`
|
||||
用户说"摘要/总结/会议纪要" → `get summary`
|
||||
用户说"关键字/关键词" → `get keywords`
|
||||
用户说"原文/转写/录音文字" → `get transcription`
|
||||
用户说"会议待办/听记待办" → `get todos`
|
||||
用户说"改听记标题/重命名听记" → `update title`
|
||||
用户说"发起听记/开始录音" → `record start`
|
||||
用户说"暂停听记/暂停录音" → `record pause`
|
||||
用户说"继续听记/恢复录音" → `record resume`
|
||||
用户说"结束听记/结束录音" → `record stop`
|
||||
用户传入听记 URL(如 `https://shanji.dingtalk.com/app/transcribes/xxx`),从 URL 提取 taskUuid,再执行对应的 get/update 操作
|
||||
|
||||
## 核心工作流
|
||||
|
||||
```bash
|
||||
# 0. 发起听记(开始录音)
|
||||
dws minutes record start --format json
|
||||
|
||||
# 1. 查看我的听记列表 — 提取 taskUuid
|
||||
dws minutes list mine --format json
|
||||
dws minutes list mine --max 10 --next-token <nextToken> --format json
|
||||
dws minutes list mine --query "周会" --format json
|
||||
|
||||
# 1b. 查看共享给我的听记
|
||||
dws minutes list shared --max 20 --format json
|
||||
dws minutes list shared --query "日报" --format json
|
||||
|
||||
# 1c. 查看我有权限访问的所有听记(支持关键字和时间范围筛选)
|
||||
dws minutes list all --format json
|
||||
dws minutes list all --query "周会" --start "2026-03-01T00:00:00+08:00" --end "2026-03-20T23:59:59+08:00" --format json
|
||||
|
||||
# 2. 获取 AI 摘要
|
||||
dws minutes get summary --id <taskUuid> --format json
|
||||
|
||||
# 3. 查看完整转写原文
|
||||
dws minutes get transcription --id <taskUuid> --format json
|
||||
|
||||
# 4. 提取待办事项
|
||||
dws minutes get todos --id <taskUuid> --format json
|
||||
|
||||
# 5. 修改标题
|
||||
dws minutes update title --id <taskUuid> --title "新标题" --format json
|
||||
|
||||
# 6. 录音控制(基于 start 返回的 taskUuid)
|
||||
dws minutes record pause --id <taskUuid> --format json
|
||||
dws minutes record resume --id <taskUuid> --format json
|
||||
dws minutes record stop --id <taskUuid> --format json
|
||||
```
|
||||
|
||||
## 上下文传递表
|
||||
|
||||
| 操作 | 从返回中提取 | 用于 |
|
||||
|------|-------------|------|
|
||||
| `list mine` | `taskUuid`、`nextToken` | get/update 的 --id;翻页时 --next-token |
|
||||
| `list shared` | `taskUuid`、`nextToken` | get/update 的 --id;翻页时 --next-token |
|
||||
| `list all` | `taskUuid`、`nextToken` | get/update 的 --id;翻页时 --next-token |
|
||||
| `get batch` | 各听记 `taskUuid` | 进一步查询详情 |
|
||||
|
||||
## 注意事项
|
||||
- `taskUuid` 是听记的唯一标识,所有 get/update 操作均以此为入参
|
||||
- `record start` 对应 MCP 工具 `execute_listening_note_command` 的 `cmd=create`,通常会返回可继续控制录音的 `taskUuid/uuid`
|
||||
- `record pause` / `record resume` / `record stop` 对应 `cmd=pause/resume/end`,需要传入 `--id`(映射 MCP 入参 `uuid`)
|
||||
- 如果用户传入听记 URL(格式: `https://shanji.dingtalk.com/app/transcribes/<taskUuid>`),直接从路径末段提取 taskUuid 作为 `--id` 参数,无需再调用 list 查询
|
||||
- `list mine`、`list shared`、`list all` 统一走 `list_by_keyword_and_time_range` 链路,通过 `belongingConditionId` 区分(`created` / `shared` / `noLimit`)
|
||||
- 三个 list 命令均支持 `--max`、`--next-token` 分页及 `--query`、`--start`、`--end` 筛选
|
||||
- `list mine`、`list shared` 默认每页 20 条,`list all` 默认每页 10 条
|
||||
- `get summary` 返回 AI 生成的结构化 Markdown 摘要
|
||||
- `get transcription` 的 `--direction` 控制时间排序: 0=正序(默认), 1=倒序
|
||||
- `get batch` 支持一次查询多个听记,用逗号分隔 taskUuid
|
||||
|
||||
## 自动化脚本
|
||||
|
||||
| 脚本 | 场景 | 用法 |
|
||||
|------|------|------|
|
||||
| [minutes_recent_summary.py](../../scripts/minutes_recent_summary.py) | 获取最近听记的 AI 摘要并合并 | `python minutes_recent_summary.py --max 5` |
|
||||
| [minutes_extract_todos.py](../../scripts/minutes_extract_todos.py) | 从听记中提取待办事项汇总 | `python minutes_extract_todos.py --max 5` |
|
||||
+10
-5
@@ -4,8 +4,9 @@ DWS Skill Test Runner
|
||||
Validates AI Agent's ability to translate natural language prompts into DWS CLI commands.
|
||||
"""
|
||||
|
||||
import re
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from collections import defaultdict
|
||||
@@ -404,8 +405,12 @@ def generate_report(results: list[dict]) -> str:
|
||||
return '\n'.join(report)
|
||||
|
||||
def main():
|
||||
# Read the test file
|
||||
with open('/Users/tianlei.qjb/Documents/my_python_project/cli/test/skill_tests.md', 'r', encoding='utf-8') as f:
|
||||
test_dir = Path(__file__).resolve().parent
|
||||
test_cases_path = test_dir / 'skill_tests.md'
|
||||
results_path = test_dir / 'skill_tests_results.md'
|
||||
|
||||
# Read the test file from the repo instead of a developer-local absolute path.
|
||||
with open(test_cases_path, 'r', encoding='utf-8') as f:
|
||||
content = f.read()
|
||||
|
||||
# Parse test cases
|
||||
@@ -421,8 +426,8 @@ def main():
|
||||
# Generate report
|
||||
report = generate_report(results)
|
||||
|
||||
# Write results
|
||||
with open('/Users/tianlei.qjb/Documents/my_python_project/cli/test/skill_tests_results.md', 'w', encoding='utf-8') as f:
|
||||
# Write results next to the test cases so the script stays portable.
|
||||
with open(results_path, 'w', encoding='utf-8') as f:
|
||||
f.write(report)
|
||||
|
||||
print(f"\nResults written to skill_tests_results.md")
|
||||
|
||||
@@ -1,34 +0,0 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
type ContentBlock struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text,omitempty"`
|
||||
}
|
||||
|
||||
func main() {
|
||||
data := []byte(`{"content":[{"type":"text","text":"{\"summary\":\"...\",\"data\":{\"tableId\":\"abc\"},\"status\":\"success\"}"}],"structuredContent":{"summary":"...","data":{"tableId":"abc"},"status":"success"},"isError":false}`)
|
||||
|
||||
type rawResult struct {
|
||||
Content json.RawMessage `json:"content"`
|
||||
StructuredContent map[string]any `json:"structuredContent"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
}
|
||||
|
||||
var raw rawResult
|
||||
_ = json.Unmarshal(data, &raw)
|
||||
|
||||
fmt.Printf("raw.Content string: %s\n", string(raw.Content))
|
||||
|
||||
var object map[string]any
|
||||
errMap := json.Unmarshal(raw.Content, &object)
|
||||
fmt.Printf("errMap: %v\n", errMap)
|
||||
|
||||
var blocks []ContentBlock
|
||||
errBlocks := json.Unmarshal(raw.Content, &blocks)
|
||||
fmt.Printf("errBlocks: %v, len(blocks): %d\n", errBlocks, len(blocks))
|
||||
}
|
||||
Reference in New Issue
Block a user