Compare commits

...
91 Commits
Author SHA1 Message Date
ugurtafrali 1ea8bad074 Reduce MCP discovery timeouts
When MCP plugins are unreachable, the discovery times out too slowly and blocks commands. Reduced the overall discovery timeout from 10s to 4s and per-server timeout from 5s to 2s. Made the per-server timeout configurable.
2026-04-17 14:40:41 +08:00
修雨 7b7aeadbbe Merge pull request #120 from FuShu-Yang/bugfix/plugin-system-fushu
feat(plugin): stdio subprocess identity injection and lifecycle management
2026-04-17 14:15:44 +08:00
扶疏 94ad422a9f style: fix code formatting 2026-04-17 12:19:27 +08:00
扶疏 93318f4a83 fix(plugin): stop stdio child processes on exit and before removal 2026-04-17 10:03:01 +08:00
扶疏 a14fd0250c feat(plugin): inject user identity (UserID, CorpID) into stdio plugin subprocesses 2026-04-17 10:02:26 +08:00
fantiu 4bf300d862 Merge pull request #117 from wxianfeng/main
feat(skill): add skill install command with new download API and multi-agent target support
2026-04-16 14:58:35 +08:00
github-actions[bot] 1a1fc531f5 chore: update coverage badge [skip ci] 2026-04-16 06:56:02 +00:00
wxianfeng 9fc570607f Merge remote-tracking branch 'refs/remotes/origin/main' 2026-04-16 14:53:43 +08:00
wxianfeng 4851d19141 Merge remote-tracking branch 'upstream/main' 2026-04-16 14:53:07 +08:00
xianfeng wang 42fb25d150 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-16 14:45:48 +08:00
wxianfeng 416ad6571d skill opt 2026-04-16 14:45:15 +08:00
修雨 9a119fbd64 Merge pull request #115 from PeterGuy326/main
feat(plugin): add plugin system core framework with lifecycle management
2026-04-16 14:37:14 +08:00
wxianfeng c5decb2f90 error opt 2026-04-16 12:40:57 +08:00
github-actions[bot] fae2a4f5f0 chore: update coverage badge [skip ci] 2026-04-16 02:55:54 +00:00
修雨 d25b106e4f fix(plugin): harden plugin system security boundaries
- Validate git URL protocol: reject file:// and local paths, only allow https/ssh
- Reject symlink entries in ZIP extraction to prevent path traversal
- Validate build.output must be relative path within plugin directory
- Reject absolute paths in stdio server command declarations
- Block dangerous env var names (PATH, LD_PRELOAD, etc.) from plugin config injection
- Remove conference from default managed plugins (source not yet available)
2026-04-16 10:53:53 +08:00
github-actions[bot] 9f78e51ae7 chore: update coverage badge [skip ci] 2026-04-16 02:00:01 +00:00
qiweijie.qwj d2752d8b5b feat(plugin): improve CLI overlay resolution and plugin install robustness
- Support file path reference in plugin.json cli field (e.g. "cli": "overlay.json")
  in addition to inline JSON objects, resolving path relative to plugin root
- Add description field to CLIToolOverride for static command descriptions
  as fallback when MCP tools/list is unavailable (e.g. upstream server offline)
- Fix plugin install on Windows: use cmd /C instead of sh -c for build commands
- Skip copying identical files during plugin install to avoid overwriting
  locked executables (running stdio plugin processes on Windows)
- Add symlink skip and path traversal guard in copyDir for security
- Clean up stale files in destination during plugin upgrade via removeStaleFiles

🤖 Generated with [Qoder][https://qoder.com]
2026-04-16 09:38:18 +08:00
修雨 8ecbff391c feat(plugin): bootstrap default managed plugins on first run
- Add `EnsureManaged` to the plugin updater to automatically install missing default plugins on startup.
- Translate the managed plugin removal error message to English for better consistency.
- Update tests to match the new error message.
2026-04-16 09:38:18 +08:00
修雨 d259864a2b fix(plugin): skip min version check in dev mode
Ignore `MinCLIVersion` validation when the current CLI version is "dev". This allows plugins to be loaded during local development and testing without being blocked by strict semantic version constraints.
2026-04-16 09:38:18 +08:00
修雨 408098bdc1 feat(plugin): add command name conflict protection
Prevent plugins from hijacking built-in commands (auth, plugin, cache,
etc.) and detect duplicate names between plugins.

- Add reservedCommands set for protected built-in command names
- Add addPluginCommandsSafe with three conflict rules:
  - Plugin vs reserved → reject with warning
  - Plugin vs plugin → first wins, second rejected with warning
  - Plugin vs Market → plugin wins (intentional override)
- Fix plugin create scaffold format string arg count
2026-04-16 09:38:18 +08:00
修雨 658ec1676c feat(plugin): add build lifecycle for binary distribution
Add `build` field to plugin.json so stdio servers can be automatically
compiled to native binaries during install, eliminating runtime deps.

- Add BuildConfig struct to Manifest (command + output fields)
- Add BuildPlugin/runBuild in loader with output verification
- Auto-trigger build in InstallFromDir and InstallFromGit (rollback on failure)
- Add `dws plugin build <dir>` command for manual builds
- Update `plugin create` scaffold to include build template
2026-04-16 09:38:18 +08:00
ybc e36d3b3474 style: fix gofmt formatting for 4 files 2026-04-16 09:38:18 +08:00
ybc 0b9952c58d feat(plugin): add persistent plugin config via settings.json
Implement 'dws plugin config set/get/list/unset' commands that persist
plugin configuration (e.g. API keys) to ~/.dws/settings.json. Values
are automatically injected as environment variables at plugin load time,
so ${KEY} references in plugin.json headers/endpoints resolve without
manual 'export' each session.

Key changes:
- loader.go: add Get/Set/Unset/List/InjectPluginConfigEnv methods
- plugin_cmd.go: add 'plugin config' subcommand group (set/get/list/unset)
- root.go: call InjectPluginConfigEnv() before plugin loading
- loader_config_test.go: 9 unit tests covering all new functionality

User env vars take precedence over settings.json values.
Sensitive values are masked in 'plugin config list' output.
2026-04-16 09:38:18 +08:00
ybc 56af1ea091 feat: parallel service discovery + unit tests for auth-headers
- Optimize loadPlugins to discover HTTP servers in parallel (sync.WaitGroup)
  when multiple remote servers exist, reducing startup time from sequential
  N*10s to parallel max(10s)
- Add auth_registry_test.go with 7 unit tests covering:
  - PluginAuth registry CRUD and multi-product isolation
  - deriveToolCLIName conversion (6 cases)
  - registerPluginAuthFromHeaders extraction and registration
  - buildPluginAuthClient with/without Authorization header
- Add 3 test cases to plugin_test.go covering:
  - Headers -> AuthHeaders conversion with env var expansion
  - No-headers descriptor produces nil AuthHeaders
  - ParseManifest correctly parses headers field from plugin.json
2026-04-16 09:38:18 +08:00
ybc ea5859b92b feat(plugin): support custom auth headers for third-party MCP servers
Add Auth Token Registry pattern (inspired by stdio_registry.go) to allow
plugins to declare per-server HTTP headers in plugin.json. This enables
third-party streamable-http MCP servers (e.g. Alibaba Cloud Bailian) to
use their own API keys independently from the default DingTalk OAuth token.

Changes:
- New: internal/app/auth_registry.go — per-productID auth credential store
- manifest.go: MCPServer gains Headers field for custom HTTP headers
- registry.go: ServerDescriptor gains AuthHeaders field
- converter.go: resolve and pass headers through ToServerDescriptors
- root.go: inject plugin auth at discovery time + auto-generate ToolOverrides
- runner.go: route plugin-auth servers to their own Bearer token at runtime

Multi-token isolation: each server is keyed by its CLI.ID (productID),
so different servers can use independent tokens without interfering with
each other or with the default DingTalk OAuth token.
2026-04-16 09:38:18 +08:00
修雨 19f2ed5c69 fix(plugin): schema flag params, HTTP tool discovery, and integration tests
- Fix schema-derived flags not being collected into tool call params
- Add registerHTTPServer for plugin streamable-http tool discovery
- Include dev plugins in ListInstalled output
- Isolate settings path for test environments
- Add stdio and HTTP MCP end-to-end integration tests
2026-04-16 09:34:52 +08:00
修雨 efbaf7a49d feat(plugin): add create scaffold, dev mode, and skill sync
- dws plugin create: generates plugin template with plugin.json, SKILL.md, hooks.json
- dws plugin dev: registers source directory as dev plugin without copying
- SyncSkills: copies plugin skills to agent directories on startup
2026-04-16 09:34:52 +08:00
修雨 374a9e9b13 feat(plugin): register stdio server tools as CLI commands
Stdio MCP servers discovered by the plugin loader now have their tools
automatically registered as CLI subcommands. The runner dispatches
stdio:// virtual endpoints to the local StdioClient subprocess instead
of the HTTP transport, enabling plugin tools to be invoked directly
from the command line (e.g. `dws hello greet --name Peter`).

Also fixes a bug where plugin endpoints registered via AppendDynamicServer
were overwritten by the subsequent SetDynamicServers call in
loadDynamicCommands, and a context lifecycle bug where stdio subprocesses
were killed immediately after startup due to a short-lived timeout context.
2026-04-16 09:34:46 +08:00
修雨 d7d85c9e67 feat(plugin): integrate updater and stdio server into startup
- Wire updater into loadPlugins to check managed plugin updates on
  CLI startup with 10s timeout and best-effort semantics
- Add StdioClients() to Plugin for creating stdio transport clients
  with DWS_PLUGIN_ROOT/DWS_PLUGIN_DATA variable expansion
- Start and initialize stdio MCP server subprocesses during plugin load
- Reorder loadPlugins: update check → load → http inject → stdio start → hooks
2026-04-16 09:33:20 +08:00
修雨 0fa982fe91 feat(plugin): add updater, stdio transport, git install, and English output
- Add updater.go for managed plugin auto-update (version check, download,
  Y/n prompt, zip extraction with zip-slip protection)
- Add stdio.go transport for local MCP server subprocess communication
  via stdin/stdout JSON-RPC 2.0
- Implement InstallFromGit in loader.go with URL parsing for HTTPS/SSH
- Wire up --git flag in plugin install command
- Switch all plugin CLI output to English (headers, status, messages)
- Add tests for parseGitURL and promptUpdate
2026-04-16 09:33:20 +08:00
修雨 c4fb1bbd3e feat(plugin): add plugin system core framework
Introduce the plugin system for DWS CLI, enabling official and third-party
plugin management. This includes plugin manifest parsing/validation, loader
with managed/user directory-based identity, MCP server conversion and
injection into the dynamic routing registry, pipeline hook adapter for
shell-based hooks, and CLI commands (list/install/info/enable/disable/
remove/validate).
2026-04-16 09:33:20 +08:00
修雨 26d7d8f946 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-16 09:31:15 +08:00
fantiu 05ac342c4b Merge pull request #109 from wxianfeng/main
feat: add skill marketplace commands, configmeta registry, and auth JSON output
2026-04-15 12:02:54 +08:00
github-actions[bot] 5e491aef8f chore: update coverage badge [skip ci] 2026-04-15 04:01:30 +00:00
wxianfeng 202187d5e2 fix conflict 2026-04-15 11:59:23 +08:00
wxianfeng 13877b1c3a merge upstream 2026-04-15 11:53:12 +08:00
wxianfeng 0e72e89ba3 auth 2026-04-15 11:45:58 +08:00
github-actions[bot] f1b68271cc chore: update coverage badge [skip ci] 2026-04-13 09:37:28 +00:00
fantiu 83efff21cd Merge pull request #102 from fantiu/feat-pipeline
feat: add unified config management, doctor diagnostics, and structured perf reports
2026-04-12 16:49:21 +08:00
fantiu e6a4b35921 feat: upgrade perf tracing from debug tool to structured diagnostics 2026-04-12 15:14:57 +08:00
fantiu cc2d97ddba feat: add dws doctor command for one-stop diagnostics 2026-04-12 14:43:55 +08:00
fantiu b78dd19cf9 feat: centralize scattered configs and add dws config list 2026-04-12 14:23:11 +08:00
fantiu 1f0a75f836 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-pipeline
refactor: extend command execution pipeline from 3 to 5 stages
2026-04-12 13:49:38 +08:00
fantiu 16202c83a3 refactor: extend command execution pipeline from 3 to 5 stages 2026-04-12 13:48:39 +08:00
github-actions[bot] f4cc76c77d chore: update coverage badge [skip ci] 2026-04-12 04:16:59 +00:00
wxianfeng 9fef6a9c43 pat test 2026-04-12 12:14:57 +08:00
fantiu 810985b03a Merge pull request #101 from wxianfeng/main
feat(errors): add ExitCoder/RawStderrError interfaces and ClassifyToolResult hook
2026-04-12 11:29:14 +08:00
xianfeng wang 02633c6bd3 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-12 11:26:29 +08:00
wxianfeng eb9416aa16 test 2026-04-12 11:26:04 +08:00
wxianfeng 65b64af213 feat(errors): add ExitCoder/RawStderrError interfaces and ClassifyToolResult hook
- Add ExitCoder interface for edition-specific error types to provide custom exit codes
- Add RawStderrError interface for errors that bypass CLI formatting and output raw content to stderr
- Update ExitCode() to resolve exit codes via ExitCoder interface before falling back to default
- Add ClassifyToolResult hook in edition.Hooks for custom MCP tool result error classification
- Invoke ClassifyToolResult in runner before default business-error detection
- Handle RawStderrError in printExecutionError to pass raw JSON through to desktop runtime
- Add unit tests for ExitCoder and RawStderrError interfaces
2026-04-12 11:25:32 +08:00
fantiu f1d160a481 Merge pull request #100 from wxianfeng/main
feat(auth): delegate token persistence to edition hooks with keychain fallback
2026-04-12 10:19:48 +08:00
xianfeng wang f8c7f012a1 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-12 10:12:16 +08:00
github-actions[bot] 45618a55e6 chore: update coverage badge [skip ci] 2026-04-10 15:36:12 +00:00
wxianfeng c49583836b feat(auth): delegate token persistence to edition hooks with keychain fallback
Add SaveToken, LoadToken, and DeleteToken hook fields to edition.Hooks so overlay builds can provide custom token storage. Refactor SaveTokenData, LoadTokenData, and DeleteTokenData to delegate to hooks when present, falling back to default keychain-based storage.
2026-04-10 23:33:11 +08:00
fantiu 9dc8dc7065 Merge pull request #99 from wxianfeng/main
feat(auth): add edition overlay auth credentials and token marker for embedded mode
2026-04-10 17:11:49 +08:00
xianfeng wang f978e306cc Merge branch 'DingTalk-Real-AI:main' into main 2026-04-10 17:03:22 +08:00
wxianfeng aec852f971 feat(auth): add edition overlay auth credentials and token marker for embedded mode
- Add AuthClientID and AuthClientFromMCP fields to edition Hooks, allowing overlays to override the default OAuth client ID and route auth through MCP endpoints.

- Introduce token marker file (token.json) mechanism so the host application in embedded mode can detect authentication state without accessing the keychain.

- Update SaveTokenData/DeleteTokenData to write/remove the marker file when running in embedded mode.

- Update ClientID() and IsClientIDFromMCP() to respect edition overrides.
2026-04-10 17:02:18 +08:00
fantiu 143f781064 Merge pull request #97 from wxianfeng/main
feat(skill): restore legacy skill find/get and improve root help
2026-04-10 10:54:48 +08:00
github-actions[bot] 953b422295 chore: update coverage badge [skip ci] 2026-04-10 02:49:49 +00:00
xianfeng wang da1a0f1299 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-10 10:47:56 +08:00
wxianfeng bc7d19cfd8 feat(skill): restore find/get for legacy skill market API
- Add skill find (keyword search) and skill get (download to temp dir)
  using mcp.dingtalk.com endpoints; keep skill add (aihub download)
- Add hidden skill search hint for old usage
- List utility commands in root help; keep skill visible (not hidden)
- Extend tests for new subcommands and root help

Made-with: Cursor
2026-04-10 10:47:14 +08:00
github-actions[bot] df3122090f chore: update coverage badge [skip ci] 2026-04-09 09:18:24 +00:00
wxianfeng 713fdf6188 feat(schema): return structured degraded errors instead of silent empty catalog
EnvironmentLoader.Load() previously returned (empty Catalog, nil) on all
failure paths, making it impossible for callers to distinguish "no services
available" from "discovery failed due to auth/network issues".

Changes:
- Add CatalogDegraded error type with three reasons: unauthenticated,
  market_unreachable, runtime_all_failed
- Add auth pre-check: return DegradedUnauthenticated immediately when no
  token is available, avoiding doomed MCP connections
- Schema command now handles CatalogDegraded gracefully: outputs hint to
  stderr and includes degraded/reason/hint fields in JSON output
- Runner preserves graceful degradation by ignoring CatalogDegraded errors
- Skip empty-products cache to prevent stale cache from masking auth errors
- Edition-aware hint text (open-source vs embedded/wukong)

Made-with: Cursor
2026-04-09 17:16:04 +08:00
fantiu 70e21b58b4 Merge pull request #96 from wxianfeng/main
feat(auth): unify token resolution and add edition lifecycle hooks for overlay transports
2026-04-09 11:16:10 +08:00
github-actions[bot] 18ebba1bb2 chore: update coverage badge [skip ci] 2026-04-09 03:04:03 +00:00
wxianfeng 937404e6df Fix test 2026-04-09 11:01:59 +08:00
wxianfeng 88e155dd23 Fix test 2026-04-09 10:49:17 +08:00
xianfeng wang 2e2cea0973 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-08 23:49:40 +08:00
wxianfeng 9b8c13a8b6 refactor(auth): unify auxiliary token resolution with MCP cached path
Extract resolveAccessTokenFromDir from getCachedRuntimeToken so both
MCP and non-MCP clients (e.g. A2A gateway) share the same OAuth +
legacy fallback logic including host compatibility hooks.

- Add ResolveAuxiliaryAccessToken in internal/app for overlay use;
  reuses process-level token cache when configDir matches the edition
  default, avoiding repeated Keychain access.
- Simplify pkg/runtimetoken to a thin delegate to the shared impl.
- Refactor getCachedRuntimeToken to call resolveAccessTokenFromDir,
  removing duplicated provider/manager setup code.

Made-with: Cursor
2026-04-08 20:15:30 +08:00
wxianfeng 8238cc9f41 feat(edition): add AfterPersistentPreRun hook and runtime token helpers
- Invoke optional AfterPersistentPreRun after root PersistentPreRunE setup
  so overlays can wire non-MCP clients (e.g. A2A gateway).
- Add pkg/runtimetoken.ResolveAccessToken mirroring MCP auth resolution.
- Export MCP identity headers via pkg/cli for auxiliary HTTP transports.
- Fix staticcheck: SA1012 nil-context test guard, ST1005 error string in verify.

Made-with: Cursor
2026-04-08 20:03:56 +08:00
coffeeBigSir e59c4f30b8 Merge pull request #95 from DingTalk-Real-AI/coffeeBigSir-patch-1
Update SKILL.md
2026-04-08 17:44:38 +08:00
coffeeBigSir fd7ef5edc2 Update SKILL.md 2026-04-08 17:21:23 +08:00
xianfeng wang a8e1acec09 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-08 10:41:33 +08:00
coffeeBigSir 31eb10985e Merge pull request #82 from wqyenjoy/main
add command
2026-04-07 20:31:49 +08:00
coffeeBigSir 1436b62a80 Merge pull request #86 from DingTalk-Real-AI/install-yh
feat(install): align skill dirs with npm and add OpenClaw
2026-04-07 20:25:31 +08:00
tianlei.qjb ec6a27635b feat(install): align skill dirs with npm and add OpenClaw 2026-04-07 16:02:05 +08:00
玉澜 1727744691 add command 2026-04-03 20:01:05 +08:00
fantiu afdd47b5a5 Merge pull request #81 from fantiu/feat-cmd-performance
perf: optimize command timeout handling, instrumentation, and diagnostics
2026-04-03 19:52:42 +08:00
fantiu d968e8e551 perf: optimize command timeout handling, instrumentation, and diagnostics 2026-04-03 18:03:28 +08:00
fantiu c649d1a762 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-cmd-performance
perf: optimize command timeout handling, instrumentation, and diagnostics
2026-04-03 18:02:33 +08:00
fantiu a1f5d97345 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-03 18:01:52 +08:00
meng93 58062515a5 Merge pull request #78 from DingTalk-Real-AI/feat/issue-ai-table-label-change
feat: 优化标签在多维表中的展示
2026-04-03 11:14:53 +08:00
meng93 5614b508f2 feat: to #73551688 优化标签在多维表中的展示 2026-04-03 10:41:49 +08:00
github-actions[bot] 5e003a41b1 chore: update coverage badge [skip ci] 2026-04-03 02:31:09 +00:00
wxianfeng 4eaeb1dd4a fix conflict 2026-04-03 10:29:05 +08:00
github-actions[bot] 84471bd6f0 chore: update coverage badge [skip ci] 2026-04-03 02:03:16 +00:00
fantiu c8e3ac21c2 Merge pull request #76 from DingTalk-Real-AI/npm
docs: add npm install method to README
2026-04-03 09:17:55 +08:00
tianlei.qjb c38892b7cf docs: add npm install method to README 2026-04-02 22:34:42 +08:00
fantiu 1a0a5324f0 docs: note upgrade command requires v1.0.7+ 2026-04-02 20:06:55 +08:00
github-actions[bot] c1e9e9e0d6 chore: update coverage badge [skip ci] 2026-04-02 09:03:29 +00:00
github-actions[bot] cc4dd1e87b chore: update coverage badge [skip ci] 2026-03-31 07:54:38 +00:00
93 changed files with 11080 additions and 311 deletions
+1 -1
View File
@@ -1 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 52.6%"><title>coverage: 52.6%</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.6%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">52.6%</text></g></svg>
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 48.7%"><title>coverage: 48.7%</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">48.7%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">48.7%</text></g></svg>

Before

Width:  |  Height:  |  Size: 1.1 KiB

After

Width:  |  Height:  |  Size: 1.1 KiB

+1 -1
View File
@@ -34,7 +34,7 @@ jobs:
body: issue.body,
state: issue.state,
html_url: issue.html_url,
labels: (issue.labels || []).map(label => label.name)
labels: (issue.labels || []).map(label => label.name).join(', ') || '无标签'
}
};
+8
View File
@@ -66,6 +66,12 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
<details>
<summary>Other install methods</summary>
**npm** (requires Node.js (npm/npx)):
```bash
npm install -g dingtalk-workspace-cli
```
**Pre-built binary**: download from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases).
> **macOS users**: If you see "cannot be opened because Apple cannot check it for malicious software", run:
@@ -88,6 +94,8 @@ cp dws ~/.local/bin/ # install to PATH
## Upgrade
> Requires **v1.0.7** or later. For earlier versions, please re-run the [install script](#installation) to upgrade.
dws has built-in self-upgrade capability. Updates are pulled directly from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) with SHA256 integrity verification and automatic backup.
```bash
+8
View File
@@ -66,6 +66,12 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
<details>
<summary>其他安装方式</summary>
**npm**(需要 Node.js(npm/npx)):
```bash
npm install -g dingtalk-workspace-cli
```
**预编译二进制文件**:从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 下载。
> **macOS 用户注意**:如果提示“无法打开,因为 Apple 无法检查其是否包含恶意软件”,请执行:
@@ -88,6 +94,8 @@ cp dws ~/.local/bin/ # 安装到 PATH
## 升级
> 需要 **v1.0.7** 及以上版本。更早版本请重新执行[安装脚本](#安装)进行升级。
dws 内置自升级能力,直接从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 拉取更新,支持 SHA256 完整性校验和自动备份。
```bash
+1
View File
@@ -52,6 +52,7 @@ __KEG_ONLY_LINE__
Pathname.new(File.join(Dir.home, ".amp/skills/dws")),
Pathname.new(File.join(Dir.home, ".kiro/skills/dws")),
Pathname.new(File.join(Dir.home, ".trae/skills/dws")),
Pathname.new(File.join(Dir.home, ".openclaw/skills/dws")),
]
targets.each_with_index do |dest, index|
+2
View File
@@ -7,6 +7,7 @@ const os = require("os");
const path = require("path");
const childProcess = require("child_process");
// Canonical list: keep scripts/install.sh, scripts/install.ps1, scripts/install-skills.sh in sync.
const AGENT_DIRS = [
".agents/skills",
".claude/skills",
@@ -20,6 +21,7 @@ const AGENT_DIRS = [
".amp/skills",
".kiro/skills",
".trae/skills",
".openclaw/skills",
];
const PLATFORM_MAP = {
+83
View File
@@ -0,0 +1,83 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package 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")
}
+26
View File
@@ -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")
}
}
+4
View File
@@ -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] 认证信息已重置")
+1 -1
View File
@@ -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
+59
View File
@@ -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
}
+213
View File
@@ -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")
}
}
+11
View File
@@ -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"
+165
View File
@@ -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
}
+177
View File
@@ -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")
}
}
+60
View File
@@ -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
+438
View File
@@ -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")
}
+172
View File
@@ -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")
}
}
+21
View File
@@ -0,0 +1,21 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
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()
}
+5 -32
View File
@@ -16,7 +16,6 @@ package app
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net"
"net/http"
@@ -89,13 +88,6 @@ func injectStaticServers(servers []edition.ServerInfo) {
// Tests may override discoveryBaseURLOverride to redirect to a local server;
// in that case the registry cache is always bypassed.
func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.Command {
totalStart := time.Now()
defer func() {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] loadDynamicCommands total: %v\n", time.Since(totalStart))
}
}()
store := cacheStoreFromEnv()
partition := config.DefaultPartition
@@ -108,18 +100,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
// --- Cache-first server registry ---
cacheLoadStart := time.Now()
snapshot, freshness, cacheErr := store.LoadRegistry(partition)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache load: %v (err=%v)\n", time.Since(cacheLoadStart), cacheErr)
}
RecordTiming(ctx, "registry_cache", time.Since(cacheLoadStart))
var servers []market.ServerDescriptor
now := store.Now().UTC()
usingCachedRegistry := useCache && cacheErr == nil && len(snapshot.Servers) > 0
if usingCachedRegistry {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] using cached registry: servers=%d, freshness=%s\n", len(snapshot.Servers), freshness)
}
servers = snapshot.Servers
// Only trigger async revalidation in production (no URL override).
// Tests set discoveryBaseURLOverride and control cache expiry directly,
@@ -135,15 +122,10 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
if discoveryBaseURLOverride != "" {
baseURL = discoveryBaseURLOverride
}
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] fetching from market API: %s\n", baseURL)
}
fetchStart := time.Now()
client := market.NewClient(baseURL, ipv4OnlyHTTPClient())
resp, fetchErr := client.FetchServers(ctx, config.DefaultFetchServersLimit)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] market API fetch: %v (err=%v)\n", time.Since(fetchStart), fetchErr)
}
RecordTiming(ctx, "market_fetch", time.Since(fetchStart))
if fetchErr != nil {
slog.Debug("loadDynamicCommands: market API fetch failed", "error", fetchErr)
// Degrade to stale cache if available (production only).
@@ -155,18 +137,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
}
} else {
servers = market.NormalizeServers(resp, "market")
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] normalized servers: %d\n", len(servers))
}
// Persist fresh data (only in non-test mode).
if useCache {
saveStart := time.Now()
if saveErr := store.SaveRegistry(partition, cache.RegistrySnapshot{Servers: servers}); saveErr != nil {
slog.Debug("loadDynamicCommands: failed to save registry cache", "error", saveErr)
}
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache save: %v\n", time.Since(saveStart))
}
RecordTiming(ctx, "cache_save", time.Since(saveStart))
}
}
}
@@ -179,15 +156,11 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
detailStart := time.Now()
detailsByID := loadCachedDetailsFast(store, servers)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] load details: %v\n", time.Since(detailStart))
}
RecordTiming(ctx, "tool_metadata", time.Since(detailStart))
buildStart := time.Now()
cmds := compat.BuildDynamicCommands(servers, runner, detailsByID)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] build commands: %v (count=%d)\n", time.Since(buildStart), len(cmds))
}
RecordTiming(ctx, "build_commands", time.Since(buildStart))
return cmds
}
+671
View File
@@ -0,0 +1,671 @@
// 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/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
"github.com/spf13/cobra"
)
func newPluginCommand() *cobra.Command {
pluginCmd := newPlaceholderParent("plugin", "Manage plugins")
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: "List installed plugins",
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: "Install a plugin",
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: "Show plugin details",
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: "Enable a plugin",
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: "Disable a plugin (managed plugins can be disabled but not removed)",
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: "Remove a user plugin (managed plugins cannot be removed)",
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: "Validate a 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: "Scaffold a new plugin directory",
Example: ` dws plugin create my-tool
dws plugin create my-tool --type managed --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, _ := cmd.Flags().GetString("type")
if pluginType == "" {
pluginType = "user"
}
if pluginType != "managed" && pluginType != "user" {
return apperrors.NewValidation("type must be 'managed' or '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")
cmd.Flags().String("type", "user", "Plugin type: managed or user")
return cmd
}
func newPluginDevCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "dev <dir>",
Short: "Register a local directory as a dev plugin",
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", "Manage plugin configuration")
configCmd.AddCommand(
newPluginConfigSetCommand(),
newPluginConfigGetCommand(),
newPluginConfigListCommand(),
newPluginConfigUnsetCommand(),
)
return configCmd
}
func newPluginConfigSetCommand() *cobra.Command {
return &cobra.Command{
Use: "set <plugin-name> <key> <value>",
Short: "Set a plugin config value",
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: "Get a plugin config value",
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: "List all config values for a plugin",
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: "Remove a plugin config value",
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: "Build plugin's stdio server into a native binary",
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"
}
+564 -18
View File
@@ -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"
@@ -38,6 +43,7 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"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,14 +57,19 @@ type outputFileContextKey struct{}
const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
// Execute runs the root command and returns the process exit code.
func Execute() int {
totalStart := time.Now()
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() {
StopAllStdioClients() // Ensure child processes are terminated on exit
timing.PrintIfEnabled()
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] Execute total: %v\n", time.Since(totalStart))
}
timing.WriteReportIfEnabled(RawVersion(), SanitizeCommand(os.Args))
}()
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
@@ -71,24 +82,14 @@ func Execute() int {
recovery.ResetRuntimeState()
engine := newPipelineEngine()
root := NewRootCommandWithEngine(ctx, engine)
initDuration := time.Since(initStart)
timing.Record("cmd_init", initDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] command init: %v\n", initDuration)
}
timing.Record("cmd_init", time.Since(initStart))
// Run PreParse handlers on raw argv before Cobra parses flags.
// This corrects model-generated errors like --userId → --user-id
// and --limit100 → --limit 100.
pipeline.RunPreParse(root, engine)
execStart := time.Now()
executed, err := root.ExecuteC()
execDuration := time.Since(execStart)
timing.Record("cobra_exec", execDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cobra ExecuteC: %v\n", execDuration)
}
if err != nil {
if executed == nil {
executed = root
@@ -139,6 +140,11 @@ 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)
}
@@ -231,6 +237,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
AuthTokenFunc: func(ctx context.Context) string {
return resolveRuntimeAuthToken(ctx, "")
},
LoggerFunc: FileLoggerInstance,
}
runner := newCommandRunnerWithFlags(loader, flags)
@@ -257,9 +264,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)
},
@@ -277,18 +291,30 @@ 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)
}
if fn := edition.Get().RegisterExtraCommands; fn != nil {
caller := newToolCallerAdapter(runner, flags)
fn(root, caller)
@@ -632,7 +658,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,
@@ -654,6 +684,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.
@@ -933,11 +1023,461 @@ 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()
// 0a. Ensure default managed plugins are installed (first-run bootstrap).
updater := plugin.NewUpdater(pluginLoader.PluginsDir, RawVersion())
// Load TokenData once; reuse for plugin bootstrap, updates, and stdio injection.
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,
}
}
}
accessToken := ""
if tokenData != nil && tokenData.IsAccessTokenValid() {
accessToken = tokenData.AccessToken
}
if accessToken != "" {
bootstrapCtx, bootstrapCancel := context.WithTimeout(context.Background(), 2*time.Minute)
installed := updater.EnsureManaged(bootstrapCtx, accessToken, os.Stderr)
bootstrapCancel()
if len(installed) > 0 {
slog.Debug("plugin: bootstrapped managed plugins", "names", installed)
}
// 0b. Check for managed plugin updates (non-blocking, best-effort).
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
updated := updater.CheckAndUpdate(ctx, accessToken, os.Stderr)
cancel()
if len(updated) > 0 {
slog.Debug("plugin: updated managed plugins", "names", updated)
}
}
// 1. Load official plugins (always enabled)
managedPlugins := pluginLoader.LoadManaged()
// 2. Load user plugins (per settings.json)
userPlugins := pluginLoader.LoadUser()
// 3. Load dev plugins (registered via `dws plugin dev`)
devPlugins := pluginLoader.LoadDev()
allPlugins := append(managedPlugins, userPlugins...)
allPlugins = append(allPlugins, 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})
}
}
}
// Discover tools from HTTP servers in parallel when there are multiple
// servers with auth headers (third-party services with higher latency).
if len(httpServers) > 1 {
type discoveryResult struct {
commands []*cobra.Command
}
results := make([]discoveryResult, len(httpServers))
var wg sync.WaitGroup
for i, ps := range httpServers {
wg.Add(1)
go func(idx int, ps pluginServer) {
defer wg.Done()
results[idx].commands = registerHTTPServer(ps.plugin, ps.srv, tc, runner)
}(i, ps)
}
wg.Wait()
for _, r := range results {
pluginCmds = append(pluginCmds, r.commands...)
}
} else {
for _, ps := range httpServers {
cmds := registerHTTPServer(ps.plugin, ps.srv, tc, runner)
pluginCmds = append(pluginCmds, cmds...)
}
}
// 4. Start stdio MCP servers, discover tools, and build CLI commands
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
}
cmds := registerStdioServer(p, sc, runner)
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",
"managed", len(managedPlugins),
"user", len(userPlugins),
"dev", len(devPlugins),
)
}
return pluginCmds
}
// registerHTTPServer discovers tools from a streamable-http MCP server and
// builds CLI commands. Used for plugin-owned HTTP servers that provide CLI metadata.
//
// 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) []*cobra.Command {
// Use a longer timeout for servers with custom auth headers (third-party
// services may have higher latency than local/DingTalk endpoints).
timeout := 2 * time.Second
if len(srv.AuthHeaders) > 0 {
timeout = 10 * time.Second
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
// If the plugin provides custom auth headers, create a dedicated client
// so the Bearer token is sent to the third-party endpoint.
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
}
if len(toolsResult.Tools) == 0 {
return nil
}
detailsByID := make(map[string][]market.DetailTool)
var detailTools []market.DetailTool
for _, tool := range toolsResult.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 && len(toolsResult.Tools) > 0 {
srv.CLI.ToolOverrides = make(map[string]market.CLIToolOverride, len(toolsResult.Tools))
for _, tool := range toolsResult.Tools {
srv.CLI.ToolOverrides[tool.Name] = market.CLIToolOverride{
CLIName: deriveToolCLIName(tool.Name),
}
}
}
cmds := compat.BuildDynamicCommands(
[]market.ServerDescriptor{srv}, runner, detailsByID)
slog.Debug("plugin: http server registered",
"plugin", p.Manifest.Name, "server", srv.Key,
"tools", len(toolsResult.Tools), "commands", len(cmds))
return cmds
}
// 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.
func registerStdioServer(p *plugin.Plugin, sc plugin.StdioServerClient, runner executor.Runner) []*cobra.Command {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
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
}
if len(toolsResult.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 toolsResult.Tools {
overlay.ToolOverrides[tool.Name] = market.CLIToolOverride{
IsSensitive: tool.Sensitive,
}
}
}
// Construct virtual endpoint and server descriptor.
endpoint := StdioEndpoint(p.Manifest.Name, sc.Key)
source := "plugin"
if p.IsManaged {
source = "plugin-managed"
}
descriptor := market.ServerDescriptor{
Key: sc.Key,
DisplayName: p.Manifest.Name + "/" + sc.Key,
Description: p.Manifest.Description,
Endpoint: endpoint,
Source: source,
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 toolsResult.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(toolsResult.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
@@ -948,6 +1488,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
}
+96
View File
@@ -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()
@@ -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)
}
}
+41 -1
View File
@@ -26,6 +26,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 +46,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 +94,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
}
+208 -46
View File
@@ -15,9 +15,10 @@ package app
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"os"
@@ -29,11 +30,54 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"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"
@@ -76,13 +120,6 @@ type runtimeRunner struct {
}
func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
totalStart := time.Now()
defer func() {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] runtimeRunner.Run total: %v\n", time.Since(totalStart))
}
}()
if r.loader == nil || r.transport == nil {
return r.fallback.Run(ctx, invocation)
}
@@ -96,6 +133,11 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
return r.executeInvocation(ctx, endpoint, invocation)
}
// Prefetch the Keychain token in the background. Keychain access costs
// ~70ms on macOS; starting it here lets the load overlap with endpoint
// resolution and catalog loading below.
go getCachedRuntimeToken(ctx)
if shouldUseDirectRuntime(invocation) {
if endpoint, ok := directRuntimeEndpoint(invocation.CanonicalProduct, invocation.Tool); ok {
return r.executeInvocation(ctx, endpoint, invocation)
@@ -106,7 +148,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, &degraded) {
return executor.Result{}, err
}
}
product, ok := catalog.FindProduct(invocation.CanonicalProduct)
@@ -127,21 +172,60 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
return r.executeInvocation(ctx, endpoint, invocation)
}
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
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
// Lazy bind FileLogger: it may be nil at construction time because
// configureLogLevel runs later in PersistentPreRunE.
if r.transport.FileLogger == nil {
r.transport.FileLogger = FileLoggerInstance()
}
authStart := time.Now()
authToken := r.resolveAuthToken(ctx)
authDuration := time.Since(authStart)
RecordTiming(ctx, "auth_token", authDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] resolveAuthToken: %v\n", authDuration)
fl := r.transport.FileLogger
defer func() {
var errCat, errReason string
if retErr != nil {
var typed *apperrors.Error
if errors.As(retErr, &typed) {
errCat = string(typed.Category)
errReason = typed.Reason
} else {
errCat = "unknown"
errReason = retErr.Error()
}
}
logging.LogCommandEnd(fl, execID,
invocation.CanonicalProduct, invocation.Tool,
retErr == nil, time.Since(invokeStart), errCat, errReason)
}()
// 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 {
timeoutSec = r.globalFlags.Timeout
}
logging.LogCommandStart(fl, execID,
invocation.CanonicalProduct, invocation.Tool, endpoint, version, authToken != "", timeoutSec)
if invocation.DryRun {
return executor.Result{
Invocation: invocation,
@@ -182,25 +266,45 @@ 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 {
var cancel context.CancelFunc
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
defer cancel()
}
callStart := time.Now()
callResult, err := tc.CallTool(ctx, endpoint, invocation.Tool, invocation.Params)
callDuration := time.Since(callStart)
RecordTiming(ctx, "mcp_call", callDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] MCP CallTool: %v\n", callDuration)
}
callResult, err := tc.CallTool(callCtx, endpoint, invocation.Tool, invocation.Params)
RecordTiming(ctx, "mcp_call", time.Since(callStart))
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
}
}
}
captureRuntimeFailure(invocation, err, err)
return executor.Result{}, err
}
if fn := edition.Get().ClassifyToolResult; fn != nil {
if editionErr := fn(callResult.Content); editionErr != nil {
return executor.Result{}, editionErr
}
}
if callResult.IsError {
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
@@ -244,12 +348,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 {
@@ -271,38 +441,30 @@ var (
func getCachedRuntimeToken(ctx context.Context) string {
cachedRuntimeTokenOnce.Do(func() {
loadStart := time.Now()
defer func() {
loadDuration := time.Since(loadStart)
RecordTiming(ctx, "keychain_load", loadDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] getCachedRuntimeToken (first load): %v\n", loadDuration)
}
}()
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
}
// generateExecutionID returns a random 16-char hex string used to correlate
// all log entries (command_start, jsonrpc_request, command_end, etc.) belonging
// to a single command invocation.
func generateExecutionID() string {
b := make([]byte, 8)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
// This should be called after login/logout operations.
func ResetRuntimeTokenCache() {
+139
View File
@@ -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
View File
@@ -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)
+65 -13
View File
@@ -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")
+119
View File
@@ -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
}
+72
View File
@@ -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")
}
}
+215 -11
View File
@@ -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"
)
// Environment variable to enable performance timing output.
const PerfTimingEnv = "DWS_PERF_TIMING"
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{}
@@ -107,32 +136,46 @@ func (tc *TimingCollector) Entries() []TimingEntry {
return result
}
// formatDuration returns a human-friendly duration string.
// Sub-µs → "0µs", sub-ms → microsecond precision (e.g. "142µs"), else → ms.
func formatDuration(d time.Duration) string {
switch {
case d < time.Microsecond:
return "0µs"
case d < time.Millisecond:
return d.Truncate(time.Microsecond).String()
default:
return d.Truncate(time.Millisecond).String()
}
}
// Print writes a summary of all timing entries to the given writer.
func (tc *TimingCollector) Print(w io.Writer) {
if tc == nil || w == nil {
return
}
entries := tc.Entries()
total := tc.Total()
if len(entries) == 0 {
fmt.Fprintf(w, "\n[Timing] Total: %v (no detailed entries)\n", tc.Total().Truncate(time.Millisecond))
fmt.Fprintf(w, "\n[Perf] Total: %v (no detailed entries)\n", formatDuration(total))
return
}
fmt.Fprintln(w)
fmt.Fprintln(w, "[Timing] Execution breakdown:")
fmt.Fprintln(w, "[Perf] Execution breakdown:")
for _, e := range entries {
fmt.Fprintf(w, " %-30s %v\n", e.Name, e.Duration.Truncate(time.Millisecond))
fmt.Fprintf(w, " %-30s %v\n", e.Name, formatDuration(e.Duration))
}
fmt.Fprintf(w, " %-30s %v\n", "──────────────────────────────", "──────────")
fmt.Fprintf(w, " %-30s %v\n", "Total", tc.Total().Truncate(time.Millisecond))
fmt.Fprintf(w, " %-30s %v\n", "Total", formatDuration(total))
}
// PrintIfEnabled prints timing info to stderr if DWS_PERF_TIMING is set.
// PrintIfEnabled prints timing info to stderr if DWS_PERF_DEBUG is set.
func (tc *TimingCollector) PrintIfEnabled() {
if tc == nil {
return
}
if os.Getenv(PerfTimingEnv) == "" {
if os.Getenv(PerfDebugEnv) == "" {
return
}
tc.Print(os.Stderr)
@@ -171,7 +214,168 @@ func StartTiming(ctx context.Context, name string) func() {
return tc.StartTimer(name)
}
// IsPerfTimingEnabled returns true if performance timing output is enabled.
func IsPerfTimingEnabled() bool {
return os.Getenv(PerfTimingEnv) != ""
// IsPerfDebugEnabled returns true if performance debug output is enabled.
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, " ")
}
+292 -12
View File
@@ -16,7 +16,9 @@ package app
import (
"bytes"
"context"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"time"
@@ -87,8 +89,8 @@ func TestTimingCollector_Print(t *testing.T) {
tc.Print(&buf)
output := buf.String()
if !strings.Contains(output, "[Timing]") {
t.Error("output should contain [Timing] header")
if !strings.Contains(output, "[Perf]") {
t.Error("output should contain [Perf] header")
}
if !strings.Contains(output, "auth_token") {
t.Error("output should contain 'auth_token'")
@@ -103,8 +105,8 @@ func TestTimingCollector_Print(t *testing.T) {
func TestTimingCollector_PrintIfEnabled(t *testing.T) {
// Set environment variable
os.Setenv(PerfTimingEnv, "1")
defer os.Unsetenv(PerfTimingEnv)
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
tc := NewTimingCollector()
tc.Record("test_op", 10*time.Millisecond)
@@ -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")
@@ -156,18 +159,295 @@ func TestStartTiming_NoCollector(t *testing.T) {
stop()
}
func TestIsPerfTimingEnabled(t *testing.T) {
func TestIsPerfDebugEnabled(t *testing.T) {
// Clear the env var first
os.Unsetenv(PerfTimingEnv)
os.Unsetenv(PerfDebugEnv)
if IsPerfTimingEnabled() {
t.Error("IsPerfTimingEnabled should return false when env var is not set")
if IsPerfDebugEnabled() {
t.Error("IsPerfDebugEnabled should return false when env var is not set")
}
os.Setenv(PerfTimingEnv, "1")
defer os.Unsetenv(PerfTimingEnv)
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
if !IsPerfTimingEnabled() {
t.Error("IsPerfTimingEnabled should return true when env var is set")
if !IsPerfDebugEnabled() {
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 '框架开销'")
}
}
+35 -8
View File
@@ -28,6 +28,8 @@ var (
ugBoldGrn = color.New(color.Bold, color.FgGreen).SprintFunc()
)
const defaultListLimit = 10
func newUpgradeCommand() *cobra.Command {
var (
flagCheck bool
@@ -36,6 +38,7 @@ func newUpgradeCommand() *cobra.Command {
flagRollback bool
flagForce bool
flagSkipSkills bool
flagAll bool
)
cmd := &cobra.Command{
@@ -47,7 +50,8 @@ func newUpgradeCommand() *cobra.Command {
升级前会自动备份当前版本,可通过 --rollback 回滚。`,
Example: ` dws upgrade # 交互式升级到最新版本
dws upgrade --check # 仅检查是否有新版本
dws upgrade --list # 列出所有可用版本
dws upgrade --list # 列出最近版本
dws upgrade --list --all # 列出所有版本
dws upgrade --version v1.0.5 # 升级到指定版本
dws upgrade --rollback # 回滚到上一版本
dws upgrade -y # 跳过确认直接升级`,
@@ -57,7 +61,11 @@ func newUpgradeCommand() *cobra.Command {
format := resolveUpgradeFormat(cmd)
if flagList {
return runUpgradeList(cmd, format)
limit := defaultListLimit
if flagAll {
limit = 0
}
return runUpgradeList(cmd, format, limit)
}
if flagRollback {
return runUpgradeRollback(yes)
@@ -75,7 +83,8 @@ func newUpgradeCommand() *cobra.Command {
}
cmd.Flags().BoolVar(&flagCheck, "check", false, "仅检查是否有新版本")
cmd.Flags().BoolVar(&flagList, "list", false, "列出所有可用版本")
cmd.Flags().BoolVar(&flagList, "list", false, "列出可用版本")
cmd.Flags().BoolVar(&flagAll, "all", false, "与 --list 搭配,显示所有版本")
cmd.Flags().StringVar(&flagVersion, "version", "", "升级到指定版本")
cmd.Flags().BoolVar(&flagRollback, "rollback", false, "回滚到上一版本")
cmd.Flags().BoolVar(&flagForce, "force", false, "强制重新安装当前版本")
@@ -146,7 +155,9 @@ func runUpgradeCheck(cmd *cobra.Command, format string) error {
// --- dws upgrade --list ---
func runUpgradeList(cmd *cobra.Command, format string) error {
// runUpgradeList displays available versions. When limit > 0, only the most
// recent `limit` versions are shown; pass 0 to show all (--all flag).
func runUpgradeList(cmd *cobra.Command, format string, limit int) error {
client := upgrade.NewClient()
if format != "json" {
@@ -158,6 +169,13 @@ func runUpgradeList(cmd *cobra.Command, format string) error {
return fmt.Errorf("获取版本列表失败: %w", err)
}
totalCount := len(versions)
truncated := false
if limit > 0 && len(versions) > limit {
versions = versions[:limit]
truncated = true
}
currentVer := strings.TrimPrefix(version, "v")
if format == "json" {
@@ -171,13 +189,19 @@ func runUpgradeList(cmd *cobra.Command, format string) error {
"changelog": parseChangelogEntries(v.Changelog, 10),
})
}
return writeJSON(cmd.OutOrStdout(), map[string]any{
result := map[string]any{
"current_version": ensureV(version),
"versions": items,
})
"total": totalCount,
}
if truncated {
result["truncated"] = true
result["shown"] = limit
}
return writeJSON(cmd.OutOrStdout(), result)
}
if len(versions) == 0 {
if totalCount == 0 {
fmt.Printf(" %s\n", ugYellow("未找到任何版本"))
return nil
}
@@ -203,7 +227,10 @@ func runUpgradeList(cmd *cobra.Command, format string) error {
fmt.Println()
fmt.Printf(" %s %s\n", ugBold("当前版本:"), ugBoldGrn(ensureV(version)))
fmt.Printf(" %s\n", ugDim("提示: 使用 dws upgrade --version v1.0.5 安装指定版本"))
if truncated {
fmt.Printf(" %s\n", ugDim(fmt.Sprintf("显示最近 %d 个版本(共 %d 个),使用 --list --all 查看全部", limit, totalCount)))
}
fmt.Printf(" %s\n", ugDim("提示: 使用 dws upgrade --version v1.0.7 安装指定版本"))
return nil
}
+24 -1
View File
@@ -18,8 +18,28 @@ import (
"path/filepath"
"strings"
"sync"
"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,
})
}
const (
// AuthorizeURL is the DingTalk OAuth authorization page.
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
@@ -110,7 +130,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 +209,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
+65 -17
View File
@@ -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
}
+59
View File
@@ -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, &degraded) {
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 != "" {
+77
View File
@@ -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
}
+91 -8
View File
@@ -17,6 +17,7 @@ import (
"context"
"encoding/json"
"fmt"
"log/slog"
"os"
"strings"
"time"
@@ -27,15 +28,86 @@ 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,
})
}
// 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"
DefaultMarketBaseURL = "https://mcp.dingtalk.com"
// defaultDiscoveryTimeout bounds the time spent on live registry discovery.
defaultDiscoveryTimeout = 10 * time.Second
defaultDiscoveryTimeout = 4 * time.Second
)
type CatalogLoader interface {
@@ -92,6 +164,10 @@ type EnvironmentLoader struct {
// AuthTokenFunc returns an access token for MCP discovery requests
// (initialize, tools/list). When nil, discovery runs without auth.
AuthTokenFunc func(context.Context) string
// LoggerFunc returns a structured logger for discovery diagnostics.
// Called lazily because the file logger may not be initialized at
// construction time (it's set up during PersistentPreRunE).
LoggerFunc func() *slog.Logger
}
type cachedCatalogState struct {
@@ -123,17 +199,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 {
@@ -147,14 +229,15 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
transportClient,
store,
)
if l.LoggerFunc != nil {
service.Logger = l.LoggerFunc()
}
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")
@@ -184,10 +267,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))
+5 -1
View File
@@ -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
+70
View File
@@ -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{}
+99
View File
@@ -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)
}
}
}
+127 -21
View File
@@ -21,12 +21,30 @@ import (
"log/slog"
"os"
"strings"
"sync"
"time"
"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"
@@ -35,12 +53,13 @@ const (
var errCLIServerSkipped = errors.New("server marked cli.skip")
type Service struct {
MarketClient *market.Client
Transport *transport.Client
Cache *cache.Store
Tenant string
AuthIdentity string
Logger *slog.Logger
MarketClient *market.Client
Transport *transport.Client
Cache *cache.Store
Tenant string
AuthIdentity string
Logger *slog.Logger
PerServerTimeout time.Duration // overrides perServerDiscoveryTimeout when > 0
}
type RuntimeServer struct {
@@ -152,29 +171,116 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
}, nil
}
const perServerDiscoveryTimeout = 2 * time.Second
func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.ServerDescriptor) ([]RuntimeServer, []RuntimeFailure) {
results := make([]RuntimeServer, 0, len(servers))
failures := make([]RuntimeFailure, 0)
for _, server := range servers {
if server.CLI.Skip {
continue
type discoveryResult struct {
server RuntimeServer
failure *RuntimeFailure
}
filtered := make([]market.ServerDescriptor, 0, len(servers))
for _, srv := range servers {
if !srv.CLI.Skip {
filtered = append(filtered, srv)
}
runtimeServer, err := s.DiscoverServerRuntime(ctx, server)
if err != nil {
if errors.Is(err, errCLIServerSkipped) {
continue
}
if len(filtered) == 0 {
return nil, nil
}
ch := make(chan discoveryResult, len(filtered))
serverTimeout := s.PerServerTimeout
if serverTimeout <= 0 {
serverTimeout = perServerDiscoveryTimeout
}
var wg sync.WaitGroup
for _, srv := range filtered {
wg.Add(1)
go func(server market.ServerDescriptor) {
defer wg.Done()
serverCtx, cancel := context.WithTimeout(ctx, serverTimeout)
defer cancel()
start := time.Now()
rs, err := s.DiscoverServerRuntime(serverCtx, server)
elapsed := time.Since(start)
if err != nil {
if errors.Is(err, errCLIServerSkipped) {
return
}
if s.Logger != nil {
s.Logger.Warn("server_discovery_failed",
slog.String("server_key", server.Key),
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
slog.String("error", err.Error()),
slog.Bool("is_timeout", errors.Is(err, context.DeadlineExceeded)),
)
}
// Per-server sub-context timed out but parent is still alive:
// try cache fallback instead of reporting a hard failure.
if errors.Is(err, context.DeadlineExceeded) && ctx.Err() == nil {
if cached, cacheErr := s.loadServerFromCache(server); cacheErr == nil {
if s.Logger != nil {
s.Logger.Info("server_discovery_cache_fallback",
slog.String("server_key", server.Key),
slog.String("source", cached.Source),
)
}
ch <- discoveryResult{server: cached}
return
}
}
ch <- discoveryResult{failure: &RuntimeFailure{ServerKey: server.Key, Err: err}}
return
}
failures = append(failures, RuntimeFailure{
ServerKey: server.Key,
Err: err,
})
continue
if s.Logger != nil {
s.Logger.Debug("server_discovery_ok",
slog.String("server_key", server.Key),
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
slog.String("source", rs.Source),
)
}
ch <- discoveryResult{server: rs}
}(srv)
}
go func() {
wg.Wait()
close(ch)
}()
results := make([]RuntimeServer, 0, len(filtered))
failures := make([]RuntimeFailure, 0)
for dr := range ch {
if dr.failure != nil {
failures = append(failures, *dr.failure)
} else {
results = append(results, dr.server)
}
results = append(results, runtimeServer)
}
return results, failures
}
// loadServerFromCache tries to load a server's tools from cache, returning a
// degraded RuntimeServer. Used as fallback when a per-server discovery timeout
// fires but the parent context is still alive.
func (s *Service) loadServerFromCache(server market.ServerDescriptor) (RuntimeServer, error) {
partition := s.partition()
snapshot, freshness, err := s.Cache.LoadTools(partition, server.Key)
if err != nil {
return RuntimeServer{}, err
}
server.NegotiatedProtocolVersion = snapshot.ProtocolVersion
server.Source = string(freshness) + "_cache"
server.Degraded = true
return RuntimeServer{
Server: server,
NegotiatedProtocolVersion: snapshot.ProtocolVersion,
Tools: snapshot.Tools,
Source: string(freshness) + "_cache",
Degraded: true,
}, nil
}
func (s *Service) DiscoverDetail(ctx context.Context, server market.ServerDescriptor) (market.DetailResponse, error) {
partition := s.partition()
var fetchErr error
+55
View File
@@ -6,6 +6,7 @@ import (
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
@@ -435,6 +436,60 @@ func TestParseDetailSchema(t *testing.T) {
}
}
func TestDiscoverAllRuntime_TimeoutFallsBackToCache(t *testing.T) {
t.Parallel()
// done signals slow handlers to exit so srv.Close() can complete.
done := make(chan struct{})
// Server that blocks until signalled (simulates an unreachable MCP server).
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
select {
case <-done:
case <-time.After(10 * time.Second):
}
http.Error(w, "timeout", http.StatusServiceUnavailable)
}))
// LIFO: close(done) runs first so handlers exit, then srv.Close() completes.
defer srv.Close()
defer close(done)
svc := newTestService(t, srv.URL, srv)
// Override timeout so the test completes quickly.
svc.PerServerTimeout = 80 * time.Millisecond
server := market.ServerDescriptor{
Key: "slow-server",
Endpoint: srv.URL + "/mcp",
}
// Pre-populate the cache so the fallback has data to return.
partition := "test-tenant/test-identity"
_ = svc.Cache.SaveTools(partition, server.Key, cache.ToolsSnapshot{
ServerKey: server.Key,
ProtocolVersion: "2025-03-26",
Tools: []transport.ToolDescriptor{
{Name: "cached-tool", Description: "from cache"},
},
})
results, failures := svc.DiscoverAllRuntime(context.Background(), []market.ServerDescriptor{server})
if len(failures) != 0 {
t.Fatalf("failures count = %d, want 0 (expected cache fallback)", len(failures))
}
if len(results) != 1 {
t.Fatalf("results count = %d, want 1", len(results))
}
if !results[0].Degraded {
t.Fatal("cache fallback result should be degraded")
}
if len(results[0].Tools) == 0 {
t.Fatal("expected cached tools to be returned")
}
if results[0].Tools[0].Name != "cached-tool" {
t.Fatalf("tool name = %q, want cached-tool", results[0].Tools[0].Name)
}
}
func TestPartition(t *testing.T) {
t.Parallel()
svc := &Service{Tenant: "corp1", AuthIdentity: "user1"}
+19
View File
@@ -198,12 +198,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
}
+59
View File
@@ -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())
}
}
+18
View File
@@ -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"
+43 -6
View File
@@ -66,7 +66,14 @@ func (aitableHandler) Command(runner executor.Runner) *cobra.Command {
return cmd.Help()
},
}
base.AddCommand(newAitableBaseDeleteCommand(runner))
base.AddCommand(
newAitableBaseListCommand(runner),
newAitableBaseSearchCommand(runner),
newAitableBaseGetCommand(runner),
newAitableBaseCreateCommand(runner),
newAitableBaseUpdateCommand(runner),
newAitableBaseDeleteCommand(runner),
)
table := &cobra.Command{
Use: "table",
@@ -78,7 +85,12 @@ func (aitableHandler) Command(runner executor.Runner) *cobra.Command {
return cmd.Help()
},
}
table.AddCommand(newAitableTableDeleteCommand(runner))
table.AddCommand(
newAitableTableGetCommand(runner),
newAitableTableCreateCommand(runner),
newAitableTableUpdateCommand(runner),
newAitableTableDeleteCommand(runner),
)
field := &cobra.Command{
Use: "field",
@@ -90,7 +102,12 @@ func (aitableHandler) Command(runner executor.Runner) *cobra.Command {
return cmd.Help()
},
}
field.AddCommand(newAitableFieldDeleteCommand(runner))
field.AddCommand(
newAitableFieldGetCommand(runner),
newAitableFieldCreateCommand(runner),
newAitableFieldUpdateCommand(runner),
newAitableFieldDeleteCommand(runner),
)
record := &cobra.Command{
Use: "record",
@@ -102,7 +119,24 @@ func (aitableHandler) Command(runner executor.Runner) *cobra.Command {
return cmd.Help()
},
}
record.AddCommand(newAitableRecordDeleteCommand(runner))
record.AddCommand(
newAitableRecordQueryCommand(runner),
newAitableRecordCreateCommand(runner),
newAitableRecordUpdateCommand(runner),
newAitableRecordDeleteCommand(runner),
)
template := &cobra.Command{
Use: "template",
Short: i18n.T("模板搜索"),
Args: cobra.NoArgs,
TraverseChildren: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
return cmd.Help()
},
}
template.AddCommand(newAitableTemplateSearchCommand(runner))
attachment := &cobra.Command{
Use: "attachment",
@@ -114,9 +148,12 @@ func (aitableHandler) Command(runner executor.Runner) *cobra.Command {
return cmd.Help()
},
}
attachment.AddCommand(newAITableUploadFileCommand(runner))
attachment.AddCommand(
newAITableAttachmentUploadCommand(runner),
newAITableUploadFileCommand(runner),
)
root.AddCommand(base, table, field, record, attachment)
root.AddCommand(base, table, field, record, template, attachment)
return root
}
+727
View File
@@ -0,0 +1,727 @@
// 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 helpers
import (
"encoding/json"
"fmt"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/spf13/cobra"
)
// ── base ────────────────────────────────────────────────────
func newAitableBaseListCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "list",
Short: i18n.T("获取 AI 表格列表"),
Example: " dws aitable base list\n dws aitable base list --limit 5 --cursor NEXT_CURSOR",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
params := map[string]any{}
if limit, _ := cmd.Flags().GetInt("limit"); limit > 0 {
params["limit"] = limit
}
if cursor := aitableStringFlag(cmd, "cursor"); cursor != "" {
params["cursor"] = cursor
}
return runAitableTool(cmd, runner, "list_bases", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().Int("limit", 0, i18n.T("每页数量"))
cmd.Flags().String("cursor", "", i18n.T("分页游标"))
return cmd
}
func newAitableBaseSearchCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "search",
Short: i18n.T("搜索 AI 表格"),
Example: " dws aitable base search --query 项目管理",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
query := aitableFlagOrFallback(cmd, "query", "keyword")
if query == "" {
return apperrors.NewValidation("--query is required")
}
params := map[string]any{"query": query}
if cursor := aitableStringFlag(cmd, "cursor"); cursor != "" {
params["cursor"] = cursor
}
return runAitableTool(cmd, runner, "search_bases", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("query", "", i18n.T("Base 名称关键词 (必填)"))
cmd.Flags().String("keyword", "", i18n.T("--query 的别名"))
_ = cmd.Flags().MarkHidden("keyword")
cmd.Flags().String("cursor", "", i18n.T("分页游标"))
return cmd
}
func newAitableBaseGetCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "get",
Short: i18n.T("获取 AI 表格信息"),
Example: " dws aitable base get --base-id BASE_ID",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
return runAitableTool(cmd, runner, "get_base", map[string]any{
"baseId": baseID,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
return cmd
}
func newAitableBaseCreateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "create",
Short: i18n.T("创建 AI 表格"),
Example: " dws aitable base create --name 项目跟踪",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
name, err := aitableRequiredFlag(cmd, "name")
if err != nil {
return err
}
params := map[string]any{"baseName": name}
if templateID := aitableStringFlag(cmd, "template-id"); templateID != "" {
params["templateId"] = templateID
}
return runAitableTool(cmd, runner, "create_base", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("name", "", i18n.T("Base 名称 (必填)"))
cmd.Flags().String("template-id", "", i18n.T("模板 ID"))
return cmd
}
func newAitableBaseUpdateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "update",
Short: i18n.T("更新 AI 表格"),
Example: " dws aitable base update --base-id BASE_ID --name 新名称",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
name, err := aitableRequiredFlag(cmd, "name")
if err != nil {
return err
}
params := map[string]any{
"baseId": baseID,
"newBaseName": name,
}
if desc := aitableStringFlag(cmd, "desc"); desc != "" {
params["description"] = desc
}
return runAitableTool(cmd, runner, "update_base", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("name", "", i18n.T("新名称 (必填)"))
cmd.Flags().String("desc", "", i18n.T("备注文本"))
return cmd
}
// ── table ───────────────────────────────────────────────────
func newAitableTableGetCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "get",
Short: i18n.T("获取数据表"),
Example: " dws aitable table get --base-id BASE_ID\n dws aitable table get --base-id BASE_ID --table-ids tbl1,tbl2",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
params := map[string]any{"baseId": baseID}
if tableIDs := aitableStringFlag(cmd, "table-ids"); tableIDs != "" {
params["tableIds"] = parseAitableCSVValues(tableIDs)
}
return runAitableTool(cmd, runner, "get_tables", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-ids", "", i18n.T("Table ID 列表,逗号分隔"))
return cmd
}
func newAitableTableCreateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "create",
Short: i18n.T("创建数据表"),
Example: " dws aitable table create --base-id BASE_ID --name 任务表 --fields '[{\"fieldName\":\"名称\",\"type\":\"text\"}]'",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableName := aitableFlagOrFallback(cmd, "name", "table-name")
if tableName == "" {
return apperrors.NewValidation("--name is required")
}
fieldsRaw, err := aitableRequiredFlag(cmd, "fields")
if err != nil {
return err
}
fields, err := parseAitableFieldsJSON(fieldsRaw)
if err != nil {
return err
}
return runAitableTool(cmd, runner, "create_table", map[string]any{
"baseId": baseID,
"tableName": tableName,
"fields": fields,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("name", "", i18n.T("表格名称 (必填)"))
cmd.Flags().String("table-name", "", i18n.T("--name 的别名"))
_ = cmd.Flags().MarkHidden("table-name")
cmd.Flags().String("fields", "", i18n.T("字段 JSON 数组 (必填)"))
return cmd
}
func newAitableTableUpdateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "update",
Short: i18n.T("更新数据表"),
Example: " dws aitable table update --base-id BASE_ID --table-id TABLE_ID --name 新表名",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
name, err := aitableRequiredFlag(cmd, "name")
if err != nil {
return err
}
return runAitableTool(cmd, runner, "update_table", map[string]any{
"baseId": baseID,
"tableId": tableID,
"newTableName": name,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("name", "", i18n.T("新表名 (必填)"))
return cmd
}
// ── field ───────────────────────────────────────────────────
func newAitableFieldGetCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "get",
Short: i18n.T("获取字段详情"),
Example: " dws aitable field get --base-id BASE_ID --table-id TABLE_ID\n dws aitable field get --base-id BASE_ID --table-id TABLE_ID --field-ids fld1,fld2",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
params := map[string]any{
"baseId": baseID,
"tableId": tableID,
}
if fieldIDs := aitableStringFlag(cmd, "field-ids"); fieldIDs != "" {
params["fieldIds"] = parseAitableCSVValues(fieldIDs)
}
return runAitableTool(cmd, runner, "get_fields", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("field-ids", "", i18n.T("Field ID 列表,逗号分隔"))
return cmd
}
func newAitableFieldCreateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "create",
Short: i18n.T("创建字段"),
Example: " dws aitable field create --base-id BASE_ID --table-id TABLE_ID --fields '[{\"fieldName\":\"状态\",\"type\":\"singleSelect\"}]'",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
var fields []any
fieldsRaw := aitableStringFlag(cmd, "fields")
if fieldsRaw != "" {
fields, err = parseAitableFieldsJSON(fieldsRaw)
if err != nil {
return err
}
} else {
name, nameErr := aitableRequiredFlag(cmd, "name")
if nameErr != nil {
return apperrors.NewValidation("must specify either --fields or both --name and --type")
}
fieldType, typeErr := aitableRequiredFlag(cmd, "type")
if typeErr != nil {
return apperrors.NewValidation("must specify either --fields or both --name and --type")
}
field := map[string]any{
"fieldName": name,
"type": fieldType,
}
if configRaw := aitableStringFlag(cmd, "config"); configRaw != "" {
configValue, err := parseAitableJSONObject(configRaw, "config")
if err != nil {
return err
}
field["config"] = configValue
}
fields = []any{field}
}
return runAitableTool(cmd, runner, "create_fields", map[string]any{
"baseId": baseID,
"tableId": tableID,
"fields": fields,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("fields", "", i18n.T("字段 JSON 数组"))
cmd.Flags().String("name", "", i18n.T("单字段名称"))
cmd.Flags().String("type", "", i18n.T("单字段类型"))
cmd.Flags().String("config", "", i18n.T("字段配置 JSON"))
return cmd
}
func newAitableFieldUpdateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "update",
Short: i18n.T("更新字段"),
Example: " dws aitable field update --base-id BASE_ID --table-id TABLE_ID --field-id FIELD_ID --name 新字段名",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
fieldID, err := aitableRequiredFlag(cmd, "field-id")
if err != nil {
return err
}
name := aitableStringFlag(cmd, "name")
configRaw := aitableStringFlag(cmd, "config")
if name == "" && configRaw == "" {
return apperrors.NewValidation("at least one of --name or --config is required")
}
params := map[string]any{
"baseId": baseID,
"tableId": tableID,
"fieldId": fieldID,
}
if name != "" {
params["newFieldName"] = name
}
if configRaw != "" {
configValue, err := parseAitableJSONObject(configRaw, "config")
if err != nil {
return err
}
params["config"] = configValue
}
return runAitableTool(cmd, runner, "update_field", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("field-id", "", i18n.T("Field ID (必填)"))
cmd.Flags().String("name", "", i18n.T("新字段名"))
cmd.Flags().String("config", "", i18n.T("字段配置 JSON"))
return cmd
}
// ── record ──────────────────────────────────────────────────
func newAitableRecordQueryCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "query",
Short: i18n.T("查询记录"),
Example: " dws aitable record query --base-id BASE_ID --table-id TABLE_ID --keyword 关键词 --limit 50",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
params := map[string]any{
"baseId": baseID,
"tableId": tableID,
}
if recordIDs := aitableStringFlag(cmd, "record-ids"); recordIDs != "" {
params["recordIds"] = parseAitableCSVValues(recordIDs)
}
if fieldIDs := aitableStringFlag(cmd, "field-ids"); fieldIDs != "" {
params["fieldIds"] = parseAitableCSVValues(fieldIDs)
}
if filtersRaw := aitableStringFlag(cmd, "filters"); filtersRaw != "" {
filters, err := parseAitableJSONObject(filtersRaw, "filters")
if err != nil {
return err
}
params["filters"] = filters
}
if sortRaw := aitableStringFlag(cmd, "sort"); sortRaw != "" {
sortValue, err := parseAitableJSONArray(sortRaw, "sort")
if err != nil {
return err
}
params["sort"] = sortValue
}
if keyword := aitableFlagOrFallback(cmd, "query", "keyword"); keyword != "" {
params["keyword"] = keyword
}
if limit, _ := cmd.Flags().GetInt("limit"); limit > 0 {
params["limit"] = limit
}
if cursor := aitableStringFlag(cmd, "cursor"); cursor != "" {
params["cursor"] = cursor
}
return runAitableTool(cmd, runner, "query_records", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("record-ids", "", i18n.T("Record ID 列表,逗号分隔"))
cmd.Flags().String("field-ids", "", i18n.T("Field ID 列表,逗号分隔"))
cmd.Flags().String("filters", "", i18n.T("过滤条件 JSON"))
cmd.Flags().String("sort", "", i18n.T("排序 JSON 数组"))
cmd.Flags().String("query", "", i18n.T("全文关键词"))
cmd.Flags().String("keyword", "", i18n.T("--query 的别名"))
_ = cmd.Flags().MarkHidden("keyword")
cmd.Flags().Int("limit", 0, i18n.T("单次最大记录数"))
cmd.Flags().String("cursor", "", i18n.T("分页游标"))
return cmd
}
func newAitableRecordCreateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "create",
Short: i18n.T("新增记录"),
Example: " dws aitable record create --base-id BASE_ID --table-id TABLE_ID --records '[{\"cells\":{\"fld1\":\"hello\"}}]'",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
recordsRaw, err := aitableRequiredFlag(cmd, "records")
if err != nil {
return err
}
records, err := parseAitableJSONArray(recordsRaw, "records")
if err != nil {
return err
}
return runAitableTool(cmd, runner, "create_records", map[string]any{
"baseId": baseID,
"tableId": tableID,
"records": records,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("records", "", i18n.T("记录 JSON 数组 (必填)"))
return cmd
}
func newAitableRecordUpdateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "update",
Short: i18n.T("更新记录"),
Example: " dws aitable record update --base-id BASE_ID --table-id TABLE_ID --records '[{\"recordId\":\"rec1\",\"cells\":{\"fld1\":\"updated\"}}]'",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlagOrFallback(cmd, "base-id", "base")
if err != nil {
return err
}
tableID, err := aitableRequiredFlag(cmd, "table-id")
if err != nil {
return err
}
recordsRaw, err := aitableRequiredFlag(cmd, "records")
if err != nil {
return err
}
records, err := parseAitableJSONArray(recordsRaw, "records")
if err != nil {
return err
}
return runAitableTool(cmd, runner, "update_records", map[string]any{
"baseId": baseID,
"tableId": tableID,
"records": records,
})
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("table-id", "", i18n.T("Table ID (必填)"))
cmd.Flags().String("records", "", i18n.T("记录 JSON 数组 (必填)"))
return cmd
}
// ── template ────────────────────────────────────────────────
func newAitableTemplateSearchCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "search",
Short: i18n.T("搜索模板"),
Example: " dws aitable template search --query 项目管理",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
query := aitableFlagOrFallback(cmd, "query", "keyword")
if query == "" {
return apperrors.NewValidation("--query is required")
}
params := map[string]any{"query": query}
if limit, _ := cmd.Flags().GetInt("limit"); limit > 0 {
params["limit"] = limit
}
if cursor := aitableStringFlag(cmd, "cursor"); cursor != "" {
params["cursor"] = cursor
}
return runAitableTool(cmd, runner, "search_templates", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("query", "", i18n.T("模板关键词 (必填)"))
cmd.Flags().String("keyword", "", i18n.T("--query 的别名"))
_ = cmd.Flags().MarkHidden("keyword")
cmd.Flags().Int("limit", 0, i18n.T("每页数量"))
cmd.Flags().String("cursor", "", i18n.T("分页游标"))
return cmd
}
// ── attachment ──────────────────────────────────────────────
func newAITableAttachmentUploadCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "upload",
Short: i18n.T("准备附件上传"),
Example: " dws aitable attachment upload --base-id BASE_ID --file-name report.pdf --size 1024",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
baseID, err := aitableRequiredFlag(cmd, "base-id")
if err != nil {
return err
}
fileName, err := aitableRequiredFlag(cmd, "file-name")
if err != nil {
return err
}
params := map[string]any{
"baseId": baseID,
"fileName": fileName,
}
if size, _ := cmd.Flags().GetInt64("size"); size > 0 {
params["size"] = size
}
if mimeType := aitableStringFlag(cmd, "mime-type"); mimeType != "" {
params["mimeType"] = mimeType
}
return runAitableTool(cmd, runner, "prepare_attachment_upload", params)
},
}
preferLegacyLeaf(cmd)
cmd.Flags().String("base-id", "", i18n.T("Base ID (必填)"))
cmd.Flags().String("file-name", "", i18n.T("文件名 (必填)"))
cmd.Flags().Int64("size", 0, i18n.T("文件大小(字节)"))
cmd.Flags().String("mime-type", "", i18n.T("文件 MIME Type"))
return cmd
}
// ── helpers ────────────────────────────────────────────────
func runAitableTool(cmd *cobra.Command, runner executor.Runner, tool string, params map[string]any) error {
invocation := executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd),
"aitable",
tool,
params,
)
invocation.DryRun = commandDryRun(cmd)
result, err := runner.Run(cmd.Context(), invocation)
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
}
func aitableStringFlag(cmd *cobra.Command, name string) string {
if cmd == nil {
return ""
}
if value, err := cmd.Flags().GetString(name); err == nil && strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
if value, err := cmd.InheritedFlags().GetString(name); err == nil && strings.TrimSpace(value) != "" {
return strings.TrimSpace(value)
}
return ""
}
func aitableFlagOrFallback(cmd *cobra.Command, primary string, aliases ...string) string {
if value := aitableStringFlag(cmd, primary); value != "" {
return value
}
for _, alias := range aliases {
if value := aitableStringFlag(cmd, alias); value != "" {
return value
}
}
return ""
}
func aitableRequiredFlag(cmd *cobra.Command, name string) (string, error) {
if value := aitableStringFlag(cmd, name); value != "" {
return value, nil
}
return "", apperrors.NewValidation(fmt.Sprintf("--%s is required", name))
}
func aitableRequiredFlagOrFallback(cmd *cobra.Command, primary string, aliases ...string) (string, error) {
if value := aitableFlagOrFallback(cmd, primary, aliases...); value != "" {
return value, nil
}
return "", apperrors.NewValidation(fmt.Sprintf("--%s is required", primary))
}
func parseAitableCSVValues(raw string) []string {
parts := strings.Split(raw, ",")
values := make([]string, 0, len(parts))
for _, part := range parts {
if trimmed := strings.TrimSpace(part); trimmed != "" {
values = append(values, trimmed)
}
}
return values
}
func parseAitableFieldsJSON(raw string) ([]any, error) {
var fields []any
if err := json.Unmarshal([]byte(raw), &fields); err == nil {
return fields, nil
}
var wrapper map[string]any
if err := json.Unmarshal([]byte(raw), &wrapper); err == nil {
if wrappedFields, ok := wrapper["fields"].([]any); ok {
return wrappedFields, nil
}
}
return nil, apperrors.NewValidation("--fields JSON parse failed: expect a JSON array")
}
func parseAitableJSONArray(raw, flagName string) ([]any, error) {
var value []any
if err := json.Unmarshal([]byte(raw), &value); err != nil {
return nil, apperrors.NewValidation(fmt.Sprintf("--%s JSON parse failed: %v", flagName, err))
}
return value, nil
}
func parseAitableJSONObject(raw, flagName string) (map[string]any, error) {
var value map[string]any
if err := json.Unmarshal([]byte(raw), &value); err != nil {
return nil, apperrors.NewValidation(fmt.Sprintf("--%s JSON parse failed: %v", flagName, err))
}
return value, nil
}
+11
View 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
+14 -6
View File
@@ -16,6 +16,7 @@ package logging
import (
"context"
"log/slog"
"runtime"
"time"
)
@@ -52,14 +53,15 @@ func LogRequestBody(logger *slog.Logger, method, executionId string, toolName st
)
}
// LogResponse logs a JSON-RPC response at Debug level.
func LogResponse(logger *slog.Logger, method, endpoint string, statusCode int, respSize int, duration time.Duration, err error) {
// LogResponse logs a JSON-RPC response at Debug level (Warn on error).
func LogResponse(logger *slog.Logger, method, endpoint, executionId string, statusCode int, respSize int, duration time.Duration, err error) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("method", method),
slog.String("endpoint", redactEndpoint(endpoint)),
slog.String("execution_id", executionId),
slog.Int("status", statusCode),
slog.Int("resp_size", respSize),
slog.String("duration", duration.Truncate(time.Millisecond).String()),
@@ -137,18 +139,24 @@ func LogErrorClassified(logger *slog.Logger, method, executionId, category, reas
}
// LogCommandStart logs the beginning of a command execution.
func LogCommandStart(logger *slog.Logger, executionId, command, product, tool, version string, authPresent bool) {
func LogCommandStart(logger *slog.Logger, executionId, product, tool, endpoint, version string, authPresent bool, timeoutSec int) {
if logger == nil {
return
}
logger.Info("command_start",
attrs := []slog.Attr{
slog.String("execution_id", executionId),
slog.String("command", command),
slog.String("product", product),
slog.String("tool", tool),
slog.String("endpoint", redactEndpoint(endpoint)),
slog.String("cli_version", version),
slog.String("os", runtime.GOOS),
slog.String("arch", runtime.GOARCH),
slog.Bool("auth_token_present", authPresent),
)
}
if timeoutSec > 0 {
attrs = append(attrs, slog.Int("timeout_sec", timeoutSec))
}
logger.LogAttrs(context.TODO(), slog.LevelInfo, "command_start", attrs...)
}
// LogCommandEnd logs the end of a command execution.
+4 -4
View File
@@ -52,7 +52,7 @@ func TestLogResponseSuccess(t *testing.T) {
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", 200, 1024, 150*time.Millisecond, nil)
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", "exec-1", 200, 1024, 150*time.Millisecond, nil)
out := buf.String()
if !strings.Contains(out, "jsonrpc_response") {
@@ -72,7 +72,7 @@ func TestLogResponseError(t *testing.T) {
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", 500, 0, 2*time.Second, errors.New("connection refused"))
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", "exec-2", 500, 0, 2*time.Second, errors.New("connection refused"))
out := buf.String()
if !strings.Contains(out, "WARN") {
@@ -87,12 +87,12 @@ func TestLogRequestNilLogger(t *testing.T) {
t.Parallel()
// Should not panic
LogRequest(nil, "test", "http://localhost", "", 0)
LogResponse(nil, "test", "http://localhost", 200, 0, 0, nil)
LogResponse(nil, "test", "http://localhost", "", 200, 0, 0, nil)
LogRequestBody(nil, "tools/call", "exec-1", "tool", nil)
LogResponseBody(nil, "tools/call", "exec-1", 200, nil, "")
LogRetryAttempt(nil, "tools/call", "exec-1", 0, 2, 429, 0, nil)
LogErrorClassified(nil, "tools/call", "exec-1", "api", "timeout", 0, 0, true, "")
LogCommandStart(nil, "exec-1", "dws test", "doc", "list", "1.0.0", false)
LogCommandStart(nil, "exec-1", "doc", "list", "https://mcp.example.com", "1.0.0", false, 0)
LogCommandEnd(nil, "exec-1", "doc", "list", true, 0, "", "")
}
+18 -16
View File
@@ -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 {
+34 -9
View File
@@ -148,45 +148,70 @@ func WriteFiltered(w io.Writer, format Format, payload any, fields, jq string) e
}
// ResolveFields extracts the --fields flag value from the command.
// It ensures that we do not mistakenly grab a business parameter also named "fields"
// by matching the flag's usage string against the global root definition.
func ResolveFields(cmd *cobra.Command) string {
if cmd == nil {
return ""
}
rootFlags := rootPersistentFlags(cmd)
if rootFlags == nil {
return ""
}
globalFlag := rootFlags.Lookup("fields")
if globalFlag == nil {
return ""
}
for _, flags := range []*pflag.FlagSet{
cmd.Flags(),
cmd.InheritedFlags(),
rootPersistentFlags(cmd),
rootFlags,
} {
if flags == nil {
continue
}
if f := flags.Lookup("fields"); f != nil && f.Changed {
if v, err := flags.GetString("fields"); err == nil {
return v
// To avoid collision with business flags (e.g. table create --fields),
// verify this flag shares the same usage string as the global one.
if f.Usage == globalFlag.Usage {
if v, err := flags.GetString("fields"); err == nil {
return v
}
}
}
}
return ""
}
// ResolveJQ extracts the --jq flag value from the command. It checks
// local flags, inherited flags, and root persistent flags because
// --jq is registered as a root PersistentFlag.
// ResolveJQ extracts the --jq flag value from the command. It ensures
// that we only grab the global output filter, not a similarly named business parameter.
func ResolveJQ(cmd *cobra.Command) string {
if cmd == nil {
return ""
}
rootFlags := rootPersistentFlags(cmd)
if rootFlags == nil {
return ""
}
globalFlag := rootFlags.Lookup("jq")
if globalFlag == nil {
return ""
}
for _, flags := range []*pflag.FlagSet{
cmd.Flags(),
cmd.InheritedFlags(),
rootPersistentFlags(cmd),
rootFlags,
} {
if flags == nil {
continue
}
if f := flags.Lookup("jq"); f != nil && f.Changed {
if v, err := flags.GetString("jq"); err == nil {
return v
if f.Usage == globalFlag.Usage {
if v, err := flags.GetString("jq"); err == nil {
return v
}
}
}
}
+38
View File
@@ -0,0 +1,38 @@
package output
import (
"github.com/spf13/cobra"
"testing"
)
func TestResolveFieldsShadowing(t *testing.T) {
t.Run("global persistent flag propagates", func(t *testing.T) {
rootCmd := &cobra.Command{Use: "dws"}
rootCmd.PersistentFlags().String("fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
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)
}
})
t.Run("shadowed local flag is ignored", func(t *testing.T) {
rootCmd := &cobra.Command{Use: "dws"}
rootCmd.PersistentFlags().String("fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
bizCmd := &cobra.Command{Use: "biz"}
bizCmd.Flags().String("fields", "", "JSON string array of objects")
rootCmd.AddCommand(bizCmd)
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)
}
})
}
+40
View File
@@ -0,0 +1,40 @@
package output
import (
"bytes"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"testing"
)
func TestUnwrapAndWrite(t *testing.T) {
// Simulate the Result
result := executor.Result{
Invocation: executor.Invocation{
Implemented: true,
Kind: "compat_invocation",
},
Response: map[string]any{
"endpoint": "https://mcp-gw",
"content": map[string]any{},
},
}
var buf bytes.Buffer
Write(&buf, FormatJSON, result)
t.Logf("Output: %s", buf.String())
resultNil := executor.Result{
Invocation: executor.Invocation{
Implemented: true,
Kind: "compat_invocation",
},
Response: map[string]any{
"endpoint": "https://mcp-gw",
"content": nil,
},
}
buf.Reset()
Write(&buf, FormatJSON, resultNil)
t.Logf("Output nil: %s", buf.String())
}
@@ -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)
}
}
+41
View File
@@ -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)
}
}
+39
View File
@@ -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)
}
}
+159
View File
@@ -0,0 +1,159 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
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"
if p.IsManaged {
source = "plugin-managed"
}
// 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
}
+124
View File
@@ -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
}
}
+923
View File
@@ -0,0 +1,923 @@
// 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
}
// LoadManaged scans ~/.dws/plugins/managed/ and returns all valid
// official plugins. Managed plugins are always enabled.
func (l *Loader) LoadManaged() []*Plugin {
managedDir := filepath.Join(l.PluginsDir, "managed")
return l.scanDir(managedDir, true)
}
// 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, false)
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, false)
if p != nil {
qualifiedName := entry.Name() + "/" + p.Manifest.Name
if isPluginEnabled(settings, qualifiedName) {
plugins = append(plugins, p)
}
}
}
}
return plugins
}
// LoadAll loads both managed and user plugins.
func (l *Loader) LoadAll() []*Plugin {
managed := l.LoadManaged()
user := l.LoadUser()
return append(managed, user...)
}
// scanDir reads a directory of plugin subdirectories and loads each one.
func (l *Loader) scanDir(dir string, isManaged bool) []*Plugin {
entries, err := os.ReadDir(dir)
if err != nil {
if !os.IsNotExist(err) {
slog.Debug("plugin: cannot read dir", "path", dir, "error", err)
}
return nil
}
var plugins []*Plugin
for _, entry := range entries {
if !entry.IsDir() {
continue
}
pluginDir := filepath.Join(dir, entry.Name())
p := l.loadPlugin(pluginDir, isManaged)
if p != nil {
plugins = append(plugins, p)
}
}
return plugins
}
// loadPlugin reads and validates a single plugin directory.
func (l *Loader) loadPlugin(dir string, isManaged bool) *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,
IsManaged: isManaged,
}
}
// 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"` // "managed" or "user"
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()
// Managed plugins
managedDir := filepath.Join(l.PluginsDir, "managed")
if entries, err := os.ReadDir(managedDir); err == nil {
for _, entry := range entries {
if !entry.IsDir() {
continue
}
dir := filepath.Join(managedDir, entry.Name())
m, err := ParseManifest(filepath.Join(dir, "plugin.json"))
if err != nil {
continue
}
result = append(result, PluginInfo{
Name: m.Name,
Version: m.Version,
Type: "managed",
Enabled: true, // managed plugins always enabled
Path: dir,
Description: m.Description,
})
}
}
// 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,
IsManaged: false,
}, 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)
}
// Determine install path based on workspace.
var destDir string
var isManaged bool
if workspace == config.OfficialPluginWorkspace {
destDir = filepath.Join(l.PluginsDir, config.PluginManagedDir, manifest.Name)
isManaged = true
} else {
destDir = filepath.Join(l.PluginsDir, config.PluginUserDir, workspace, manifest.Name)
isManaged = false
}
// 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)
}
}
if !isManaged {
qualifiedName := workspace + "/" + manifest.Name
l.setPluginEnabled(qualifiedName, true)
}
return &Plugin{
Manifest: *manifest,
Root: destDir,
IsManaged: isManaged,
}, 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 a user plugin. Returns an error if it's managed.
func (l *Loader) RemovePlugin(name string, keepData bool) error {
// Check managed first — official plugins cannot be removed.
managedDir := filepath.Join(l.PluginsDir, "managed", name)
if _, err := os.Stat(managedDir); err == nil {
return fmt.Errorf("%s is a managed plugin (DingTalk-Real-AI/%s) and cannot be removed.\n To disable it, run: dws plugin disable %s", name, name, name)
}
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, "data", name)
_ = os.RemoveAll(dataDir)
}
l.setPluginEnabled(name, false)
return nil
}
// SetEnabled enables or disables a plugin in settings.json.
func (l *Loader) SetEnabled(name string, enabled bool) error {
// Verify plugin exists
if l.findUserPluginDir(name) == "" {
managedDir := filepath.Join(l.PluginsDir, "managed", name)
if _, err := os.Stat(managedDir); err != nil {
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, false)
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
})
}
+193
View File
@@ -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")
}
}
+258
View File
@@ -0,0 +1,258 @@
// 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
IsManaged bool // true for official (DingTalk-Real-AI) plugins
}
// 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
}
+628
View File
@@ -0,0 +1,628 @@
// 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"
"strings"
"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",
IsManaged: true,
}
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-managed" {
t.Errorf("source = %q, want plugin-managed", 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 (non-managed)", 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",
}
managed := loader.LoadManaged()
if len(managed) != 0 {
t.Errorf("expected 0 managed plugins, got %d", len(managed))
}
user := loader.LoadUser()
if len(user) != 0 {
t.Errorf("expected 0 user plugins, got %d", len(user))
}
}
func TestLoaderLoadManaged(t *testing.T) {
dir := t.TempDir()
managedDir := filepath.Join(dir, "managed", "conference")
if err := os.MkdirAll(managedDir, 0o755); err != nil {
t.Fatal(err)
}
manifest := `{
"name": "conference",
"version": "1.0.0",
"type": "managed",
"mcpServers": {
"conference": {
"type": "streamable-http",
"endpoint": "https://example.com"
}
}
}`
if err := os.WriteFile(filepath.Join(managedDir, "plugin.json"), []byte(manifest), 0o644); err != nil {
t.Fatal(err)
}
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
plugins := loader.LoadManaged()
if len(plugins) != 1 {
t.Fatalf("expected 1 managed plugin, got %d", len(plugins))
}
if plugins[0].Manifest.Name != "conference" {
t.Errorf("name = %q, want conference", plugins[0].Manifest.Name)
}
if !plugins[0].IsManaged {
t.Error("expected IsManaged = true")
}
}
func TestRemoveManagedPluginBlocked(t *testing.T) {
dir := t.TempDir()
managedDir := filepath.Join(dir, "managed", "conference")
if err := os.MkdirAll(managedDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(managedDir, "plugin.json"), []byte(`{"name":"conference","version":"1.0.0"}`), 0o644); err != nil {
t.Fatal(err)
}
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
err := loader.RemovePlugin("conference", false)
if err == nil {
t.Fatal("expected error when removing managed plugin")
}
if !contains(err.Error(), "managed plugin") {
t.Errorf("error message should mention managed plugin, got: %v", err)
}
}
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 TestPromptUpdate(t *testing.T) {
tests := []struct {
name string
input string
want bool
}{
{"empty = yes", "\n", true},
{"y = yes", "y\n", true},
{"Y = yes", "Y\n", true},
{"yes = yes", "yes\n", true},
{"n = no", "n\n", false},
{"no = no", "no\n", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var buf strings.Builder
r := strings.NewReader(tt.input)
got := promptUpdate(&buf, r, "test-plugin", "1.0.0", "2.0.0", "")
if got != tt.want {
t.Errorf("promptUpdate() = %v, want %v", got, tt.want)
}
})
}
}
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
}
+430
View File
@@ -0,0 +1,430 @@
// 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 (
"archive/zip"
"bufio"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// pluginDownloadEndpoint is the API endpoint for downloading plugin updates.
const pluginDownloadEndpoint = "https://aihub.dingtalk.com/cli/download"
// lastCheckFileName stores the last update check timestamp.
const lastCheckFileName = ".last-update-check"
// pluginDownloadTimeout is the timeout for plugin download operations.
const pluginDownloadTimeout = 5 * time.Minute
// Updater checks and applies updates for managed plugins.
type Updater struct {
PluginsDir string
CLIVersion string
Platform string // e.g. "darwin-arm64", "linux-amd64"
mu sync.Mutex
}
// NewUpdater creates an Updater with auto-detected platform.
func NewUpdater(pluginsDir, cliVersion string) *Updater {
return &Updater{
PluginsDir: pluginsDir,
CLIVersion: cliVersion,
Platform: runtime.GOOS + "-" + runtime.GOARCH,
}
}
// remoteVersionInfo holds version metadata returned by the download API.
type remoteVersionInfo struct {
Version string `json:"version"`
DownloadURL string `json:"downloadUrl"`
FileName string `json:"fileName"`
Changelog string `json:"changelog,omitempty"`
}
// pluginDownloadResponse represents the API response from the plugin
// download endpoint.
type pluginDownloadResponse struct {
Success bool `json:"success"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMsg string `json:"errorMsg,omitempty"`
Result *remoteVersionInfo `json:"result,omitempty"`
}
// CheckAndUpdate checks for updates for all managed plugins.
// It reads a last-check timestamp file to avoid checking too frequently.
// Returns the list of updated plugin names.
func (u *Updater) CheckAndUpdate(ctx context.Context, accessToken string, w io.Writer) []string {
u.mu.Lock()
defer u.mu.Unlock()
if !u.shouldCheck() {
slog.Debug("plugin: skipping update check (checked recently)")
return nil
}
managedDir := filepath.Join(u.PluginsDir, config.PluginManagedDir)
entries, err := os.ReadDir(managedDir)
if err != nil {
if !os.IsNotExist(err) {
slog.Warn("plugin: cannot read managed dir for update check",
"path", managedDir, "error", err)
}
u.recordCheckTime()
return nil
}
var updated []string
for _, entry := range entries {
if !entry.IsDir() {
continue
}
pluginDir := filepath.Join(managedDir, entry.Name())
pluginName := config.OfficialPluginWorkspace + "/" + entry.Name()
result := u.checkAndUpdateOne(ctx, accessToken, pluginDir, pluginName, w)
if result != "" {
updated = append(updated, result)
}
}
u.recordCheckTime()
return updated
}
// EnsureManaged checks that every plugin in config.DefaultManagedPlugins
// exists locally under ~/.dws/plugins/managed/. Missing plugins are
// downloaded from the remote API and extracted automatically.
// This runs once on first launch (or after a user deletes the managed dir).
func (u *Updater) EnsureManaged(ctx context.Context, accessToken string, w io.Writer) []string {
if len(config.DefaultManagedPlugins) == 0 {
return nil
}
managedDir := filepath.Join(u.PluginsDir, config.PluginManagedDir)
var installed []string
for _, shortName := range config.DefaultManagedPlugins {
pluginDir := filepath.Join(managedDir, shortName)
// Already exists locally — skip.
if _, err := os.Stat(filepath.Join(pluginDir, "plugin.json")); err == nil {
continue
}
qualifiedName := config.OfficialPluginWorkspace + "/" + shortName
fmt.Fprintf(w, "📦 Pulling built-in plugin %s ...\n", qualifiedName)
remote, err := u.checkRemoteVersion(ctx, accessToken, qualifiedName)
if err != nil {
slog.Warn("plugin: failed to fetch remote info for default plugin",
"plugin", qualifiedName, "error", err)
fmt.Fprintf(w, " ⚠️ Failed to fetch %s info: %v\n", qualifiedName, err)
continue
}
if remote == nil || remote.DownloadURL == "" {
slog.Warn("plugin: no download URL for default plugin",
"plugin", qualifiedName)
fmt.Fprintf(w, " ⚠️ No version available for %s\n", qualifiedName)
continue
}
if err := u.downloadAndInstall(ctx, remote.DownloadURL, pluginDir); err != nil {
slog.Warn("plugin: failed to install default plugin",
"plugin", qualifiedName, "error", err)
fmt.Fprintf(w, " ❌ Failed to install %s: %v\n", qualifiedName, err)
continue
}
fmt.Fprintf(w, " ✅ Installed %s (%s)\n", qualifiedName, remote.Version)
installed = append(installed, qualifiedName)
}
return installed
}
// checkAndUpdateOne checks and potentially updates a single managed plugin.
func (u *Updater) checkAndUpdateOne(
ctx context.Context,
accessToken, pluginDir, pluginName string,
w io.Writer,
) string {
manifest, err := ParseManifest(filepath.Join(pluginDir, "plugin.json"))
if err != nil {
slog.Warn("plugin: cannot parse manifest for update check",
"plugin", pluginName, "error", err)
return ""
}
remote, err := u.checkRemoteVersion(ctx, accessToken, pluginName)
if err != nil {
slog.Warn("plugin: failed to check remote version",
"plugin", pluginName, "error", err)
return ""
}
if remote == nil || remote.Version == "" || remote.DownloadURL == "" {
slog.Debug("plugin: no remote version info available",
"plugin", pluginName)
return ""
}
if compareSemver(remote.Version, manifest.Version) <= 0 {
slog.Debug("plugin: already up to date",
"plugin", pluginName,
"local", manifest.Version,
"remote", remote.Version)
return ""
}
if !promptUpdate(w, os.Stdin, pluginName, manifest.Version, remote.Version, remote.Changelog) {
return ""
}
if err := u.downloadAndInstall(ctx, remote.DownloadURL, pluginDir); err != nil {
slog.Warn("plugin: failed to download and install update",
"plugin", pluginName, "error", err)
fmt.Fprintf(w, " Update failed: %v\n", err)
return ""
}
fmt.Fprintf(w, " ✅ Updated %s to %s\n", pluginName, remote.Version)
return pluginName
}
// shouldCheck returns true if enough time has elapsed since the last check.
func (u *Updater) shouldCheck() bool {
checkFile := filepath.Join(u.PluginsDir, lastCheckFileName)
data, err := os.ReadFile(checkFile)
if err != nil {
return true
}
lastCheck, err := time.Parse(time.RFC3339, strings.TrimSpace(string(data)))
if err != nil {
return true
}
return time.Since(lastCheck) >= config.PluginUpdateCheckInterval
}
// recordCheckTime writes the current time to the last-check file.
func (u *Updater) recordCheckTime() {
checkFile := filepath.Join(u.PluginsDir, lastCheckFileName)
_ = os.MkdirAll(filepath.Dir(checkFile), config.DirPerm)
_ = os.WriteFile(checkFile, []byte(time.Now().Format(time.RFC3339)), config.FilePerm)
}
// checkRemoteVersion queries the aihub API for the latest version.
func (u *Updater) checkRemoteVersion(ctx context.Context, accessToken, pluginName string) (*remoteVersionInfo, error) {
apiURL := fmt.Sprintf("%s?pluginName=%s&platform=%s",
pluginDownloadEndpoint,
url.QueryEscape(pluginName),
url.QueryEscape(u.Platform),
)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
if err != nil {
return nil, fmt.Errorf("create request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("x-user-access-token", accessToken)
client := &http.Client{Timeout: config.HTTPTimeout}
resp, err := client.Do(req)
if err != nil {
return nil, fmt.Errorf("check remote version: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("download API returned HTTP %d", resp.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if err != nil {
return nil, fmt.Errorf("read response: %w", err)
}
var result pluginDownloadResponse
if err := json.Unmarshal(body, &result); err != nil {
return nil, fmt.Errorf("parse response: %w", err)
}
if !result.Success {
errMsg := result.ErrorMsg
if errMsg == "" {
errMsg = result.ErrorCode
}
if errMsg == "" {
errMsg = "unknown error"
}
return nil, fmt.Errorf("API error: %s", errMsg)
}
return result.Result, nil
}
// downloadAndInstall downloads a plugin zip and extracts it, replacing
// the previous version.
func (u *Updater) downloadAndInstall(ctx context.Context, downloadURL, pluginDir string) error {
tempFile, err := os.CreateTemp("", "dws-plugin-update-*.zip")
if err != nil {
return fmt.Errorf("create temp file: %w", err)
}
tempPath := tempFile.Name()
defer os.Remove(tempPath)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
if err != nil {
tempFile.Close()
return fmt.Errorf("create download request: %w", err)
}
client := &http.Client{Timeout: pluginDownloadTimeout}
resp, err := client.Do(req)
if err != nil {
tempFile.Close()
return fmt.Errorf("download plugin: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
tempFile.Close()
return fmt.Errorf("download returned HTTP %d", resp.StatusCode)
}
if _, err := io.Copy(tempFile, resp.Body); err != nil {
tempFile.Close()
return fmt.Errorf("write temp file: %w", err)
}
tempFile.Close()
// Remove old plugin directory contents before extracting.
if err := os.RemoveAll(pluginDir); err != nil {
return fmt.Errorf("remove old plugin: %w", err)
}
if err := extractPluginZip(tempPath, pluginDir); err != nil {
return fmt.Errorf("extract plugin: %w", err)
}
return nil
}
// extractPluginZip extracts a zip archive to the destination directory
// with zip slip protection.
func extractPluginZip(zipPath, destDir string) error {
if err := os.MkdirAll(destDir, 0o755); err != nil {
return fmt.Errorf("create destination directory: %w", err)
}
reader, err := zip.OpenReader(zipPath)
if err != nil {
return fmt.Errorf("open zip: %w", err)
}
defer reader.Close()
cleanDest := filepath.Clean(destDir) + string(os.PathSeparator)
for _, file := range reader.File {
filePath := filepath.Join(destDir, file.Name)
if !strings.HasPrefix(filepath.Clean(filePath), cleanDest) {
return fmt.Errorf("invalid file path in zip: %s", file.Name)
}
// Reject symlinks in ZIP to prevent path traversal attacks.
if file.FileInfo().Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("symlinks are not allowed in plugin zip: %s", file.Name)
}
if file.FileInfo().IsDir() {
if err := os.MkdirAll(filePath, 0o755); err != nil {
return fmt.Errorf("create directory: %w", err)
}
continue
}
if err := os.MkdirAll(filepath.Dir(filePath), 0o755); err != nil {
return fmt.Errorf("create parent directory: %w", err)
}
if err := extractOneFile(file, filePath); err != nil {
return err
}
}
return nil
}
// extractOneFile extracts one file from a zip archive to disk.
func extractOneFile(file *zip.File, destPath string) error {
srcFile, err := file.Open()
if err != nil {
return fmt.Errorf("open file in zip: %w", err)
}
defer srcFile.Close()
fileMode := file.Mode()
if fileMode&0o600 == 0 {
fileMode = 0o644
}
destFile, err := os.OpenFile(destPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
if err != nil {
return fmt.Errorf("create file: %w", err)
}
defer destFile.Close()
if _, err := io.Copy(destFile, srcFile); err != nil {
return fmt.Errorf("extract file: %w", err)
}
return nil
}
// promptUpdate asks the user for confirmation before applying an update.
// Returns true if the user accepts (Y or empty input means yes).
func promptUpdate(w io.Writer, r io.Reader, pluginName, oldVer, newVer, changelog string) bool {
fmt.Fprintf(w, "🔄 %s %s → %s", pluginName, oldVer, newVer)
if changelog != "" {
fmt.Fprintf(w, "\n %s", changelog)
}
fmt.Fprintf(w, "\n Update? [Y/n] ")
scanner := bufio.NewScanner(r)
if !scanner.Scan() {
return false // EOF or error: non-interactive, skip
}
answer := strings.TrimSpace(strings.ToLower(scanner.Text()))
return answer == "" || answer == "y" || answer == "yes"
}
+188
View File
@@ -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 plugin
import (
"archive/zip"
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// makePluginZip creates an in-memory zip containing a valid plugin.json.
func makePluginZip(t *testing.T, name, version string) []byte {
t.Helper()
var buf bytes.Buffer
w := zip.NewWriter(&buf)
manifest := map[string]any{
"name": name,
"version": version,
"mcpServers": map[string]any{
name: map[string]any{
"type": "streamable-http",
"endpoint": "https://example.com/" + name,
},
},
}
data, _ := json.Marshal(manifest)
f, err := w.Create("plugin.json")
if err != nil {
t.Fatal(err)
}
if _, err := f.Write(data); err != nil {
t.Fatal(err)
}
if err := w.Close(); err != nil {
t.Fatal(err)
}
return buf.Bytes()
}
func TestEnsureManaged_PullsMissing(t *testing.T) {
pluginName := "conference"
zipData := makePluginZip(t, pluginName, "1.0.0")
// Serve the zip file.
zipServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/zip")
w.Write(zipData)
}))
defer zipServer.Close()
// Serve the download API returning the zip URL.
apiServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
resp := pluginDownloadResponse{
Success: true,
Result: &remoteVersionInfo{
Version: "1.0.0",
DownloadURL: zipServer.URL + "/conference.zip",
},
}
json.NewEncoder(w).Encode(resp)
}))
defer apiServer.Close()
// Override the download endpoint for this test.
origEndpoint := pluginDownloadEndpoint
defer func() {
// pluginDownloadEndpoint is a const, so we use a workaround:
// we won't restore it — instead we accept the const limitation
// and test via a helper that injects the endpoint.
_ = origEndpoint
}()
tmpDir := t.TempDir()
u := &Updater{
PluginsDir: tmpDir,
CLIVersion: "1.0.0",
Platform: "darwin-arm64",
}
// Patch checkRemoteVersion by using a custom updater method —
// since checkRemoteVersion uses the const endpoint, we test
// downloadAndInstall + EnsureManaged logic directly.
managedDir := filepath.Join(tmpDir, config.PluginManagedDir)
// Verify plugin does not exist yet.
pluginDir := filepath.Join(managedDir, pluginName)
if _, err := os.Stat(filepath.Join(pluginDir, "plugin.json")); err == nil {
t.Fatal("plugin should not exist before EnsureManaged")
}
// Simulate what EnsureManaged does: downloadAndInstall for missing plugin.
if err := os.MkdirAll(managedDir, 0o755); err != nil {
t.Fatal(err)
}
err := u.downloadAndInstall(context.Background(), zipServer.URL+"/conference.zip", pluginDir)
if err != nil {
t.Fatalf("downloadAndInstall: %v", err)
}
// Verify plugin.json was extracted.
m, err := ParseManifest(filepath.Join(pluginDir, "plugin.json"))
if err != nil {
t.Fatalf("ParseManifest after install: %v", err)
}
if m.Name != pluginName {
t.Errorf("name = %q, want %q", m.Name, pluginName)
}
if m.Version != "1.0.0" {
t.Errorf("version = %q, want 1.0.0", m.Version)
}
}
func TestEnsureManaged_SkipsExisting(t *testing.T) {
tmpDir := t.TempDir()
managedDir := filepath.Join(tmpDir, config.PluginManagedDir)
// Pre-create the plugin directory with a valid manifest.
pluginDir := filepath.Join(managedDir, "conference")
if err := os.MkdirAll(pluginDir, 0o755); err != nil {
t.Fatal(err)
}
manifest := `{"name":"conference","version":"1.0.0","mcpServers":{"conference":{"type":"streamable-http","endpoint":"https://example.com"}}}`
if err := os.WriteFile(filepath.Join(pluginDir, "plugin.json"), []byte(manifest), 0o644); err != nil {
t.Fatal(err)
}
u := &Updater{
PluginsDir: tmpDir,
CLIVersion: "1.0.0",
Platform: "darwin-arm64",
}
var output bytes.Buffer
// EnsureManaged should not attempt any download (no token needed since it skips).
installed := u.EnsureManaged(context.Background(), "fake-token", &output)
if len(installed) != 0 {
t.Errorf("expected 0 installs for existing plugin, got %d: %v", len(installed), installed)
}
// Should produce no output since nothing was downloaded.
if strings.Contains(output.String(), "Pulling") {
t.Errorf("unexpected download attempt for existing plugin: %s", output.String())
}
}
func TestExtractPluginZip_ZipSlipProtection(t *testing.T) {
// Create a zip with a path traversal entry.
var buf bytes.Buffer
w := zip.NewWriter(&buf)
f, _ := w.Create("../../etc/passwd")
f.Write([]byte("malicious"))
w.Close()
tmpZip := filepath.Join(t.TempDir(), "bad.zip")
os.WriteFile(tmpZip, buf.Bytes(), 0o644)
destDir := filepath.Join(t.TempDir(), "dest")
err := extractPluginZip(tmpZip, destDir)
if err == nil {
t.Fatal("expected zip slip error, got nil")
}
if !strings.Contains(err.Error(), "invalid file path") {
t.Errorf("unexpected error: %v", err)
}
}
+122 -14
View File
@@ -16,9 +16,13 @@ package transport
import (
"bytes"
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"os"
@@ -26,17 +30,31 @@ import (
"sync"
"time"
"io"
"log/slog"
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/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"
@@ -45,9 +63,9 @@ const (
defaultHTTPTimeout = 30 * time.Second
// Default retry parameters for JSON-RPC calls.
defaultMaxRetries = 2
defaultRetryDelay = 10 * time.Millisecond
defaultRetryMaxDelay = 80 * time.Millisecond
defaultMaxRetries = 1
defaultRetryDelay = 500 * time.Millisecond
defaultRetryMaxDelay = 5 * time.Second
// Security headers
HeaderSource = "X-Cli-Source"
@@ -188,15 +206,40 @@ func (r *ToolCallResult) UnmarshalJSON(data []byte) error {
return fmt.Errorf("unsupported tools/call content shape")
}
// defaultTransport returns a tuned http.Transport for MCP JSON-RPC calls.
// Compared to http.DefaultTransport it adds ResponseHeaderTimeout to detect
// "accepted but never responded" servers faster, and explicit TLS/dial timeouts.
func defaultTransport() *http.Transport {
return &http.Transport{
DialContext: (&net.Dialer{
Timeout: 3 * time.Second,
KeepAlive: 30 * time.Second,
}).DialContext,
TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12},
TLSHandshakeTimeout: 10 * time.Second,
ResponseHeaderTimeout: 20 * time.Second,
ExpectContinueTimeout: 1 * time.Second,
MaxIdleConns: 100,
MaxIdleConnsPerHost: 10,
IdleConnTimeout: 90 * time.Second,
ForceAttemptHTTP2: true,
}
}
func NewClient(httpClient *http.Client) *Client {
if httpClient == nil {
httpClient = &http.Client{
Timeout: defaultHTTPTimeout,
Transport: defaultTransport(),
CheckRedirect: safeRedirectPolicy,
}
} else if httpClient.CheckRedirect == nil {
// Wrap existing client with safe redirect policy
httpClient.CheckRedirect = safeRedirectPolicy
} else {
if httpClient.Transport == nil {
httpClient.Transport = defaultTransport()
}
if httpClient.CheckRedirect == nil {
httpClient.CheckRedirect = safeRedirectPolicy
}
}
return &Client{
HTTPClient: httpClient,
@@ -362,7 +405,7 @@ func (c *Client) callJSONRPC(ctx context.Context, endpoint string, request reque
headerTraceID := ExtractTraceIDFromHeaders(resp.Header)
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
logging.LogResponse(c.FileLogger, request.Method, endpoint, resp.StatusCode, len(data), time.Since(callStart), err)
logging.LogResponse(c.FileLogger, request.Method, endpoint, c.ExecutionId, resp.StatusCode, len(data), time.Since(callStart), err)
if err != nil {
return apperrors.NewDiscovery(
"failed to read JSON-RPC response",
@@ -469,6 +512,9 @@ func (c *Client) doWithRetry(ctx context.Context, endpoint string, body []byte)
resp, err := c.HTTPClient.Do(req)
if err != nil {
lastErr = err
if isTimeoutError(err) {
break
}
} else if !retryable(resp.StatusCode) || attempt == c.MaxRetries {
return resp, nil
} else {
@@ -499,12 +545,16 @@ func (c *Client) doWithRetry(ctx context.Context, endpoint string, body []byte)
}
}
}
reason, hint := classifyRequestFailure(lastErr)
logging.LogErrorClassified(c.FileLogger, "jsonrpc", c.ExecutionId,
string(apperrors.CategoryDiscovery), reason, 0, 0,
!isTimeoutError(lastErr), "")
return nil, apperrors.NewDiscovery(
fmt.Sprintf("request to %s failed: %v", RedactURL(endpoint), lastErr),
apperrors.WithOperation("jsonrpc"),
apperrors.WithReason("request_failed"),
apperrors.WithRetryable(true),
apperrors.WithHint(i18n.T("请检查网络连通性和 MCP 服务状态后重试。")),
apperrors.WithReason(reason),
apperrors.WithRetryable(!isTimeoutError(lastErr)),
apperrors.WithHint(hint),
apperrors.WithActions(discoveryActions("")...),
apperrors.WithCause(&CallError{
Stage: CallStageRequest,
@@ -517,6 +567,64 @@ func retryable(statusCode int) bool {
return statusCode == http.StatusTooManyRequests || statusCode >= http.StatusInternalServerError
}
// isTimeoutError returns true for errors caused by context deadline or HTTP
// client timeout. These are typically deterministic (server overloaded or
// unreachable) and retrying immediately is unlikely to help.
func isTimeoutError(err error) bool {
if err == nil {
return false
}
if errors.Is(err, context.DeadlineExceeded) {
return true
}
if errors.Is(err, context.Canceled) {
return true
}
if os.IsTimeout(err) {
return true
}
msg := err.Error()
return strings.Contains(msg, "Client.Timeout exceeded") ||
strings.Contains(msg, "TLS handshake timeout")
}
// classifyRequestFailure returns a machine-readable reason and a user-facing
// hint tailored to the specific failure type, so users get actionable guidance
// instead of opaque Go error strings.
func classifyRequestFailure(err error) (reason, hint string) {
if err == nil {
return "request_failed", i18n.T("请检查网络连通性和 MCP 服务状态后重试。")
}
msg := err.Error()
switch {
case errors.Is(err, context.DeadlineExceeded):
return "request_timeout",
i18n.T("请求超时(上下文截止时间已到)。可通过 --timeout 增大超时时间,或检查网络连接。")
case errors.Is(err, context.Canceled):
return "request_cancelled",
i18n.T("请求已取消。如果非手动取消,请检查调用侧超时设置。")
case strings.Contains(msg, "Client.Timeout exceeded"):
return "http_client_timeout",
i18n.T("HTTP 请求超时(等待服务端响应超时)。可通过 --timeout 增大超时时间,或检查服务端是否正常。")
case strings.Contains(msg, "TLS handshake timeout"):
return "tls_timeout",
i18n.T("TLS 握手超时。请检查网络连接或代理设置。")
case strings.Contains(msg, "connection refused"):
return "connection_refused",
i18n.T("连接被拒绝。请确认服务端已启动并正在监听。")
case strings.Contains(msg, "no such host"):
return "dns_resolution_failed",
i18n.T("DNS 解析失败。请检查域名拼写和网络 DNS 配置。")
case strings.Contains(msg, "i/o timeout"):
return "io_timeout",
i18n.T("网络 I/O 超时。可通过 --timeout 增大超时时间,或检查网络连接。")
default:
return "request_failed",
i18n.T("请检查网络连通性和 MCP 服务状态后重试。")
}
}
func respRetryAfter(resp *http.Response) string {
if resp == nil {
return ""
+190
View File
@@ -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, &params)
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")
}
}
+245
View File
@@ -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
View File
@@ -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)
}
+24
View File
@@ -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"
+1
View File
@@ -35,6 +35,7 @@ var knownSkillDirs = []string{
".amp/skills",
".kiro/skills",
".trae/skills",
".openclaw/skills",
}
// skillDirBlacklist contains parent directories whose skills are managed by
+1 -1
View File
@@ -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
}
+22
View File
@@ -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()
}
+33
View File
@@ -106,3 +106,36 @@ const (
// MaxUploadFileSize is the maximum file size for attachment uploads.
MaxUploadFileSize int64 = 100 * 1024 * 1024 // 100 MB
)
// ── Plugin system ──────────────────────────────────────────────────────
const (
// PluginManagedDir is the subdirectory under ~/.dws/plugins/ for
// official (DingTalk-Real-AI) plugins that are auto-pulled.
PluginManagedDir = "managed"
// PluginUserDir is the subdirectory under ~/.dws/plugins/ for
// user-installed third-party plugins.
PluginUserDir = "user"
// PluginDataDir is the subdirectory under ~/.dws/plugins/ for
// plugin persistent data that survives across version updates.
PluginDataDir = "data"
// PluginUpdateCheckInterval is how often to check for official
// plugin updates (at most once per interval per CLI invocation).
PluginUpdateCheckInterval = 1 * time.Hour
// PluginHookTimeout is the default timeout for plugin hook commands.
PluginHookTimeout = 30 * time.Second
// OfficialPluginWorkspace is the workspace name that identifies
// official plugins. Plugins under this workspace are auto-pulled.
OfficialPluginWorkspace = "DingTalk-Real-AI"
)
// DefaultManagedPlugins lists the official plugins that should be
// automatically pulled on first run if not already present locally.
// Each entry is the short plugin name (without the workspace prefix);
// the full qualified name is OfficialPluginWorkspace + "/" + name.
var DefaultManagedPlugins = []string{}
+155
View File
@@ -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
}
+164
View File
@@ -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")
}
}
+23
View File
@@ -77,10 +77,33 @@ type Hooks struct {
OnAuthError func(configDir string, err error) error
TokenProvider func(ctx context.Context, fallback func() (string, error)) (string, error)
// --- token persistence (overlay-only) ---
// When non-nil, these override the default keychain-based token storage.
// The data parameter is JSON-serialized TokenData.
SaveToken func(configDir string, data []byte) error
LoadToken func(configDir string) ([]byte, error)
DeleteToken func(configDir string) error
// --- auth credentials (overlay-only) ---
AuthClientID string // non-empty overrides DefaultClientID
AuthClientFromMCP bool // true routes OAuth through MCP endpoints
// --- 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
// ClassifyToolResult is called before the framework's default business-error
// detection on MCP tool results. If it returns a non-nil error, that error
// is used instead of the generic CategoryAPI business error. Editions use
// this to return custom error types with specific exit codes (e.g. PAT
// authorization errors with exit code 4).
ClassifyToolResult func(content map[string]any) error
}
var (
+29
View File
@@ -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)
}
+108 -26
View File
@@ -1,21 +1,23 @@
#!/bin/sh
set -eu
# Install DWS agent skills from GitHub into detected agent directories.
# Install DWS agent skills from GitHub Releases into agent skill directories.
# Usage:
# curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-skills.sh | sh
#
# The script downloads the dws-skills.zip release asset from GitHub Releases
# and copies it into every detected agent skills directory in the current
# project.
# Downloads dws-skills.zip from GitHub Releases and copies it under each target
# path using the same rules as build/npm/install.js installSkillsToHomes
# (AGENT_DIRS + parent-directory gate), with root defaulting to the current
# directory. Set DWS_SKILLS_ROOT=$HOME to match npm install layout exactly.
#
# Environment variables (optional):
# DWS_VERSION — release tag (default: latest)
# DWS_SKILLS_ROOT — base path for agent dirs (default: $PWD)
REPO="DingTalk-Real-AI/dingtalk-workspace-cli"
VERSION="${DWS_VERSION:-latest}"
SKILL_NAME="dws"
# ── Agent directory to install skills into ───────────────────────────────────
# Only install to .agents/skills — most agents can fall back to this directory.
AGENT_DIR=".agents/skills"
ROOT="${DWS_SKILLS_ROOT:-$PWD}"
# ── Helpers ──────────────────────────────────────────────────────────────────
@@ -51,14 +53,107 @@ extract_zip() {
exit 1
}
# One-line summary copy (2nd+ targets).
_copy_skill_summary() {
_src="$1"
_dest="$2"
_label="$3"
if [ -d "$_dest" ]; then
rm -rf "$_dest"
fi
mkdir -p "$_dest"
cp -R "$_src/"* "$_dest/" 2>/dev/null || cp -r "$_src/"* "$_dest/"
file_count="$(find "$_dest" -type f | wc -l | tr -d ' ')"
printf ' ✅ Skills → %s (%s files)\n' "$_label" "$file_count"
}
# Full copy with top-level listing (1st target).
_copy_skill() {
_src="$1"
_dest="$2"
_label="$3"
if [ -d "$_dest" ]; then
rm -rf "$_dest"
fi
mkdir -p "$_dest"
cp -R "$_src/"* "$_dest/" 2>/dev/null || cp -r "$_src/"* "$_dest/"
file_count="$(find "$_dest" -type f | wc -l | tr -d ' ')"
printf ' ✅ Skills → %s (%s files)\n' "$_label" "$file_count"
for entry in "$_dest"/*; do
entry_name="$(basename "$entry")"
if [ -d "$entry" ]; then
sub_count="$(find "$entry" -type f | wc -l | tr -d ' ')"
printf ' 📁 %s/ (%s files)\n' "$entry_name" "$sub_count"
else
printf ' 📄 %s\n' "$entry_name"
fi
done
}
# Same semantics as build/npm/install.js installSkillsToHomes (root = DWS_SKILLS_ROOT or PWD).
install_skills_to_root() {
skill_src="$1"
root="$2"
installed=0
idx=0
for agent_dir in \
".agents/skills" \
".claude/skills" \
".cursor/skills" \
".gemini/skills" \
".codex/skills" \
".github/skills" \
".windsurf/skills" \
".augment/skills" \
".cline/skills" \
".amp/skills" \
".kiro/skills" \
".trae/skills" \
".openclaw/skills"
do
base_dir="$root/$agent_dir"
parent_gate="$(dirname "$base_dir")"
if [ "$idx" -gt 0 ] && [ ! -e "$parent_gate" ]; then
idx=$((idx + 1))
continue
fi
dest="$base_dir/$SKILL_NAME"
if [ "$root" = "$HOME" ]; then
label="~/$agent_dir/$SKILL_NAME"
else
label="$root/$agent_dir/$SKILL_NAME"
fi
if [ "$installed" -eq 0 ]; then
_copy_skill "$skill_src" "$dest" "$label"
else
_copy_skill_summary "$skill_src" "$dest" "$label"
fi
installed=$((installed + 1))
idx=$((idx + 1))
done
if [ "$installed" -eq 0 ]; then
if [ "$root" = "$HOME" ]; then
flabel="~/.agents/skills/$SKILL_NAME"
else
flabel="$root/.agents/skills/$SKILL_NAME"
fi
_copy_skill "$skill_src" "$root/.agents/skills/$SKILL_NAME" "$flabel"
fi
}
# ── Main ─────────────────────────────────────────────────────────────────────
main() {
need_cmd curl
resolve_version
CWD="$(pwd)"
printf '\n'
printf ' ┌──────────────────────────────────────┐\n'
printf ' │ DWS Skill Installer │\n'
@@ -66,7 +161,6 @@ main() {
printf ' └──────────────────────────────────────┘\n'
printf '\n'
# Download the tarball to a temp directory
TMPDIR_WORK="$(mktemp -d)"
trap 'rm -rf "$TMPDIR_WORK"' EXIT INT TERM
@@ -85,21 +179,9 @@ main() {
exit 1
fi
# Install to .agents/skills only
dest="$CWD/$AGENT_DIR/$SKILL_NAME"
# Remove existing installation
if [ -d "$dest" ]; then
rm -rf "$dest"
fi
# Copy skill files
mkdir -p "$dest"
cp -R "$SKILL_SRC/"* "$dest/"
file_count="$(find "$dest" -type f | wc -l | tr -d ' ')"
printf ' ✅ Universal (.agents)\n'
printf ' → %s/%s (%s files)\n' "$AGENT_DIR" "$SKILL_NAME" "$file_count"
printf '\n'
printf ' Installing under root: %s\n' "$ROOT"
install_skills_to_root "$SKILL_SRC" "$ROOT"
printf '\n'
printf ' 📖 Skill includes:\n'
+69 -8
View File
@@ -17,6 +17,8 @@
# DWS_ARCH — architecture override (amd64 or arm64)
# DWS_NO_SKILLS — set to 1 to skip skills install
# DWS_SKILLS_ONLY — set to 1 to install only skills
#
# Agent skills paths follow build/npm/install.js AGENT_DIRS (order and entries must match).
$ErrorActionPreference = "Stop"
@@ -28,8 +30,22 @@ $NoSkills = $env:DWS_NO_SKILLS -eq "1"
$SkillsOnly = $env:DWS_SKILLS_ONLY -eq "1"
$SkillName = "dws"
# Agent directory to install skills into — most agents can fall back to .agents\skills
$AgentDir = ".agents\skills"
# Agent skill base directories (same order as build/npm/install.js AGENT_DIRS).
$AgentDirs = @(
".agents\skills",
".claude\skills",
".cursor\skills",
".gemini\skills",
".codex\skills",
".github\skills",
".windsurf\skills",
".augment\skills",
".cline\skills",
".amp\skills",
".kiro\skills",
".trae\skills",
".openclaw\skills"
)
# ── Helpers ──────────────────────────────────────────────────────────────────
@@ -170,6 +186,17 @@ function Copy-SkillToDir {
}
}
function Copy-SkillToDirSummary {
param([string]$SkillSrc, [string]$Dest, [string]$Label)
if (Test-Path $Dest) {
Remove-Item -Path $Dest -Recurse -Force
}
$fileCount = Copy-DirRecursive -Source $SkillSrc -Destination $Dest
Write-Say "✅ Skills → $Label ($fileCount files)"
}
function Resolve-SourceRoot {
$scriptPath = $PSScriptRoot
if (-not $scriptPath) { return $null }
@@ -280,9 +307,45 @@ function Install-SkillsLocal {
Write-Say ""
Write-Say "📦 Installing agent skills from local source: $skillSrc"
$dest = Join-Path (Join-Path $HOME $AgentDir) $SkillName
$label = "~\$AgentDir\$SkillName"
Copy-SkillToDir -SkillSrc $skillSrc -Dest $dest -Label $label
Install-SkillsToHomes -SkillSrc $skillSrc -Root $HOME
}
function Install-SkillsToHomes {
param(
[string]$SkillSrc,
[string]$Root = $HOME
)
$installed = 0
for ($i = 0; $i -lt $AgentDirs.Count; $i++) {
$agentDir = $AgentDirs[$i]
$baseDir = Join-Path $Root $agentDir
$parentGate = Split-Path $baseDir -Parent
if ($i -gt 0 -and !(Test-Path $parentGate)) {
continue
}
$dest = Join-Path $baseDir $SkillName
if ($Root -eq $HOME) {
$label = "~\$agentDir\$SkillName"
} else {
$label = Join-Path $Root (Join-Path $agentDir $SkillName)
}
if ($installed -eq 0) {
Copy-SkillToDir -SkillSrc $SkillSrc -Dest $dest -Label $label
} else {
Copy-SkillToDirSummary -SkillSrc $SkillSrc -Dest $dest -Label $label
}
$installed++
}
if ($installed -eq 0) {
$fallback = Join-Path (Join-Path $Root ".agents\skills") $SkillName
if ($Root -eq $HOME) {
$flabel = "~\.agents\skills\$SkillName"
} else {
$flabel = Join-Path $Root (Join-Path ".agents\skills" $SkillName)
}
Copy-SkillToDir -SkillSrc $SkillSrc -Dest $fallback -Label $flabel
}
}
# ── Install Binary from Source ───────────────────────────────────────────────
@@ -359,9 +422,7 @@ function Install-Skills {
return
}
$dest = Join-Path (Join-Path $HOME $AgentDir) $SkillName
$label = "~\$AgentDir\$SkillName"
Copy-SkillToDir -SkillSrc $skillSrc -Dest $dest -Label $label
Install-SkillsToHomes -SkillSrc $skillSrc -Root $HOME
} finally {
Remove-Item -Path $tmpDir -Recurse -Force -ErrorAction SilentlyContinue
}
+78 -10
View File
@@ -14,6 +14,8 @@
# DWS_VERSION — version to install (default: latest)
# DWS_NO_SKILLS — set to 1 to skip skills install
# DWS_SKILLS_ONLY — set to 1 to install only skills (skip binary)
#
# Agent skills paths follow build/npm/install.js AGENT_DIRS (order and entries must match).
set -eu
@@ -26,10 +28,6 @@ NO_SKILLS="${DWS_NO_SKILLS:-0}"
SKILLS_ONLY="${DWS_SKILLS_ONLY:-0}"
SKILL_NAME="dws"
# ── Agent directory to install skills into ───────────────────────────────────
# Only install to .agents/skills — most agents can fall back to this directory.
AGENT_DIR=".agents/skills"
# ── Helpers ──────────────────────────────────────────────────────────────────
say() {
@@ -179,13 +177,85 @@ install_skills_local() {
say ""
say "📦 Installing agent skills from local source: ${skill_src}"
dest="$HOME/$AGENT_DIR/$SKILL_NAME"
display_path="~/$AGENT_DIR/$SKILL_NAME"
_copy_skill "$skill_src" "$dest" "$display_path"
install_skills_to_homes "$skill_src"
return 0
}
# Install skill tree into all agent homes (same rules as build/npm/install.js installSkillsToHomes).
install_skills_to_homes() {
skill_src="$1"
root="${HOME}"
installed=0
idx=0
for agent_dir in \
".agents/skills" \
".claude/skills" \
".cursor/skills" \
".gemini/skills" \
".codex/skills" \
".github/skills" \
".windsurf/skills" \
".augment/skills" \
".cline/skills" \
".amp/skills" \
".kiro/skills" \
".trae/skills" \
".openclaw/skills"
do
base_dir="$root/$agent_dir"
parent_gate="$(dirname "$base_dir")"
if [ "$idx" -gt 0 ] && [ ! -e "$parent_gate" ]; then
idx=$((idx + 1))
continue
fi
dest="$base_dir/$SKILL_NAME"
case "$root" in
"$HOME")
label="~/$agent_dir/$SKILL_NAME"
;;
*)
label="$root/$agent_dir/$SKILL_NAME"
;;
esac
if [ "$installed" -eq 0 ]; then
_copy_skill "$skill_src" "$dest" "$label"
else
_copy_skill_summary "$skill_src" "$dest" "$label"
fi
installed=$((installed + 1))
idx=$((idx + 1))
done
if [ "$installed" -eq 0 ]; then
case "$root" in
"$HOME")
flabel="~/.agents/skills/$SKILL_NAME"
;;
*)
flabel="$root/.agents/skills/$SKILL_NAME"
;;
esac
_copy_skill "$skill_src" "$root/.agents/skills/$SKILL_NAME" "$flabel"
fi
}
# One-line summary copy (used for 2nd+ agent targets).
_copy_skill_summary() {
_src="$1"
_dest="$2"
_label="$3"
if [ -d "$_dest" ]; then
rm -rf "$_dest"
fi
mkdir -p "$_dest"
cp -R "$_src/"* "$_dest/" 2>/dev/null || cp -r "$_src/"* "$_dest/"
file_count="$(find "$_dest" -type f | wc -l | tr -d ' ')"
say "✅ Skills → ${_label} (${file_count} files)"
}
# Helper: copy skill files to a destination and print details
_copy_skill() {
_src="$1"
@@ -349,9 +419,7 @@ install_skills() {
fi
fi
dest="$HOME/$AGENT_DIR/$SKILL_NAME"
display_path="~/$AGENT_DIR/$SKILL_NAME"
_copy_skill "$skill_src" "$dest" "$display_path"
install_skills_to_homes "$skill_src"
rm -rf "$tmpdir_skills"
}
@@ -36,6 +36,7 @@ HOME_AGENT_PARENTS="
.amp
.kiro
.trae
.openclaw
"
HOME_SKILL_TARGETS="
.agents/skills/dws
@@ -50,6 +51,7 @@ HOME_SKILL_TARGETS="
.amp/skills/dws
.kiro/skills/dws
.trae/skills/dws
.openclaw/skills/dws
"
cleanup() {
if command -v brew >/dev/null 2>&1; then
+1 -1
View File
@@ -1,7 +1,7 @@
---
name: dws
description: 管理钉钉产品能力(AI表格/日历/通讯录/群聊与机器人/待办/审批/考勤/日志/DING消息/工作台/开放平台文档等)。当用户需要操作表格数据、管理日程会议、查询通讯录、管理群聊、机器人发消息、创建待办、提交审批、查看考勤、提交日报周报(钉钉日志模版)时使用。
cli_version: ">=1.1.0"
cli_version: ">=1.0.6"
---
# 钉钉全产品 Skill
+41
View File
@@ -0,0 +1,41 @@
package cli_compat_test
import (
"bytes"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app"
)
func TestDebugAitableTableCreate(t *testing.T) {
_ = setupTestDeps(t, "aitable")
root := app.NewRootCommand()
cliArgs := []string{"-f", "json", "aitable", "table", "create",
"--base-id", "B1", "--name", "任务表",
"--fields", `[{"fieldName":"名称","type":"text"}]`,
}
var out bytes.Buffer
var errOut bytes.Buffer
root.SetOut(&out)
root.SetErr(&errOut)
root.SetArgs(cliArgs)
err := root.Execute()
t.Logf("Execute error: %v", err)
t.Logf("Stdout: [%s]", out.String())
t.Logf("Stderr: [%s]", errOut.String())
// Check all aitable subcommands
aitableCmd, _, _ := root.Find([]string{"aitable"})
if aitableCmd != nil {
t.Logf("aitable subcommands:")
for _, grp := range aitableCmd.Commands() {
t.Logf(" %s:", grp.Use)
for _, sub := range grp.Commands() {
t.Logf(" %s (hidden=%v)", sub.Use, sub.Hidden)
}
}
}
}
+9
View File
@@ -10,6 +10,7 @@ import (
var expectedHomeSkillTargets = []string{
".agents/skills/dws",
".cursor/skills/dws",
}
func TestInstallScriptSourceModeInstallsBinary(t *testing.T) {
@@ -107,6 +108,11 @@ done
mustWriteFile(t, filepath.Join(stubRoot, "make"), []byte(makeStub), 0o755)
mustWriteFile(t, filepath.Join(stubRoot, "go"), []byte("#!/bin/sh\ntrue\n"), 0o755)
// Gate for index>0 agent dirs (matches build/npm/install.js): parent must exist.
if err := os.MkdirAll(filepath.Join(fakeHome, ".cursor"), 0o755); err != nil {
t.Fatalf("MkdirAll(.cursor) error = %v", err)
}
cmd := exec.Command("sh", scriptPath)
cmd.Env = append(os.Environ(),
"HOME="+fakeHome,
@@ -147,6 +153,9 @@ func TestInstallPowerShellScriptInstallsToAgentsDir(t *testing.T) {
if !strings.Contains(text, ".agents\\skills") {
t.Fatalf("install.ps1 missing .agents\\skills")
}
if !strings.Contains(text, ".cursor\\skills") {
t.Fatalf("install.ps1 missing .cursor\\skills (AGENT_DIRS must match build/npm/install.js)")
}
}
func TestInstallScriptsUseGitHubReleaseSkillsAsset(t *testing.T) {
+1
View File
@@ -22,6 +22,7 @@ var expectedPackagedSkillTargets = []string{
".amp/skills/dws",
".kiro/skills/dws",
".trae/skills/dws",
".openclaw/skills/dws",
}
// seedDistArtifacts creates fake goreleaser output archives (empty tar.gz/zip
+34
View File
@@ -0,0 +1,34 @@
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))
}