Compare commits

...
168 Commits
Author SHA1 Message Date
修雨 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
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
fantiu 57c93243a0 Merge pull request #75 from fantiu/feat-login-upgrade
feat: add self-upgrade command with GitHub Releases, atomic flow, and cross-platform support
2026-04-02 17:00:42 +08:00
fantiu 4259336e6d style: gofmt formatting for upgrade files 2026-04-02 16:48:06 +08:00
fantiu a18ce2e54d feat: add self-upgrade command with GitHub Releases, atomic flow, and cross-platform support 2026-04-02 16:34:51 +08:00
fantiu a0dc5d6183 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-login-upgrade
feat: add upgrade command.
2026-04-02 15:44:49 +08:00
fantiu 5149f6808f Merge branch 'DingTalk-Real-AI:main' into main 2026-04-02 15:43:07 +08:00
meng93 918db33a8b Merge pull request #74 from DingTalk-Real-AI/feat/issue-label
feat/issue label
2026-04-02 14:54:08 +08:00
meng93 95cbde9187 Merge branch 'main' of https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli into feat/issue-label 2026-04-02 14:48:42 +08:00
meng93 5e194393ff feat: to #73551688 支持上报用例集标签 2026-04-02 14:48:02 +08:00
fantiu 3da572a76a Merge remote-tracking branch 'origin/main' into feat-login-upgrade 2026-04-02 11:47:58 +08:00
github-actions[bot] a116cba8ba chore: update coverage badge [skip ci] 2026-04-02 03:32:07 +00:00
fantiu fe5952fe14 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-02 11:30:30 +08:00
coffeeBigSir 2ab45ffd90 Merge pull request #72 from fantiu/feat-check-enable
feat(auth): switch CLI auth check to fail-closed with unified retry
2026-04-02 11:27:59 +08:00
fantiu 7c4932154c fix:auth check test file gofmt. 2026-04-02 11:20:06 +08:00
fantiu e2dbaa7c78 feat(auth): unify MCP retry constant and add retry to remaining endpoints 2026-04-02 11:10:49 +08:00
fantiu 11a0dbc84a feat(auth): unify MCP retry constant and add retry to remaining endpoints 2026-04-02 10:55:51 +08:00
fantiu ee441643dd fix(auth): switch CLI auth check from fail-open to fail-closed
Previously, when the /cli/cliAuthEnabled API was unreachable (network
error, timeout, 5xx, etc.), both OAuth and Device Flow login modes
silently assumed CLI access was enabled (fail-open). This allowed users
to "successfully" log in even when their organization had not granted
CLI data access, leading to confusing failures on subsequent API calls.

Changes:
- Reverse the check logic in OAuth callback: treat any error as
  "not enabled" instead of "enabled", showing the permission request
  page so users can take action.
- Block login in Device Flow when the check fails, with a clear
  error message asking users to verify network connectivity.
- Add retry with backoff (3 attempts, 0s/1s/2s) to
  CheckCLIAuthEnabled to tolerate transient network issues.
- Add retry with backoff (3 attempts) to FetchClientIDFromMCP
  (/cli/clientId) for the same transient-error resilience.
- Add i18n entries (en/zh) for new error messages.
- Add 19 tests covering server error, connection refused, malformed
  JSON, timeout, transient-then-recovery, business error, and full
  Device Flow loginOnce integration scenarios for both endpoints.

Made-with: Cursor
2026-04-02 10:55:51 +08:00
github-actions[bot] 4c5affba99 chore: update coverage badge [skip ci] 2026-04-02 02:44:56 +00:00
fantiu c8148ef2cc Merge pull request #71 from wxianfeng/feature/wk_merge
Feature/wk merge
2026-04-02 10:22:49 +08:00
fantiu b89740bad6 feat: add self-upgrade command with GitHub Releases as data source 2026-04-01 21:26:01 +08:00
xianfeng wang 5c0d2b274c Merge branch 'DingTalk-Real-AI:main' into feature/wk_merge 2026-04-01 21:05:35 +08:00
fantiu 0489cd67c8 Merge pull request #70 from DingTalk-Real-AI/feat-css
style(auth): redesign OAuth authorization pages UI
2026-04-01 20:38:00 +08:00
fantiu 15d495e56e style(auth): redesign OAuth authorization pages UI 2026-04-01 20:21:05 +08:00
coffeeBigSir 28b775198d Merge pull request #69 from DingTalk-Real-AI/version-unify
Version unify
2026-04-01 17:11:40 +08:00
tianlei.qjb e976bd5fc9 - release.yml: add DWS_PACKAGE_VERSION=${{ github.ref_name }} 2026-04-01 17:04:34 +08:00
tianlei.qjb 86355413fc fix: unify version management - pass git tag to post-goreleaser.sh 2026-04-01 17:03:52 +08:00
fantiu c650afa6eb Merge pull request #66 from DingTalk-Real-AI/feat-performance
feat(auth): fail-fast for unauthenticated requests before MCP call
2026-04-01 15:53:28 +08:00
fantiu 14f558facf feat(auth): fail-fast for unauthenticated requests before MCP call
- Add token validation in executeInvocation to reject requests early
- Return clear Chinese error message: '未登录,请先执行 dws auth login'
- Include actionable hint and suggested command in error response
- Avoid wasting network resources on HTTP 400 from MCP gateway

Also includes:
- Add DWS_PERF_TIMING env for CLI execution timing analysis
- Add TimingCollector to track cmd_init, auth_token, mcp_call durations

Test changes:
- Add --token flag to tests that require authentication
- Add TestRuntimeRunnerRejectsUnauthenticatedRequest test case
2026-04-01 15:48:23 +08:00
coffeeBigSir 0257d1f084 Merge pull request #65 from DingTalk-Real-AI/releasescript-yh
refactor: unify version management - Git tag as SSOT
2026-04-01 15:11:23 +08:00
tianlei.qjb b1b4730536 refactor: unify version management - Git tag as SSOT 2026-04-01 15:05:31 +08:00
fantiu f732dcd2ba Merge pull request #64 from DingTalk-Real-AI/feat-auth-improvement
fix(auth): prevent auth code expiration error on OAuth callback page …
2026-04-01 14:54:30 +08:00
fantiu 8406355e7f fix(auth): prevent auth code expiration error on OAuth callback page refresh 2026-04-01 14:48:58 +08:00
coffeeBigSir 3c5f40648e Merge pull request #63 from audanye-sudo/docs/optimize-readme
docs: simplify getting started section in README
2026-04-01 11:50:32 +08:00
audanye-sudo 794e168008 docs: simplify getting started section with inline login commands
- Show both browser and device-flow login commands upfront
- Remove authorization screenshot image for cleaner layout
2026-04-01 11:47:47 +08:00
fantiu 35548e4780 Merge pull request #62 from DingTalk-Real-AI/feat-login-optimization
feat(auth): persist OAuth credentials for reliable token refresh
2026-04-01 11:05:47 +08:00
fantiu 2a056cc5d0 feat(auth): persist OAuth credentials for reliable token refresh 2026-04-01 11:03:08 +08:00
coffeeBigSir 3baadb99ce Merge pull request #60 from audanye-sudo/feat/onboarding-official-app-mode
docs: add Official App onboarding mode to Getting Started
2026-04-01 11:01:17 +08:00
audanye-sudo fcb8b2c782 docs: add Official App mode to Getting Started and update onboarding flow
- Add two authentication modes: Official App (recommended) and Custom App
- Official App mode: direct login without creating an app, admin-controlled access
- Include authorization flow screenshots for both enabled/disabled org scenarios
- Add collapsible admin guide for enabling CLI access and handling member requests
- Move Custom App steps into collapsible sections to reduce visual noise
- Remove whitelist references from IMPORTANT banner
- Update Agent usage section to show both auth modes
2026-04-01 10:59:07 +08:00
xianfeng wang c74f1eeb26 Merge branch 'DingTalk-Real-AI:main' into feature/wk_merge 2026-04-01 09:53:58 +08:00
wxianfeng 933615003c merge upstream main 2026-04-01 09:53:18 +08:00
coffeeBigSir fc49f3dc7a Merge pull request #59 from audanye-sudo/feat/base-experience-upgrade
fix: error diagnostics enhancement and logging reliability
2026-04-01 09:37:21 +08:00
audanye-sudo 8fb0ecb9ba fix: resolve verbosity flag lookup, FileLogger lazy binding, and business error logging
- resolveVerbosity: use cmd.Flags() instead of PersistentFlags() to correctly
  pick up inherited --verbose/--debug flags on subcommands
- FileLogger lazy binding: bind in executeInvocation since configureLogLevel
  initializes it after runner construction
- logBusinessError: log MCP tool errors and business errors (HTTP 200 +
  success=false) to file logger for offline diagnosis
2026-04-01 09:33:09 +08:00
audanye-sudo ce6c32bf53 feat: enhance error diagnostics with trace ID, server error code, and comprehensive logging
Improve CLI error output and local logging to enable offline issue diagnosis:

- Add ServerDiagnostics struct to extract and propagate trace_id, server_error_code,
  technical_detail, and server_retryable from MCP server responses
- Extract diagnostics from JSON-RPC error.data, tool call result content, and HTTP
  response headers (X-Trace-Id, X-Request-Id, x-dingtalk-trace-id)
- Redesign PrintHuman with three verbosity levels (Normal/Verbose/Debug):
  Normal shows trace ID + server code; Verbose adds technical detail;
  Debug adds internal diagnostics (RPC code, operation, reason)
- Add PrintHumanAt for explicit verbosity control, PrintHuman defaults to Normal
- Enhance local logging with request body (sanitized), response body (on error),
  retry attempts, and error classification events
- Add TruncateBody, SanitizeArguments, RedactHeaders logging utilities with
  sensitive key detection via substring matching
- Fix respRetryAfter called before nil guard in doWithRetry retry loop
- Use GetBool instead of string comparison for flag resolution
2026-04-01 09:32:59 +08:00
fantiu 7dbef98dd0 Merge pull request #58 from fantiu/feat-login-upgrade
Feat login upgrade
2026-04-01 00:53:52 +08:00
fantiu 773804ee80 feat(auth): enhance device flow with CLI auth check and admin guidance 2026-04-01 00:47:49 +08:00
fantiu 1df56cc99b Merge branch 'DingTalk-Real-AI:main' into main 2026-03-31 21:06:17 +08:00
coffeeBigSir 0606762c29 Merge pull request #51 from Gyyshen/feature/addskill
feature: add skill management command for downloading and installing skills
2026-03-31 19:29:20 +08:00
meng93 e7677df541 Merge pull request #53 from DingTalk-Real-AI/feat/support-issue-notification
feat: 支持issuie消息通知
2026-03-31 18:58:47 +08:00
meng93 dda4dacb1c feat: to #73551688 支持消息通知 2026-03-31 17:38:08 +08:00
shenyunliang 345156c605 fix: correct import path for errors package in skill_command.go 2026-03-31 16:21:22 +08:00
github-actions[bot] cc4dd1e87b chore: update coverage badge [skip ci] 2026-03-31 07:54:38 +00:00
wxianfeng 3c75c66d4d chore: group pkg imports per goimports
Move pkg/config next to other pkg/* imports after local internal packages.

Made-with: Cursor
2026-03-31 15:32:50 +08:00
wxianfeng 93d6fdb17e feat: edition layer for Wukong overlay; promote shared code to pkg/
Introduce a build-time edition hook so downstream overlays (e.g. Wukong)
can customize auth UX, config dir, static server list, visible products,
and extra root commands while keeping the default open-source behavior.

- Add pkg/edition (defaults) and pkg/editiontest contract tests; Makefile
  target edition-test; CI job edition-tests.
- Root: optional RegisterExtraCommands with tool caller adapter; inject
  static servers to skip market discovery when configured; deduplicate
  top-level commands so overlay wins; extend hideNonDirectRuntimeCommands
  with edition VisibleProducts and static commands (recovery/schema/mcp).
- Auth/config: gate login subcommand and auth login hints for embedded
  editions; optional token auto-purge; edition ConfigDir override.
- version: human-readable multi-line output plus JSON with edition,
  architecture, build, commit.
- Relocate internal/{config,convert,validate} to pkg/; add pkg/cli and
  pkg/cmdutil (flags/time helpers).
- CI: optional notify-downstream job (GitLab trigger via secrets) on main
  push; open_source_policy_test adjusted for new layout.

Secrets (repo settings): WUKONG_TRIGGER_TOKEN, WUKONG_TRIGGER_URL — optional;
if unset, downstream step is skipped.

Made-with: Cursor
2026-03-31 15:29:21 +08:00
github-actions[bot] 75b873d231 chore: update coverage badge [skip ci] 2026-03-31 05:47:53 +00:00
coffeeBigSir a912cbc52b Merge pull request #43 from audanye-sudo/feat/stability-and-hint-optimization
docs: update DingTalk co-creation group QR code
2026-03-31 10:03:38 +08:00
audanye-sudo 5a99b84c25 docs: update DingTalk co-creation group QR code image 2026-03-31 10:00:56 +08:00
audanye-sudoandClaude Opus 4.6 7fda120d5a feat: enhance error diagnostics with trace ID, server error code, and comprehensive logging
Improve CLI error output and local logging to enable offline issue diagnosis:

- Add ServerDiagnostics struct to extract and propagate trace_id, server_error_code,
  technical_detail, and server_retryable from MCP server responses
- Extract diagnostics from JSON-RPC error.data, tool call result content, and HTTP
  response headers (X-Trace-Id, X-Request-Id, x-dingtalk-trace-id)
- Redesign PrintHuman with three verbosity levels (Normal/Verbose/Debug):
  Normal shows trace ID + server code; Verbose adds technical detail;
  Debug adds internal diagnostics (RPC code, operation, reason)
- Add PrintHumanAt for explicit verbosity control, PrintHuman defaults to Normal
- Enhance local logging with request body (sanitized), response body (on error),
  retry attempts, and error classification events
- Add TruncateBody, SanitizeArguments, RedactHeaders logging utilities with
  sensitive key detection via substring matching
- Fix respRetryAfter called before nil guard in doWithRetry retry loop
- Use GetBool instead of string comparison for flag resolution

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-03-31 09:01:41 +08:00
fantiu 3b8233e6ba Merge pull request #40 from fantiu/feat-readme-typo
docs(readme): combine redirect URLs into single line for easier copy-pasteFeat readme typo
2026-03-30 16:08:15 +08:00
fantiu 675ce53c06 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-readme-typo
docs(readme): combine redirect URLs into single line for easier copy-paste
2026-03-30 16:01:32 +08:00
fantiu d51c1ff131 docs(readme): combine redirect URLs into single line for easier copy-paste 2026-03-30 16:01:03 +08:00
fantiu 8b423b97e8 Merge pull request #39 from fantiu/feat-performance
perf(auth): cache resolved credentials to avoid repeated keychain access
2026-03-30 15:22:15 +08:00
fantiu 1a6d129fe3 调整回调地址格式 2026-03-30 15:17:56 +08:00
fantiu a0bf715ddf Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-performance
解决Windows issue bug,修改readme回调地址格式
2026-03-30 15:11:29 +08:00
xianfeng wang 110f887181 Merge branch 'DingTalk-Real-AI:main' into main 2026-03-30 15:10:33 +08:00
github-actions[bot] d2c5a027d4 chore: update coverage badge [skip ci] 2026-03-30 07:09:07 +00:00
fantiu 077665e27b Merge branch 'DingTalk-Real-AI:main' into main 2026-03-30 15:08:10 +08:00
coffeeBigSir 5c41d2b8f4 Merge pull request #29 from DingTalk-Real-AI/xtyh
Xtyh
2026-03-30 14:55:46 +08:00
tianlei.qjb 7d9e592f84 去除message history 2026-03-30 14:49:30 +08:00
tianlei.qjb 11199e6848 add interactive confirmation for destructive dynamic commands 2026-03-30 14:23:50 +08:00
tianlei.qjb 19c38d4b94 add report helper with flexible date parsing and defaults 2026-03-30 13:44:01 +08:00
tianlei.qjb 851cf43180 Merge branch 'main' of https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli into xtyh 2026-03-30 13:32:31 +08:00
github-actions[bot] e004df38c7 chore: update coverage badge [skip ci] 2026-03-30 04:01:24 +00:00
fantiu 3bb504bb1f feat(auth): persist client credentials and optimize keychain access
- Auto-persist client-id/client-secret after successful login (stored in keychain)
- Enhanced logout to clean up app.json, keychain secrets, and token.json
- Fixed OAuth callback race condition (write response before sending code)
- Added credentials cache to avoid repeated keychain access (improves CLI speed)
- Updated README with credential priority docs and bumped to v1.0.5
2026-03-30 11:59:03 +08:00
fantiu 26263f8a17 feat(auth): persist client credentials and optimize keychain access
- Auto-persist client-id/client-secret after successful login (stored in keychain)
- Enhanced logout to clean up app.json, keychain secrets, and token.json
- Fixed OAuth callback race condition (write response before sending code)
- Added credentials cache to avoid repeated keychain access (improves CLI speed)
- Updated README with credential priority docs and bumped to v1.0.4
2026-03-30 11:57:07 +08:00
fantiu fc22f53b92 Merge pull request #33 from fantiu/feat-application
feat(auth): persist client credentials for token refresh
2026-03-30 11:41:56 +08:00
fantiu 654bcc4ecb feat(auth): persist client credentials for token refresh
When logging in with --client-id and --client-secret, credentials are now
automatically persisted to enable automatic token refresh after expiration.
Client secret is securely stored in system Keychain, config file only
contains a reference.
2026-03-30 11:34:11 +08:00
fantiu 2175f2fe59 merge 2026-03-30 11:15:53 +08:00
fantiu 81ca098db1 调整readme 2026-03-30 11:10:15 +08:00
github-actions[bot] 25118d1ec7 chore: update coverage badge [skip ci] 2026-03-30 03:09:54 +00:00
fantiu 3e3c17d686 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-application
feat(auth): persist client credentials for token refresh

When logging in with --client-id and --client-secret, credentials are now
automatically persisted to enable automatic token refresh after expiration.
Client secret is securely stored in system Keychain, config file only
contains a reference.
2026-03-30 10:58:37 +08:00
fantiu 964855373e feat(auth): persist client credentials for token refresh
When logging in with --client-id and --client-secret, credentials are now
automatically persisted to enable automatic token refresh after expiration.
Client secret is securely stored in system Keychain, config file only
contains a reference.
2026-03-30 10:53:26 +08:00
tianlei.qjb 3c83c0cff2 Merge branch 'xtyh' of https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli into xtyh 2026-03-29 23:57:09 +08:00
tianlei.qjb 4b555abefe docs: remove v1.0.1 version tags from feature sections 2026-03-29 23:56:35 +08:00
tianlei.qjb 4742112000 feature modified 2026-03-29 23:49:18 +08:00
tianlei.qjb 7f0567aa39 Merge branch 'main' of https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli into xtyh 2026-03-29 23:48:02 +08:00
coffeeBigSir 8fb1dcda86 Merge pull request #28 from audanye-sudo/feat/onboarding-experience-improvement
docs: improve onboarding experience and skill reference clarity
2026-03-29 23:38:27 +08:00
audanye-sudo 94deafbaa9 docs: update release badge version to v1.0.3 2026-03-29 23:37:31 +08:00
audanye-sudo 2555447c7b docs: enhance README onboarding flow and getting started guide
- Improve installation and getting started sections for new users
- Add clearer step-by-step guidance for first-time setup
- Update examples with more realistic use cases
2026-03-29 23:34:35 +08:00
audanye-sudo 6e91b2d142 docs: improve agent skill references and intent guide clarity
- Enhance product skill docs with clearer examples and descriptions
- Update intent-guide with better routing patterns
- Expand simple.md with comprehensive onboarding examples
- Fix inconsistent command references across skill docs
2026-03-29 23:33:41 +08:00
tianlei.qjb 54145b65ec fix test script 2026-03-29 22:30:03 +08:00
tianlei.qjb 2fad9c95db test 2026-03-29 22:29:15 +08:00
176 changed files with 22678 additions and 618 deletions
+1 -1
View File
@@ -1 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 53.9%"><title>coverage: 53.9%</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">53.9%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">53.9%</text></g></svg>
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 48.8%"><title>coverage: 48.8%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">48.8%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">48.8%</text></g></svg>

Before

Width:  |  Height:  |  Size: 1.1 KiB

After

Width:  |  Height:  |  Size: 1.1 KiB

+38
View File
@@ -141,3 +141,41 @@ jobs:
- name: Generated Drift
run: ./scripts/policy/check-generated-drift.sh
edition-tests:
name: Edition Contract Tests
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- name: Check out repository
uses: actions/checkout@v4
- name: Set up Go
uses: actions/setup-go@v5
with:
go-version-file: go.mod
- name: Run edition contract tests
run: go test -v -count=1 ./pkg/editiontest/...
notify-downstream:
name: Notify Wukong Overlay
needs: [test, policy, edition-tests]
runs-on: ubuntu-latest
if: github.ref == 'refs/heads/main' && github.event_name == 'push'
steps:
- name: Trigger downstream CI
run: |
# Trigger internal GitLab CI pipeline via webhook.
# WUKONG_TRIGGER_TOKEN is a repository secret.
if [ -n "${{ secrets.WUKONG_TRIGGER_TOKEN }}" ]; then
curl --fail --silent --show-error \
-X POST \
-F "token=${{ secrets.WUKONG_TRIGGER_TOKEN }}" \
-F "ref=main" \
-F "variables[UPSTREAM_SHA]=${{ github.sha }}" \
"${{ secrets.WUKONG_TRIGGER_URL }}"
echo "Downstream CI triggered."
else
echo "No WUKONG_TRIGGER_TOKEN configured, skipping downstream notification."
fi
+53
View File
@@ -0,0 +1,53 @@
# Issue 变更推送到 Webhook
# 当有 Issue 变更时,发送指定格式的数据到 webhook
name: 📤 Issue Webhook Notification
on:
issues:
types: [opened, reopened, closed, edited, labeled, unlabeled]
jobs:
notify:
runs-on: ubuntu-latest
steps:
- name: 📬 Send Issue to Webhook
uses: actions/github-script@v7
with:
script: |
const webhook = process.env.ISSUE_WEBHOOK_URL;
if (!webhook) {
console.log('⚠️ ISSUE_WEBHOOK_URL not set, skipping notification');
return;
}
const payload = context.payload;
const issue = payload.issue;
const action = payload.action;
// 构建指定格式的数据
const webhookPayload = {
action: action,
issue: {
id: issue.id,
number: issue.number,
title: issue.title,
body: issue.body,
state: issue.state,
html_url: issue.html_url,
labels: (issue.labels || []).map(label => label.name).join(', ') || '无标签'
}
};
const response = await fetch(webhook, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify(webhookPayload)
});
if (response.ok) {
console.log('✅ Webhook notification sent successfully');
} else {
console.log('❌ Failed to send webhook notification:', response.status, response.statusText);
}
env:
ISSUE_WEBHOOK_URL: ${{ secrets.DINGTALK_AI_TABLE_WEBHOOK }}
+50
View File
@@ -0,0 +1,50 @@
# Issue 自动同步到钉钉群
# 当有新 Issue 时,自动发送到钉钉群(不包括 comment)
name: 🔔 Issue Notification to DingTalk
on:
issues:
types: [opened, reopened, closed, labeled]
jobs:
notify:
runs-on: ubuntu-latest
steps:
- name: 📬 Send Issue to DingTalk
uses: actions/github-script@v7
with:
script: |
const webhook = process.env.DINGTALK_WEBHOOK;
if (!webhook) {
console.log('⚠️ DINGTALK_WEBHOOK not set, skipping notification');
return;
}
const payload = context.payload;
const issue = payload.issue;
const action = payload.action;
// 构建消息标题和内容(确保包含关键字 "issue" 以支持 Custom Keywords 模式)
const title = `[${action.toUpperCase()}] Issue #${issue.number}: ${issue.title}`;
const content = issue.body?.substring(0, 500) || 'No description';
const url = issue.html_url;
const labelsText = (issue.labels || []).map(label => label.name).join(', ') || '无标签';
// 消息内容必须包含关键字(如 "issue")以支持 Custom Keywords 安全模式
const message = {
msgtype: 'markdown',
markdown: {
title: 'GitHub Issue 通知',
text: `## 🔔 GitHub Issue 通知\n\n**${title}**\n\n🏷️ **Labels**: ${labelsText}\n\n${content}${content.length >= 500 ? '...' : ''}\n\n[点击查看详情](${url})\n\n---\n📦 ${context.repo.owner}/${context.repo.repo}\n\n**关键词**: issue`
}
};
await fetch(webhook, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify(message)
});
console.log('✅ DingTalk notification sent');
env:
DINGTALK_WEBHOOK: ${{ secrets.DINGTALK_WEBHOOK }}
+8
View File
@@ -37,6 +37,14 @@ jobs:
- name: Post-release packaging
run: ./scripts/release/post-goreleaser.sh
env:
DWS_PACKAGE_VERSION: ${{ github.ref_name }}
- name: Upload dws-skills.zip to release
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
gh release upload "${{ github.ref_name }}" dist/dws-skills.zip --clobber
- name: Setup Node.js
uses: actions/setup-node@v4
+1
View File
@@ -27,3 +27,4 @@ test/cli_compat/testdata/
credentials*
plans
_docs
dws.zip
+1
View File
@@ -69,3 +69,4 @@ release:
draft: false
prerelease: auto
name_template: "v{{.Version}}"
mode: replace
+4 -1
View File
@@ -1,6 +1,6 @@
GO ?= go
.PHONY: all help build rebuild test lint fmt policy package release publish-homebrew-formula setup-hooks
.PHONY: all help build rebuild test lint fmt policy edition-test package release publish-homebrew-formula setup-hooks
all: setup-hooks fmt lint build test rebuild
@@ -34,6 +34,9 @@ policy:
@./scripts/policy/check-open-source-assets.sh
@./scripts/policy/check-command-surface.sh --strict
edition-test:
$(GO) test -v -count=1 ./pkg/editiontest/...
package:
@./scripts/dev/build-all.sh
@./scripts/release/post-goreleaser.sh
+138 -70
View File
@@ -9,7 +9,7 @@
<p align="center">
<img src="https://img.shields.io/badge/Go-1.25+-green?logo=go&logoColor=white" alt="Go 1.25+">
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/blob/main/LICENSE"><img src="https://img.shields.io/badge/License-Apache_2.0-blue" alt="License Apache-2.0"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/badge/release-v1.1.0-red" alt="v1.1.0"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/github/v/release/DingTalk-Real-AI/dingtalk-workspace-cli?color=red&label=release" alt="Latest Release"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml"><img src="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml/badge.svg" alt="CI"></a>
<img src=".github/badges/coverage.svg" alt="Coverage">
</p>
@@ -19,15 +19,16 @@
</p>
> [!IMPORTANT]
> **Co-creation Phase**: This project accesses DingTalk enterprise data and requires enterprise admin authorization. Please join the DingTalk DWS co-creation group to complete whitelist configuration. See [Getting Started](#getting-started) below.
> **Co-creation Phase**: This project accesses DingTalk enterprise data and requires enterprise admin authorization. Join the DingTalk DWS co-creation group for support and updates. See [Getting Started](#getting-started) below.
>
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i1/O1CN01ZqtgeV1cImFmTZPAH_!!6000000003578-2-tps-398-372.png" alt="DingTalk Group QR Code" width="150"></a>
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i4/O1CN01Rijgk81gKqVSKMzdx_!!6000000004124-2-tps-654-644.png" alt="DingTalk Group QR Code" width="150"></a>
<details>
<summary><strong>Table of Contents</strong></summary>
- [Why dws?](#why-dws)
- [Installation](#installation)
- [Upgrade](#upgrade)
- [Getting Started](#getting-started)
- [Quick Start](#quick-start)
- [Using with Agents](#using-with-agents)
@@ -39,6 +40,7 @@
</details>
---
<h2 id="why-dws">Why dws?</h2>
@@ -64,8 +66,19 @@ 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:
> ```bash
> xattr -d com.apple.quarantine /path/to/dws
> ```
**Build from source**:
```bash
@@ -79,67 +92,95 @@ cp dws ~/.local/bin/ # install to PATH
</details>
## 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
dws upgrade # interactive upgrade to latest version
dws upgrade --check # check for new versions without installing
dws upgrade --list # list all available versions
dws upgrade --version v1.0.7 # upgrade to a specific version
dws upgrade --rollback # rollback to the previous version
dws upgrade -y # skip confirmation prompt
```
<details>
<summary><strong>How it works</strong></summary>
The upgrade process follows a two-phase atomic flow to ensure consistency:
1. **Prepare** — downloads the platform-specific binary and skill packages to a temporary directory, verifies SHA256 checksums, and extracts/validates all files. If any step fails, the upgrade aborts without modifying the existing installation.
2. **Apply** — only after all preparations succeed, the binary is replaced and skill packages are installed to all detected agent directories (`~/.agents/skills/dws`, `~/.claude/skills/dws`, `~/.cursor/skills/dws`, etc.).
A backup of the current version is automatically created before each upgrade. Use `dws upgrade --rollback` to restore the previous version if needed.
| Flag | Description |
|------|-------------|
| `--check` | Check for updates without installing |
| `--list` | List all available versions with changelogs |
| `--version` | Upgrade to a specific version (e.g. `v1.0.7`) |
| `--rollback` | Rollback to the previous backed-up version |
| `--force` | Force reinstall even if already on the latest version |
| `--skip-skills` | Skip skill package update |
| `-y` | Skip confirmation prompt |
</details>
## Getting Started
### Step 1: Create a DingTalk Application
Go to the [Open Platform Console](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app). Under "Internal Enterprise Apps - DingTalk Apps", click **Create App**.
<details>
<summary>View screenshot</summary>
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="Create Application" width="600">
</p>
</details>
### Step 2: Configure Redirect URL
Go to app settings → **Security Settings**. Add the following redirect URLs and save:
```
http://127.0.0.1
https://login.dingtalk.com
```bash
dws auth login # browser opens automatically
dws auth login --device # for headless environments (Docker, SSH, CI)
```
> `http://127.0.0.1` is for local browser login; `https://login.dingtalk.com` is for `--device` device-flow login (Docker containers, remote servers, and other headless environments). We recommend configuring both.
Select your organization and authorize. That's it.
> If your organization hasn't enabled CLI access, you'll be prompted to send an access request to your admin. Once approved, re-run `dws auth login`.
<details>
<summary>View screenshot</summary>
<summary><strong>Organization hasn't enabled CLI access?</strong></summary>
1. After selecting your organization, click "Apply Now" to notify the admin
2. The admin receives a request card and can approve with one click
3. Once approved, re-run `dws auth login`
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="Configure Redirect URL" width="600">
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="Apply for Access" width="600">
</p>
</details>
### Step 3: Publish the Application
Click "App Release - Version Management & Release" to publish and go live.
<details>
<summary>View screenshot</summary>
<summary><strong>Admin: Enable CLI access for your organization</strong></summary>
Go to [Developer Platform](https://open-dev.dingtalk.com) → "CLI Access Management" → Enable.
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="Publish Application" width="600">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI Access Management" width="600">
</p>
</details>
### Step 4: Request Whitelist Access
<details>
<summary><strong>Custom App mode (CI/CD, ISV integration)</strong></summary>
Join the DingTalk DWS co-creation group and provide your **Client ID** and **admin confirmation** to complete whitelist setup.
For enterprise-managed scenarios, create your own DingTalk app:
### Step 5: Authenticate
1. [Open Platform Console](https://open-dev.dingtalk.com/fe/app#/corp/app) → Create App
2. Security Settings → Add redirect URLs: `http://127.0.0.1,https://login.dingtalk.com`
3. Publish the app
4. Login:
```bash
dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
```
Or via environment variables:
Credentials are securely persisted after first login (Keychain). Subsequent runs auto-refresh tokens.
```bash
export DWS_CLIENT_ID=<your-app-key>
export DWS_CLIENT_SECRET=<your-app-secret>
dws auth login
```
> CLI flags take precedence over environment variables. Credentials are used for DingTalk's OAuth device flow.
</details>
## Quick Start
@@ -154,13 +195,6 @@ dws todo task list --dry-run # preview without executing
dws is designed as an AI-native CLI. Complete [Installation](#installation) and [Getting Started](#getting-started) first, then configure your agent:
```bash
# Configure auth via environment variables (recommended for agents, no interactive login)
export DWS_CLIENT_ID=<your-app-key>
export DWS_CLIENT_SECRET=<your-app-secret>
dws auth login
```
### Agent Invocation Patterns
```bash
@@ -183,7 +217,7 @@ Agents don't need pre-built knowledge of every command. Use `dws schema` to dyna
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
# Step 2: Inspect target tool's parameter schema
dws schema aitable.query_records --jq '.tool.input_schema'
dws schema aitable.query_records --jq '.tool.parameters'
# Step 3: Construct the correct call
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
@@ -191,7 +225,7 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
### Agent Skills
The repo ships Agent Skills (`SKILL.md` files) for every DingTalk product. After installing, tools like Claude Code / Cursor can use DingTalk capabilities directly:
The repo ships a complete Agent Skill system (`skills/`). After installing, AI tools like Claude Code / Cursor can operate DingTalk directly through natural language:
```bash
# Install skills into current project
@@ -200,12 +234,45 @@ curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace
> `install.sh` installs to `$HOME/.agents/skills/dws` (global); `install-skills.sh` installs to `./.agents/skills/dws` (current project).
Author your own Agent Skills and orchestrate them with dws skills for cross-product workflows: **ISV Skill → dws Skill → DingTalk Open Platform API (enforced auth + full audit)**.
**What's included:**
| Component | Path | Description |
|-----------|------|-------------|
| Master Skill | `SKILL.md` | Intent routing, decision tree, safety rules, error handling |
| Product references | `references/products/*.md` | Per-product command reference (aitable, chat, calendar, etc.) |
| Intent guide | `references/intent-guide.md` | Disambiguation for confusing scenarios (e.g. report vs todo) |
| Global reference | `references/global-reference.md` | Auth, output formats, global flags |
| Error codes | `references/error-codes.md` | Error codes + debugging workflows |
| Recovery guide | `references/recovery-guide.md` | `RECOVERY_EVENT_ID` handling |
| Ready-made scripts | `scripts/*.py` | 13 batch operation scripts (see below) |
<details>
<summary><strong>Ready-made scripts</strong> — 13 Python scripts for common multi-step workflows</summary>
| Script | Description |
|--------|-------------|
| `calendar_schedule_meeting.py` | Create event + add participants + find & book available meeting room |
| `calendar_free_slot_finder.py` | Find common free slots across multiple people, recommend best meeting time |
| `calendar_today_agenda.py` | View today/tomorrow/this week's schedule |
| `import_records.py` | Batch import records from CSV/JSON into AITable |
| `bulk_add_fields.py` | Batch add fields to an AITable data table |
| `upload_attachment.py` | Upload attachment to AITable attachment field |
| `todo_batch_create.py` | Batch create todos from JSON (with priority, due date, executors) |
| `todo_daily_summary.py` | Summarize today/this week's incomplete todos |
| `todo_overdue_check.py` | Scan overdue todos and output overdue list |
| `contact_dept_members.py` | Search department by name and list all members |
| `attendance_my_record.py` | View my attendance records for today/this week/specific date |
| `attendance_team_shift.py` | Query team shift schedules and attendance statistics |
| `report_inbox_today.py` | View today's received reports with details |
</details>
**ISV Integration**: Author your own Agent Skills and orchestrate them with dws skills for cross-product workflows: **ISV Skill → dws Skill → DingTalk Open Platform API (enforced auth + full audit)**.
## Features
<details>
<summary><strong>Smart Input Correction</strong> — auto-corrects common AI model parameter mistakes <code>v1.0.1</code></summary>
<summary><strong>Smart Input Correction</strong> — auto-corrects common AI model parameter mistakes</summary>
Built-in pipeline engine that normalizes flag names, splits sticky arguments, and fuzzy-matches typos:
@@ -234,7 +301,7 @@ dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-i
</details>
<details>
<summary><strong>jq Filtering & Field Selection</strong> — fine-grained output control to reduce token consumption <code>v1.0.1</code></summary>
<summary><strong>jq Filtering & Field Selection</strong> — fine-grained output control to reduce token consumption</summary>
```bash
# Built-in jq expressions
@@ -248,19 +315,19 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocati
</details>
<details>
<summary><strong>Schema Introspection</strong> — query parameter schemas before making calls <code>v1.0.1</code></summary>
<summary><strong>Schema Introspection</strong> — query parameter schemas before making calls</summary>
```bash
dws schema # list all products and tools
dws schema aitable.query_records # view parameter schema
dws schema aitable.query_records --jq '.tool.input_schema.required' # view required fields
dws schema aitable.query_records --jq '.tool.required' # view required fields
dws schema --jq '.products[].id' # extract all product IDs
```
</details>
<details>
<summary><strong>Pipe & File Input</strong> — read flag values from files or stdin <code>v1.0.1</code></summary>
<summary><strong>Pipe & File Input</strong> — read flag values from files or stdin</summary>
```bash
# Read message body from a file
@@ -280,21 +347,22 @@ dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
## Key Services
| Service | Command | Description |
|---------|---------|-------------|
| Contact | `contact` | Users / departments |
| Chat | `chat` | Group management / members / bot messaging / webhook |
| Calendar | `calendar` | Events / meeting rooms / free-busy |
| Todo | `todo` | Task management |
| Approval | `oa` | Processes / forms / instances |
| Attendance | `attendance` | Clock-in / shifts / statistics |
| Ding | `ding` | DING messages / send / recall |
| Report | `report` | Reports / templates / statistics |
| AITable | `aitable` | AI table operations |
| Workbench | `workbench` | App query |
| DevDoc | `devdoc` | Open platform docs search |
| Service | Command | Commands | Subcommands | Description |
|---------|---------|:--------:|-------------|-------------|
| Contact | `contact` | 6 | `user` `dept` | Search users by name/mobile, batch query, departments, current user profile |
| Chat | `chat` | 10 | `message` `group` `search` | Group CRUD, member management, bot messaging, webhook |
| Bot | `chat bot` | 6 | `bot` `group` `message` `search` | Robot creation/search, group/single messaging, webhook, message recall |
| Calendar | `calendar` | 13 | `event` `room` `participant` `busy` | Events CRUD, meeting room booking, free-busy query, participant management |
| Todo | `todo` | 6 | `task` | Create, list, update, done, get detail, delete |
| Approval | `oa` | 9 | `approval` | Approve/reject/revoke, pending tasks, initiated instances, process list |
| Attendance | `attendance` | 4 | `record` `shift` `summary` `rules` | Clock-in records, shift schedules, attendance summary, group rules |
| Ding | `ding` | 2 | `message` | Send/recall DING messages |
| Report | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | Create reports, sent/received list, templates, statistics |
| AITable | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | Full CRUD for bases/tables/records/fields, templates |
| Workbench | `workbench` | 2 | `app` | Batch query app details |
| DevDoc | `devdoc` | 1 | `article` | Search platform docs and error codes |
Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
> 86 commands across 12 products. Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
<details>
<summary>Coming soon</summary>
+138 -70
View File
@@ -9,7 +9,7 @@
<p align="center">
<img src="https://img.shields.io/badge/Go-1.25+-green?logo=go&logoColor=white" alt="Go 1.25+">
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/blob/main/LICENSE"><img src="https://img.shields.io/badge/License-Apache_2.0-blue" alt="License Apache-2.0"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/badge/release-v1.1.0-red" alt="v1.1.0"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/github/v/release/DingTalk-Real-AI/dingtalk-workspace-cli?color=red&label=release" alt="Latest Release"></a>
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml"><img src="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml/badge.svg" alt="CI"></a>
<img src=".github/badges/coverage.svg" alt="Coverage">
</p>
@@ -19,15 +19,16 @@
</p>
> [!IMPORTANT]
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。当前为灰度共创阶段,请加入钉钉 DWS 共创群完成白名单配置。详见下方 [开始使用](#开始使用)。
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。欢迎加入钉钉 DWS 共创群获取支持与最新动态。详见下方 [开始使用](#开始使用)。
>
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i1/O1CN01ZqtgeV1cImFmTZPAH_!!6000000003578-2-tps-398-372.png" alt="DingTalk Group QR Code" width="150"></a>
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i4/O1CN01Rijgk81gKqVSKMzdx_!!6000000004124-2-tps-654-644.png" alt="DingTalk Group QR Code" width="150"></a>
<details>
<summary><strong>目录</strong></summary>
- [为什么选择 dws?](#why-dws)
- [安装](#安装)
- [升级](#升级)
- [开始使用](#开始使用)
- [快速开始](#快速开始)
- [在 Agent 中使用](#在-agent-中使用)
@@ -39,6 +40,7 @@
</details>
---
<h2 id="why-dws">为什么选择 dws?</h2>
@@ -64,8 +66,19 @@ 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 无法检查其是否包含恶意软件”,请执行:
> ```bash
> xattr -d com.apple.quarantine /path/to/dws
> ```
**从源码构建**:
```bash
@@ -79,67 +92,95 @@ cp dws ~/.local/bin/ # 安装到 PATH
</details>
## 升级
> 需要 **v1.0.7** 及以上版本。更早版本请重新执行[安装脚本](#安装)进行升级。
dws 内置自升级能力,直接从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 拉取更新,支持 SHA256 完整性校验和自动备份。
```bash
dws upgrade # 交互式升级到最新版本
dws upgrade --check # 仅检查是否有新版本
dws upgrade --list # 列出所有可用版本
dws upgrade --version v1.0.7 # 升级到指定版本
dws upgrade --rollback # 回滚到上一版本
dws upgrade -y # 跳过确认直接升级
```
<details>
<summary><strong>工作原理</strong></summary>
升级过程采用两阶段原子流程,确保一致性:
1. **准备阶段** — 将平台对应的二进制文件和技能包下载到临时目录,校验 SHA256 校验和,解压并验证所有文件。任何步骤失败则立即中止,不会修改现有安装。
2. **执行阶段** — 仅在所有准备工作成功后,替换二进制文件并将技能包安装到所有已检测到的 Agent 目录(`~/.agents/skills/dws`、`~/.claude/skills/dws`、`~/.cursor/skills/dws` 等)。
每次升级前自动备份当前版本,可通过 `dws upgrade --rollback` 随时回滚。
| Flag | 说明 |
|------|------|
| `--check` | 仅检查更新,不安装 |
| `--list` | 列出所有可用版本及更新日志 |
| `--version` | 升级到指定版本(如 `v1.0.7`) |
| `--rollback` | 回滚到上一个备份版本 |
| `--force` | 强制重新安装,即使已是最新版本 |
| `--skip-skills` | 跳过技能包更新 |
| `-y` | 跳过确认提示 |
</details>
## 开始使用
### 步骤 1:创建钉钉应用
进入 [开放平台应用开发后台](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app),在「企业内部应用 - 钉钉应用」点击**创建应用**。
<details>
<summary>查看截图</summary>
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="创建应用" width="600">
</p>
</details>
### 步骤 2:配置重定向 URL
进入应用 → **安全设置**,在「重定向 URL」中添加以下地址并保存:
```
http://127.0.0.1
https://login.dingtalk.com
```bash
dws auth login # 自动唤起浏览器
dws auth login --device # 无浏览器环境(Docker、SSH、CI)
```
> `http://127.0.0.1` 用于本地浏览器登录;`https://login.dingtalk.com` 用于 `--device` 设备流登录(Docker 容器、远程服务器等无浏览器环境)。建议两个都配置。
选择组织并授权即可。
> 如果组织尚未开启 CLI 访问权限,系统会引导你向管理员发送申请。审批通过后重新执行 `dws auth login` 即可。
<details>
<summary>查看截图</summary>
<summary><strong>组织未开启 CLI 访问权限?</strong></summary>
1. 选择组织后,点击「立即申请」通知管理员
2. 管理员收到申请卡片,一键审批
3. 审批通过后,重新执行 `dws auth login`
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="配置重定向URL" width="600">
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="申请权限" width="600">
</p>
</details>
### 步骤 3:发布应用
点击「应用发布 - 版本管理与发布」,发布版本使应用上线。
<details>
<summary>查看截图</summary>
<summary><strong>管理员:为组织开启 CLI 访问权限</strong></summary>
进入 [开发者平台](https://open-dev.dingtalk.com) →「CLI 访问管理」→ 开启。
<p align="center">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="发布应用" width="600">
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI访问管理" width="600">
</p>
</details>
### 步骤 4:申请白名单
<details>
<summary><strong>自建应用模式(CI/CD、ISV 集成)</strong></summary>
加入钉钉 DWS 共创群,提供 **Client ID** 和**管理员确认凭证**完成白名单配置。
企业自主管控场景,可创建自有钉钉应用:
### 步骤 5:登录认证
1. [开放平台应用开发后台](https://open-dev.dingtalk.com/fe/app#/corp/app) → 创建应用
2. 安全设置 → 添加重定向 URL:`http://127.0.0.1,https://login.dingtalk.com`
3. 发布应用
4. 登录:
```bash
dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
```
或通过环境变量:
首次登录后凭证安全存储(Keychain),后续自动刷新 Token。
```bash
export DWS_CLIENT_ID=<your-app-key>
export DWS_CLIENT_SECRET=<your-app-secret>
dws auth login
```
> CLI 参数优先于环境变量。凭证用于钉钉 OAuth 设备流认证。
</details>
## 快速开始
@@ -154,13 +195,6 @@ dws todo task list --dry-run # 预览操作但不执行
dws 是为 AI Agent 设计的 CLI 工具。请先完成[安装](#安装)和[开始使用](#开始使用),然后配置 Agent 环境:
```bash
# 通过环境变量配置认证(Agent 推荐方式,无需交互式登录)
export DWS_CLIENT_ID=<your-app-key>
export DWS_CLIENT_SECRET=<your-app-secret>
dws auth login
```
### Agent 调用模式
```bash
@@ -183,7 +217,7 @@ Agent 无需预置所有命令知识,通过 `dws schema` 动态发现可用能
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
# 第二步:查看目标工具的参数结构
dws schema aitable.query_records --jq '.tool.input_schema'
dws schema aitable.query_records --jq '.tool.parameters'
# 第三步:构造正确的调用
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
@@ -191,7 +225,7 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
### Agent Skills
仓库为每个钉钉产品提供 Agent Skill(`SKILL.md`),安装后 Claude Code / Cursor 等工具可直接使用钉钉能力:
仓库内置完整的 Agent Skill 体系(`skills/`),安装后 Claude Code / Cursor 等 AI 工具可通过自然语言直接操作钉钉:
```bash
# 安装 skills 到当前项目
@@ -200,12 +234,45 @@ curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace
> `install.sh` 安装到 `$HOME/.agents/skills/dws`(全局);`install-skills.sh` 安装到 `./.agents/skills/dws`(当前项目)。
编写您自己的 Agent Skill,与 dws 内置 Skill 搭配构建跨产品工作流:**ISV Skill → dws Skill → 钉钉开放平台 API(强制鉴权 + 全链路审计)**。
**包含内容:**
| 组件 | 路径 | 说明 |
|------|------|------|
| 主 Skill | `SKILL.md` | 意图路由、决策树、安全规则、错误处理 |
| 产品参考 | `references/products/*.md` | 各产品命令详细参考(aitable、chat、calendar 等) |
| 意图指南 | `references/intent-guide.md` | 易混淆场景消歧(如 report vs todo) |
| 全局参考 | `references/global-reference.md` | 认证、输出格式、全局 flag |
| 错误码 | `references/error-codes.md` | 错误码 + 调试流程 |
| Recovery 指南 | `references/recovery-guide.md` | `RECOVERY_EVENT_ID` 处理 |
| 现成脚本 | `scripts/*.py` | 13 个批量操作脚本(见下方) |
<details>
<summary><strong>现成脚本</strong> — 13 个 Python 脚本,覆盖常见多步工作流</summary>
| 脚本 | 说明 |
|------|------|
| `calendar_schedule_meeting.py` | 一键创建日程 + 添加参与者 + 搜索并预定空闲会议室 |
| `calendar_free_slot_finder.py` | 查询多人共同空闲时段,推荐最佳会议时间 |
| `calendar_today_agenda.py` | 查看今天/明天/本周的日程安排 |
| `import_records.py` | 从 CSV/JSON 批量导入记录到 AI 表格 |
| `bulk_add_fields.py` | 批量添加字段到 AI 表格数据表 |
| `upload_attachment.py` | 上传附件到 AI 表格 attachment 字段 |
| `todo_batch_create.py` | 从 JSON 文件批量创建待办(含优先级、截止时间、执行者) |
| `todo_daily_summary.py` | 汇总今天/本周未完成的待办 |
| `todo_overdue_check.py` | 扫描已过截止时间但未完成的待办,输出逾期清单 |
| `contact_dept_members.py` | 按部门名称搜索并列出所有成员 |
| `attendance_my_record.py` | 查看我今天/本周/指定日期的考勤记录 |
| `attendance_team_shift.py` | 查询团队成员本周排班和出勤统计 |
| `report_inbox_today.py` | 查看今天收到的日志列表及详情 |
</details>
**ISV 集成**:编写您自己的 Agent Skill,与 dws 内置 Skill 搭配构建跨产品工作流:**ISV Skill → dws Skill → 钉钉开放平台 API(强制鉴权 + 全链路审计)**。
## 功能特性
<details>
<summary><strong>智能输入纠错</strong> — 自动修正 AI 模型常见的参数错误 <code>v1.0.1</code></summary>
<summary><strong>智能输入纠错</strong> — 自动修正 AI 模型常见的参数错误</summary>
内置 Pipeline 纠错引擎,支持命名风格转换、粘连参数拆分、拼写模糊匹配:
@@ -234,7 +301,7 @@ dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-i
</details>
<details>
<summary><strong>jq 过滤 & 字段筛选</strong> — 精确控制输出,减少 token 消耗 <code>v1.0.1</code></summary>
<summary><strong>jq 过滤 & 字段筛选</strong> — 精确控制输出,减少 token 消耗</summary>
```bash
# 内置 jq 表达式
@@ -248,19 +315,19 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocati
</details>
<details>
<summary><strong>Schema 自省</strong> — 调用前查询任意工具的参数结构 <code>v1.0.1</code></summary>
<summary><strong>Schema 自省</strong> — 调用前查询任意工具的参数结构</summary>
```bash
dws schema # 列出所有产品和工具
dws schema aitable.query_records # 查看参数 Schema
dws schema aitable.query_records --jq '.tool.input_schema.required' # 查看必填字段
dws schema aitable.query_records --jq '.tool.required' # 查看必填字段
dws schema --jq '.products[].id' # 提取所有产品 ID
```
</details>
<details>
<summary><strong>管道 & 文件输入</strong> — 从文件或 stdin 读取 flag 值 <code>v1.0.1</code></summary>
<summary><strong>管道 & 文件输入</strong> — 从文件或 stdin 读取 flag 值</summary>
```bash
# 从文件读取消息内容
@@ -280,21 +347,22 @@ dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
## 核心服务
| 服务 | 命令 | 描述 |
|---------|---------|-------------|
| 通讯录 | `contact` | 用户 / 部门 |
| 群聊 | `chat` | 群管理 / 群成员 / 机器人消息 / Webhook |
| 日历 | `calendar` | 日程 / 会议室 / 闲忙 |
| 待办 | `todo` | 任务管理 |
| 审批 | `oa` | 流程 / 表单 / 实例 |
| 考勤 | `attendance` | 打卡 / 排班 / 统计 |
| DING | `ding` | DING 消息 / 发送 / 撤回 |
| 日志 | `report` | 日志 / 模版 / 统计 |
| 智能表格 | `aitable` | AI 表格操作 |
| 工作台 | `workbench` | 应用查询 |
| 开发者文档 | `devdoc` | 开放平台文档搜索 |
| 服务 | 命令 | 命令数 | 子命令 | 描述 |
|------|------|:------:|--------|------|
| 通讯录 | `contact` | 6 | `user` `dept` | 按姓名/手机号搜索、批量查询、部门树、当前用户信息 |
| 群聊 | `chat` | 10 | `message` `group` `search` | 群增删改查、成员管理、机器人消息、Webhook |
| 机器人 | `chat bot` | 6 | `bot` `group` `message` `search` | 机器人创建/搜索、群聊/单聊消息、Webhook、消息撤回 |
| 日历 | `calendar` | 13 | `event` `room` `participant` `busy` | 日程增删改查、会议室预订、闲忙查询、参与者管理 |
| 待办 | `todo` | 6 | `task` | 创建、列表、修改、完成、详情、删除 |
| 审批 | `oa` | 9 | `approval` | 同意/拒绝/撤销、待我审批、我发起的、流程列表 |
| 考勤 | `attendance` | 4 | `record` `shift` `summary` `rules` | 打卡记录、排班查询、考勤摘要、考勤组规则 |
| DING | `ding` | 2 | `message` | 发送/撤回 DING 消息 |
| 日志 | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | 创建日志、收发列表、模版、统计 |
| 智能表格 | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | 多维表/数据表/记录/字段全量 CRUD、模板 |
| 工作台 | `workbench` | 2 | `app` | 批量查询应用详情 |
| 开发者文档 | `devdoc` | 1 | `article` | 搜索开放平台文档与错误码 |
运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看子命令。
> 12 个产品,86 个命令。运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看子命令。
<details>
<summary>即将推出</summary>
+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 = {
+1 -1
View File
@@ -5,6 +5,7 @@ go 1.25.8
require (
github.com/fatih/color v1.18.0
github.com/google/uuid v1.6.0
github.com/itchyny/gojq v0.12.18
github.com/spf13/cobra v1.10.2
github.com/zalando/go-keyring v0.2.8
golang.org/x/crypto v0.49.0
@@ -15,7 +16,6 @@ require (
require (
github.com/danieljoos/wincred v1.2.3 // indirect
github.com/godbus/dbus/v5 v5.2.2 // indirect
github.com/itchyny/gojq v0.12.18 // indirect
github.com/itchyny/timefmt-go v0.1.7 // indirect
github.com/mattn/go-colorable v0.1.13 // indirect
github.com/mattn/go-isatty v0.0.20 // indirect
+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")
}
}
+33 -5
View File
@@ -24,8 +24,9 @@ import (
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
@@ -48,7 +49,9 @@ func buildAuthCommand() *cobra.Command {
},
}
cmd.AddCommand(newAuthLoginCommand())
if !edition.Get().HideAuthLogin {
cmd.AddCommand(newAuthLoginCommand())
}
cmd.AddCommand(
newAuthLogoutCommand(),
newAuthStatusCommand(),
@@ -118,6 +121,7 @@ func newAuthLoginCommand() *cobra.Command {
}
}
ResetRuntimeTokenCache()
clearCompatCache()
w := cmd.OutOrStdout()
@@ -185,15 +189,31 @@ func newAuthLogoutCommand() *cobra.Command {
defer cancel()
_ = authpkg.RevokeTokenRemote(revokeCtx)
// Load token data to get associated clientId before deletion
var storedClientID string
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
storedClientID = tokenData.ClientID
}
if err := authpkg.DeleteTokenData(configDir); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to clear token data: %v", err))
}
// Clean up associated client secret from keychain
if storedClientID != "" {
_ = authpkg.DeleteClientSecret(storedClientID)
}
// Clean up app credentials (app.json + keychain secret)
_ = authpkg.DeleteAppConfig(configDir)
_ = 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] 已清除所有认证信息")
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
if !edition.Get().IsEmbedded {
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
}
return nil
},
}
@@ -223,6 +243,8 @@ func newAuthStatusCommand() *cobra.Command {
tokenData = updatedData
refreshed = true
}
} else if edition.Get().AutoPurgeToken {
_ = authpkg.DeleteTokenData(configDir)
}
}
if authStatusAuthenticated(tokenData) {
@@ -250,7 +272,9 @@ func newAuthStatusCommand() *cobra.Command {
}
} else {
fmt.Fprintf(w, "%-16s%s\n", "状态:", "未登录")
fmt.Fprintln(w, "运行 dws auth login 进行登录")
if !edition.Get().IsEmbedded {
fmt.Fprintln(w, "运行 dws auth login 进行登录")
}
}
return nil
},
@@ -286,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()
@@ -325,10 +350,13 @@ 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] 认证信息已重置")
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
if !edition.Get().IsEmbedded {
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
}
return nil
},
}
+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")
}
}
+16 -1
View File
@@ -16,8 +16,21 @@ package app
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"
@@ -28,7 +41,9 @@ func defaultConfigDir() string {
if envDir := os.Getenv("DWS_CONFIG_DIR"); envDir != "" {
return envDir
}
if fn := edition.Get().ConfigDir; fn != nil {
return fn()
}
homeDir, err := os.UserHomeDir()
if err != nil {
return exeRelativeConfigDir()
+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")
}
}
+3 -3
View File
@@ -216,10 +216,10 @@ func TestRootHelpCustomizationDoesNotAffectSubcommandHelp(t *testing.T) {
}
}
func TestRootCommandDoesNotRegisterUpgradeCommand(t *testing.T) {
func TestRootCommandRegistersUpgradeCommand(t *testing.T) {
root := NewRootCommand()
if cmd := lookupCommand(root, "upgrade"); cmd != nil {
t.Fatalf("findCommand(upgrade) = %q, want nil", cmd.CommandPath())
if cmd := lookupCommand(root, "upgrade"); cmd == nil {
t.Fatal("upgrade command should be registered on root, but was not found")
}
}
+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()
}
+43 -7
View File
@@ -29,16 +29,25 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/compat"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/helpers"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cobra.Command {
if fn := edition.Get().StaticServers; fn != nil {
injectStaticServers(fn())
// Static servers provided by the edition hook — skip Market discovery
// entirely. The overlay registers its own product commands via
// RegisterExtraCommands; we only add the open-source helpers here.
commands := helpers.NewPublicCommands(runner)
return mergeTopLevelCommands(commands)
}
var commands []*cobra.Command
// Generate commands dynamically from the market discovery API.
if dynamicCmds := loadDynamicCommands(ctx, runner); len(dynamicCmds) > 0 {
commands = append(commands, dynamicCmds...)
}
@@ -46,6 +55,26 @@ func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cob
return mergeTopLevelCommands(commands)
}
// injectStaticServers converts edition.ServerInfo entries into
// market.ServerDescriptor and feeds them into SetDynamicServers so the
// direct-runtime endpoint resolver can find them.
func injectStaticServers(servers []edition.ServerInfo) {
descriptors := make([]market.ServerDescriptor, 0, len(servers))
for _, s := range servers {
descriptors = append(descriptors, market.ServerDescriptor{
Key: s.ID,
DisplayName: s.Name,
Endpoint: s.Endpoint,
CLI: market.CLIOverlay{
ID: s.ID,
Command: s.ID,
Prefixes: s.Prefixes,
},
})
}
SetDynamicServers(descriptors)
}
// loadDynamicCommands loads the server registry and generates CLI commands
// dynamically from CLIOverlay metadata. It consults the disk cache first.
// Within the short revalidation window it uses the cached registry directly;
@@ -59,7 +88,6 @@ func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cob
// 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 {
store := cacheStoreFromEnv()
partition := config.DefaultPartition
@@ -70,13 +98,15 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
useCache := strings.TrimSpace(os.Getenv(cli.CatalogFixtureEnv)) == ""
// --- Cache-first server registry ---
cacheLoadStart := time.Now()
snapshot, freshness, cacheErr := store.LoadRegistry(partition)
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 {
slog.Debug("loadDynamicCommands: using cached registry", "servers", len(snapshot.Servers), "freshness", freshness)
servers = snapshot.Servers
// Only trigger async revalidation in production (no URL override).
// Tests set discoveryBaseURLOverride and control cache expiry directly,
@@ -92,9 +122,10 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
if discoveryBaseURLOverride != "" {
baseURL = discoveryBaseURLOverride
}
slog.Debug("loadDynamicCommands: fetching servers from market API", "base_url", baseURL)
fetchStart := time.Now()
client := market.NewClient(baseURL, ipv4OnlyHTTPClient())
resp, fetchErr := client.FetchServers(ctx, config.DefaultFetchServersLimit)
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).
@@ -106,12 +137,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
}
} else {
servers = market.NormalizeServers(resp, "market")
slog.Debug("loadDynamicCommands: normalized servers", "count", 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)
}
RecordTiming(ctx, "cache_save", time.Since(saveStart))
}
}
}
@@ -122,9 +154,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
// Inject dynamic server data for endpoint resolution
SetDynamicServers(servers)
detailStart := time.Now()
detailsByID := loadCachedDetailsFast(store, servers)
RecordTiming(ctx, "tool_metadata", time.Since(detailStart))
buildStart := time.Now()
cmds := compat.BuildDynamicCommands(servers, runner, detailsByID)
slog.Debug("loadDynamicCommands: built dynamic commands", "commands", len(cmds))
RecordTiming(ctx, "build_commands", time.Since(buildStart))
return cmds
}
+669
View File
@@ -0,0 +1,669 @@
// 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 {
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"
}
+1 -1
View File
@@ -261,7 +261,7 @@ func TestExecuteWritesRecoveryEventIDToStderrOnCapturedFailure(t *testing.T) {
oldArgs := os.Args
defer func() { os.Args = oldArgs }()
os.Args = []string{"dws", "mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`}
os.Args = []string{"dws", "mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"}
stdoutR, stdoutW, err := os.Pipe()
if err != nil {
+650 -20
View File
@@ -15,21 +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/config"
"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"
@@ -39,8 +43,11 @@ 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"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
"github.com/spf13/pflag"
)
@@ -51,12 +58,23 @@ const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
// Execute runs the root command and returns the process exit code.
func Execute() int {
timing := NewTimingCollector()
defer func() {
timing.PrintIfEnabled()
timing.WriteReportIfEnabled(RawVersion(), SanitizeCommand(os.Args))
}()
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer cancel()
// Attach timing collector to context for use by child components
ctx = WithTimingCollector(ctx, timing)
initStart := time.Now()
recovery.ResetRuntimeState()
engine := newPipelineEngine()
root := NewRootCommandWithEngine(ctx, engine)
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
@@ -114,10 +132,29 @@ func flagErrorWithSuggestions(cmd *cobra.Command, err error) error {
}
func printExecutionError(root *cobra.Command, stdout, stderr io.Writer, err error) error {
var raw apperrors.RawStderrError
if stderrors.As(err, &raw) {
_, writeErr := fmt.Fprintln(stderr, raw.RawStderr())
return writeErr
}
if wantsJSONErrors(root) {
return apperrors.PrintJSON(stdout, err)
}
return apperrors.PrintHuman(stderr, err)
return apperrors.PrintHumanAt(stderr, err, resolveVerbosity(root))
}
// resolveVerbosity derives the error verbosity level from the root command's flags.
func resolveVerbosity(cmd *cobra.Command) apperrors.Verbosity {
if cmd == nil {
return apperrors.VerbosityNormal
}
if debug, err := cmd.Flags().GetBool("debug"); err == nil && debug {
return apperrors.VerbosityDebug
}
if verbose, err := cmd.Flags().GetBool("verbose"); err == nil && verbose {
return apperrors.VerbosityVerbose
}
return apperrors.VerbosityNormal
}
func wantsJSONErrors(root *cobra.Command) bool {
@@ -192,6 +229,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
AuthTokenFunc: func(ctx context.Context) string {
return resolveRuntimeAuthToken(ctx, "")
},
LoggerFunc: FileLoggerInstance,
}
runner := newCommandRunnerWithFlags(loader, flags)
@@ -218,7 +256,13 @@ 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 {
CloseFileLogger()
@@ -236,18 +280,38 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
utilityCommands := []*cobra.Command{
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)
deduplicateCommands(root)
}
hideNonDirectRuntimeCommands(root)
configureRootHelp(root)
// Set custom flag error handler for better UX
@@ -261,6 +325,10 @@ func newAuthCommand() *cobra.Command {
return buildAuthCommand()
}
func newSkillCommand() *cobra.Command {
return buildSkillCommand()
}
func newCacheCommand() *cobra.Command {
cacheCmd := newPlaceholderParent("cache", "缓存管理")
@@ -422,24 +490,51 @@ func newVersionCommand() *cobra.Command {
Example: " dws version\n dws version --format json",
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
format, err := cmd.Flags().GetString("format")
if err != nil {
return apperrors.NewInternal("failed to read format flag")
wantJSON := cmd.Flags().Changed("format")
if wantJSON {
format, _ := cmd.Flags().GetString("format")
wantJSON = (format == "json")
}
payload := map[string]any{
"version": Version(),
"go": "1.24+",
editionName := edition.Get().Name
if editionName == "" {
editionName = "open"
}
if format == "json" {
ver := RawVersion()
bt := BuildTime()
gc := GitCommit()
goVer := "1.24+"
arch := "MCP Dynamic Aggregation"
if wantJSON {
payload := map[string]any{
"version": ver,
"edition": editionName,
"architecture": arch,
"go": goVer,
}
if bt != "unknown" {
payload["build"] = bt
}
if gc != "unknown" {
payload["commit"] = gc
}
return output.WriteJSON(cmd.OutOrStdout(), payload)
}
_, err = fmt.Fprintf(
cmd.OutOrStdout(),
"版本: %s\nGo: %s\n",
Version(),
"1.24+",
)
return err
w := cmd.OutOrStdout()
fmt.Fprintf(w, "%-16s%s\n", "Version:", ver)
fmt.Fprintf(w, "%-16s%s\n", "Edition:", editionName)
if bt != "unknown" {
fmt.Fprintf(w, "%-16s%s\n", "Build:", bt)
}
if gc != "unknown" {
fmt.Fprintf(w, "%-16s%s\n", "Commit:", gc)
}
fmt.Fprintf(w, "%-16s%s\n", "Architecture:", arch)
fmt.Fprintf(w, "%-16s%s\n", "Go:", goVer)
return nil
},
}
}
@@ -536,17 +631,34 @@ func newMCPCommand(ctx context.Context, loader cli.CatalogLoader, runner executo
}
// hideNonDirectRuntimeCommands marks top-level product commands as hidden
// unless they correspond to a product discovered via dynamic server discovery.
// unless they correspond to a product discovered via dynamic server discovery
// or listed in the edition's VisibleProducts hook.
// Public utility commands (auth, cache, completion, version) are always kept
// visible; explicitly hidden commands stay hidden.
func hideNonDirectRuntimeCommands(root *cobra.Command) {
allowedProducts := DirectRuntimeProductIDs()
var allowedProducts map[string]bool
if fn := edition.Get().VisibleProducts; fn != nil {
products := fn()
allowedProducts = make(map[string]bool, len(products))
for _, p := range products {
allowedProducts[p] = true
}
} else {
allowedProducts = DirectRuntimeProductIDs()
}
staticCommands := map[string]bool{
"auth": true,
"cache": true,
"config": true,
"doctor": true,
"completion": true,
"skill": true,
"plugin": true,
"version": true,
"help": true,
"recovery": true,
"schema": true,
"mcp": true,
}
for _, cmd := range root.Commands() {
name := cmd.Name()
@@ -563,6 +675,84 @@ 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.
func deduplicateCommands(root *cobra.Command) {
seen := make(map[string]*cobra.Command)
var dups []*cobra.Command
for _, cmd := range root.Commands() {
name := cmd.Name()
if prev, ok := seen[name]; ok {
dups = append(dups, prev)
}
seen[name] = cmd
}
for _, dup := range dups {
root.RemoveCommand(dup)
}
}
func cacheStoreFromEnv() *cache.Store {
cacheDir := strings.TrimSpace(os.Getenv(cli.CacheDirEnv))
return cache.NewStore(cacheDir)
@@ -824,11 +1014,445 @@ 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())
accessToken, tokenErr := loadSkillAccessToken()
if tokenErr == nil && 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() {
// 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)
RegisterStdioClient(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
@@ -839,6 +1463,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
}
+70 -4
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()
@@ -172,8 +180,8 @@ func TestVersionCommandDoesNotRequirePINOrLogin(t *testing.T) {
if err := root.Execute(); err != nil {
t.Fatalf("Execute(version) error = %v", err)
}
if !strings.Contains(out.String(), "\"version\"") {
t.Fatalf("version output missing version key:\n%s", out.String())
if !strings.Contains(out.String(), "Version:") {
t.Fatalf("version output missing Version line:\n%s", out.String())
}
}
@@ -213,8 +221,8 @@ func TestVersionCommandUsesCachedRegistryWithoutBlockingAgedDiscovery(t *testing
if elapsed := time.Since(start); elapsed >= 200*time.Millisecond {
t.Fatalf("Execute(version) took %v, want cached startup under 200ms", elapsed)
}
if !strings.Contains(out.String(), "\"version\"") {
t.Fatalf("version output missing version key:\n%s", out.String())
if !strings.Contains(out.String(), "Version:") {
t.Fatalf("version output missing Version line:\n%s", out.String())
}
}
@@ -255,6 +263,11 @@ func TestRootHelpDoesNotRequirePINOrLogin(t *testing.T) {
if !strings.Contains(out.String(), "Discovered MCP Services:") {
t.Fatalf("root help output missing MCP summary:\n%s", out.String())
}
for _, want := range []string{"Utility Commands:", "skill", "auth", "version"} {
if !strings.Contains(out.String(), want) {
t.Fatalf("root help output missing %q:\n%s", want, out.String())
}
}
}
func TestRootShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
@@ -342,3 +355,56 @@ 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())
}
}
+52 -2
View File
@@ -5,6 +5,7 @@ import (
"strings"
"text/tabwriter"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
@@ -25,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 {
@@ -44,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 {
@@ -53,7 +68,16 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
return nil
}
allowed := DirectRuntimeProductIDs()
var allowed map[string]bool
if fn := edition.Get().VisibleProducts; fn != nil {
products := fn()
allowed = make(map[string]bool, len(products))
for _, p := range products {
allowed[p] = true
}
} else {
allowed = DirectRuntimeProductIDs()
}
if len(allowed) == 0 {
return nil
}
@@ -70,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
}
+317 -25
View File
@@ -15,23 +15,69 @@ package app
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"strings"
"sync"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"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"
@@ -87,15 +133,25 @@ 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)
}
}
catalogStart := time.Now()
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)
@@ -116,8 +172,59 @@ 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) {
tc := r.transport.WithAuth(r.resolveAuthToken(ctx), resolveIdentityHeaders())
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()
}
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{
@@ -148,19 +255,66 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
}, nil
}
callResult, err := tc.CallTool(ctx, endpoint, invocation.Tool, invocation.Params)
// Fail-fast: reject unauthenticated requests before making network calls.
// This provides a clear error message instead of cryptic HTTP 400 from MCP.
if strings.TrimSpace(authToken) == "" {
return executor.Result{}, apperrors.NewAuth(
"未登录,请先执行 dws auth login",
apperrors.WithReason("not_authenticated"),
apperrors.WithHint("运行 'dws auth login' 完成登录后重试"),
apperrors.WithActions("dws auth login"),
)
}
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(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 {
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)
mcpErr := apperrors.NewAPI(
extractMCPErrorMessage(callResult),
apperrors.WithOperation("tools/call"),
apperrors.WithReason("mcp_tool_error"),
apperrors.WithServerKey(invocation.CanonicalProduct),
apperrors.WithHint("MCP tool returned a business error; check tool parameters and refer to skill documentation."),
apperrors.WithServerDiag(diag),
)
captureRuntimeFailure(invocation, mcpErr, mcpErr)
return executor.Result{}, mcpErr
@@ -172,11 +326,14 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
}
if bizErr := detectBusinessError(callResult.Content); bizErr != "" {
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
logBusinessError(r.transport.FileLogger, "business_error", invocation, callResult.Content, diag)
return executor.Result{}, apperrors.NewAPI(bizErr,
apperrors.WithOperation("tools/call"),
apperrors.WithReason("business_error"),
apperrors.WithServerKey(invocation.CanonicalProduct),
apperrors.WithHint("The API returned a business-level error. Check required parameters and values."),
apperrors.WithServerDiag(diag),
)
}
@@ -191,37 +348,128 @@ 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 {
if token := strings.TrimSpace(explicitToken); token != "" {
return token
}
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) != "" {
return strings.TrimSpace(token)
}
// If the error is a decryption failure (corrupted data), surface
// it immediately instead of falling back to empty token.
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
slog.Error(tokenErr.Error())
return ""
}
manager := authpkg.NewManager(configDir, nil)
configureLegacyAuthManagerCompatibility(manager)
if token, _, err := manager.GetToken(); err == nil && strings.TrimSpace(token) != "" {
return strings.TrimSpace(token)
}
return ""
// Use cached token to avoid repeated Keychain access (~70ms per call)
return getCachedRuntimeToken(ctx)
}
// Cached token state for process lifetime
var (
cachedRuntimeToken string
cachedRuntimeTokenOnce sync.Once
)
// getCachedRuntimeToken returns a cached access token, loading it only once per process.
// This avoids repeated Keychain access which takes ~70ms each time.
func getCachedRuntimeToken(ctx context.Context) string {
cachedRuntimeTokenOnce.Do(func() {
loadStart := time.Now()
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
configDir := defaultConfigDir()
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
slog.Error(tokenErr.Error())
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() {
cachedRuntimeTokenOnce = sync.Once{}
cachedRuntimeToken = ""
}
func newRuntimeContentScanner() safety.Scanner {
@@ -255,6 +503,14 @@ func runtimeFlagEnabled(raw string, defaultValue bool) bool {
}
}
func isAuthError(err error) bool {
var appErr *apperrors.Error
if errors.As(err, &appErr) {
return appErr.Category == apperrors.CategoryAuth
}
return false
}
func productEndpointOverride(productID string) (string, bool) {
key := "DINGTALK_" + strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(productID), "-", "_")) + "_MCP_URL"
value := strings.TrimSpace(os.Getenv(key))
@@ -285,6 +541,9 @@ func resolveIdentityHeaders() map[string]string {
headers[k] = v
}
}
if fn := edition.Get().MergeHeaders; fn != nil {
headers = fn(headers)
}
return headers
}
@@ -328,3 +587,36 @@ func extractMCPErrorMessage(result transport.ToolCallResult) string {
}
return "MCP tool returned an error response"
}
// logBusinessError logs MCP tool errors and business errors to the file logger
// so they can be diagnosed offline. These errors arrive as HTTP 200 responses
// and would otherwise not be captured by transport-level logging.
func logBusinessError(logger *slog.Logger, reason string, inv executor.Invocation, content map[string]any, diag apperrors.ServerDiagnostics) {
if logger == nil {
return
}
attrs := []any{
"product", inv.CanonicalProduct,
"tool", inv.Tool,
"reason", reason,
}
if diag.TraceID != "" {
attrs = append(attrs, "trace_id", diag.TraceID)
}
if diag.ServerErrorCode != "" {
attrs = append(attrs, "server_error_code", diag.ServerErrorCode)
}
if diag.TechnicalDetail != "" {
attrs = append(attrs, "technical_detail", diag.TechnicalDetail)
}
if msg, ok := content["error"].(string); ok {
attrs = append(attrs, "error", msg)
}
if msg, ok := content["errorMsg"].(string); ok {
attrs = append(attrs, "errorMsg", msg)
}
if msg, ok := content["message"].(string); ok {
attrs = append(attrs, "message", msg)
}
logger.Warn("business_error", attrs...)
}
+176 -6
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) {
@@ -45,7 +95,7 @@ func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
@@ -90,7 +140,7 @@ func TestRuntimeRunnerBlocksUnsafeContentWhenEnforced(t *testing.T) {
cmd := NewRootCommand()
cmd.SetOut(&bytes.Buffer{})
cmd.SetErr(&bytes.Buffer{})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
err := cmd.Execute()
if err == nil {
@@ -121,7 +171,7 @@ func TestCanonicalCommandUsesRuntimeRunnerWhenEnabled(t *testing.T) {
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--yes"})
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--yes", "--token", "test-token"})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
@@ -255,6 +305,37 @@ func TestRuntimeRunnerInjectsAuthTokenFromFlag(t *testing.T) {
}
}
// TestRuntimeRunnerRejectsUnauthenticatedRequest verifies that requests without
// a valid token are rejected with a clear error before making any network call.
func TestRuntimeRunnerRejectsUnauthenticatedRequest(t *testing.T) {
setupRuntimeCommandTest(t)
server := mockmcp.DefaultServer()
defer server.Close()
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.RemoteURL("/server/doc"), false))
cmd := NewRootCommand()
var stdout, stderr bytes.Buffer
cmd.SetOut(&stdout)
cmd.SetErr(&stderr)
// No --token flag, should be rejected
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
err := cmd.Execute()
if err == nil {
t.Fatal("Execute() error = nil, want authentication error")
}
// Verify we get a clear auth error, not a cryptic HTTP 400
errMsg := err.Error()
if !strings.Contains(errMsg, "未登录") {
t.Fatalf("Execute() error = %v, want error containing '未登录'", err)
}
if !strings.Contains(errMsg, "auth login") {
t.Fatalf("Execute() error = %v, want error containing 'auth login'", err)
}
}
func TestRuntimeRunnerFallsBackForUnavailableProduct(t *testing.T) {
setupRuntimeCommandTest(t)
server := mockmcp.DefaultServer()
@@ -415,7 +496,7 @@ func TestCanonicalSensitiveToolAcceptsInteractiveConfirmation(t *testing.T) {
cmd.SetOut(&out)
cmd.SetErr(&errOut)
cmd.SetIn(strings.NewReader("yes\n"))
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`})
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--token", "test-token"})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
@@ -465,7 +546,7 @@ func TestRuntimeRunnerUsesProductEndpointOverride(t *testing.T) {
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
@@ -565,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")
@@ -628,7 +798,7 @@ func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
cmd := NewRootCommand()
cmd.SetOut(&bytes.Buffer{})
cmd.SetErr(&bytes.Buffer{})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
err := cmd.Execute()
if err == nil {
+613
View File
@@ -0,0 +1,613 @@
// 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 (
"archive/zip"
"context"
"encoding/json"
"fmt"
"io"
"mime"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/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.
skillDownloadTimeout = 5 * time.Minute
)
// downloadSkillResponse represents the API response for skill download.
type downloadSkillResponse struct {
Success bool `json:"success"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMsg string `json:"errorMsg,omitempty"`
Result *downloadSkillResult `json:"result,omitempty"`
}
// downloadSkillResult contains the download URL and file name.
type downloadSkillResult struct {
DownloadURL string `json:"downloadUrl"`
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 find`.
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{
"qoder": ".qoder/skills",
"claude": ".claude/skills",
"cursor": ".cursor/skills",
"codex": ".codex/skills",
"opencode": filepath.Join(".config", "opencode", "skills"),
}
// supportedTargets returns a comma-separated list of supported targets.
func supportedTargets() string {
targets := make([]string, 0, len(agentSkillPaths)+1)
for target := range agentSkillPaths {
targets = append(targets, target)
}
targets = append(targets, ".")
return strings.Join(targets, ", ")
}
func buildSkillCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "skill",
Short: "技能管理",
Long: "管理钉钉技能市场的技能。支持搜索、下载与安装到指定 Agent 目录。",
Args: cobra.NoArgs,
TraverseChildren: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
return cmd.Help()
},
}
cmd.AddCommand(
newSkillAddCommand(),
newSkillGetCommand(),
newSkillFindCommand(),
newSkillSearchHintCommand(),
)
return cmd
}
func newSkillGetCommand() *cobra.Command {
cmd := &cobra.Command{
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 newSkillFindCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "find",
Short: "从钉钉技能市场搜索技能",
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
Example: " dws skill find --context 关键词",
DisableAutoGenTag: true,
RunE: runSkillFind,
}
cmd.Flags().String("context", "", "搜索关键词(必填)")
_ = cmd.MarkFlagRequired("context")
return cmd
}
func newSkillSearchHintCommand() *cobra.Command {
return &cobra.Command{
Use: "search",
Short: "兼容旧用法,提示使用 skill find",
Hidden: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill find --context <关键词>")
return nil
},
}
}
func newSkillAddCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "add <skillId> <target>",
Short: "下载并安装技能到指定目录",
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
参数:
skillId 技能 ID(必填),可从钉钉技能市场获取
target 安装目标(必填),支持: %s
安装路径:
qoder -> ~/.qoder/skills/
claude -> ~/.claude/skills/
cursor -> ~/.cursor/skills/
codex -> ~/.codex/skills/
opencode -> ~/.config/opencode/skills/
. -> 当前目录
示例:
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
Args: cobra.ExactArgs(2),
DisableAutoGenTag: true,
RunE: runSkillAdd,
}
return cmd
}
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("context")
accessToken, err := loadSkillAccessToken()
if err != nil {
return err
}
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
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])
if skillID == "" {
return apperrors.NewValidation("skillId is required")
}
// Resolve target path
destPath, err := resolveSkillTargetPath(target)
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
}
accessToken, err := loadSkillAccessToken()
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(cmd.Context(), skillDownloadTimeout)
defer cancel()
w := cmd.OutOrStdout()
// Step 1: Get download URL from API
fmt.Fprintf(w, "正在获取技能信息...\n")
downloadResp, err := fetchSkillDownloadInfo(ctx, accessToken, skillID)
if err != nil {
return err
}
if !downloadResp.Success {
errMsg := downloadResp.ErrorMsg
if errMsg == "" {
errMsg = downloadResp.ErrorCode
}
if errMsg == "" {
errMsg = "unknown error"
}
return apperrors.NewAPI(fmt.Sprintf("failed to get skill download info: %s", errMsg),
apperrors.WithReason(downloadResp.ErrorCode))
}
if downloadResp.Result == nil || downloadResp.Result.DownloadURL == "" {
return apperrors.NewAPI("skill download URL not found in response")
}
// Step 2: Download the skill zip file
fmt.Fprintf(w, "正在下载技能...\n")
tempZipPath, err := downloadSkillFile(ctx, downloadResp.Result.DownloadURL, downloadResp.Result.FileName)
if err != nil {
return err
}
defer cleanupTempFile(tempZipPath)
// Step 3: Extract zip to destination
fmt.Fprintf(w, "正在解压到 %s...\n", destPath)
if err := extractSkillZip(tempZipPath, destPath); err != nil {
return err
}
fmt.Fprintf(w, "\n[OK] 技能安装成功!\n")
fmt.Fprintf(w, "安装路径: %s\n", destPath)
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)
if target == "" {
return "", fmt.Errorf("target is required")
}
// Special case: current directory
if target == "." {
return os.Getwd()
}
// Look up predefined agent paths
relPath, ok := agentSkillPaths[strings.ToLower(target)]
if !ok {
return "", fmt.Errorf("unsupported target")
}
homeDir, err := os.UserHomeDir()
if err != nil {
return "", fmt.Errorf("failed to get home directory: %w", err)
}
return filepath.Join(homeDir, relPath), nil
}
// fetchSkillDownloadInfo calls the download API to get the skill download URL.
func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*downloadSkillResponse, error) {
url := fmt.Sprintf("%s?skillId=%s", skillDownloadEndpoint, skillID)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("x-user-access-token", accessToken)
client := &http.Client{Timeout: 30 * time.Second}
resp, err := client.Do(req)
if err != nil {
return nil, apperrors.NewAPI(fmt.Sprintf("failed to call download API: %v", err),
apperrors.WithRetryable(true))
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
return nil, skillAuthError()
}
if resp.StatusCode != http.StatusOK {
return nil, apperrors.NewAPI(fmt.Sprintf("download API returned HTTP %d", resp.StatusCode),
apperrors.WithRetryable(resp.StatusCode >= 500))
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024)) // 10MB limit
if err != nil {
return nil, apperrors.NewAPI(fmt.Sprintf("failed to read response: %v", err))
}
var result downloadSkillResponse
if err := json.Unmarshal(body, &result); err != nil {
return nil, apperrors.NewAPI(fmt.Sprintf("failed to parse response: %v", err))
}
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)
if err != nil {
return "", apperrors.NewInternal(fmt.Sprintf("failed to create download request: %v", err))
}
client := &http.Client{Timeout: skillDownloadTimeout}
resp, err := client.Do(req)
if err != nil {
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill: %v", err),
apperrors.WithRetryable(true))
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", apperrors.NewAPI(fmt.Sprintf("download returned HTTP %d", resp.StatusCode),
apperrors.WithRetryable(resp.StatusCode >= 500))
}
// Create temp file
if fileName == "" {
fileName = "skill.zip"
}
tempFile, err := os.CreateTemp("", "dws-skill-*.zip")
if err != nil {
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
}
tempPath := tempFile.Name()
// Copy response body to temp file
_, err = io.Copy(tempFile, resp.Body)
closeErr := tempFile.Close()
if err != nil {
os.Remove(tempPath)
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
}
if closeErr != nil {
os.Remove(tempPath)
return "", apperrors.NewInternal(fmt.Sprintf("failed to close temp file: %v", closeErr))
}
return tempPath, nil
}
// extractSkillZip extracts a zip file to the destination directory.
func extractSkillZip(zipPath, destDir string) error {
// Ensure destination directory exists
if err := os.MkdirAll(destDir, 0755); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to create destination directory: %v", err))
}
reader, err := zip.OpenReader(zipPath)
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to open zip file: %v", err))
}
defer reader.Close()
for _, file := range reader.File {
if err := extractZipFile(file, destDir); err != nil {
return err
}
}
return nil
}
// extractZipFile extracts a single file from the zip archive.
func extractZipFile(file *zip.File, destDir string) error {
// Sanitize file path to prevent zip slip attacks
filePath := filepath.Join(destDir, file.Name)
if !strings.HasPrefix(filepath.Clean(filePath), filepath.Clean(destDir)+string(os.PathSeparator)) {
return apperrors.NewValidation(fmt.Sprintf("invalid file path in zip: %s", file.Name))
}
if file.FileInfo().IsDir() {
// Use 0755 to ensure we have write permission for creating files inside
return os.MkdirAll(filePath, 0755)
}
// Ensure parent directory exists with write permission
if err := os.MkdirAll(filepath.Dir(filePath), 0755); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
}
// Extract file
srcFile, err := file.Open()
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to open file in zip: %v", err))
}
defer srcFile.Close()
// Use file mode from zip but ensure at least 0644 for files
fileMode := file.Mode()
if fileMode&0600 == 0 {
fileMode = 0644
}
destFile, err := os.OpenFile(filePath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to create file: %v", err))
}
defer destFile.Close()
if _, err := io.Copy(destFile, srcFile); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to extract file: %v", err))
}
return nil
}
// cleanupTempFile removes a temporary file, ignoring errors.
func cleanupTempFile(path string) {
if path != "" {
os.Remove(path)
}
}
+784
View File
@@ -0,0 +1,784 @@
// 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 (
"archive/zip"
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
)
func TestResolveSkillTargetPath(t *testing.T) {
homeDir, err := os.UserHomeDir()
if err != nil {
t.Fatalf("failed to get home directory: %v", err)
}
tests := []struct {
name string
target string
wantSuffix string
wantErr bool
}{
{
name: "qoder target",
target: "qoder",
wantSuffix: filepath.Join(".qoder", "skills"),
wantErr: false,
},
{
name: "claude target",
target: "claude",
wantSuffix: filepath.Join(".claude", "skills"),
wantErr: false,
},
{
name: "cursor target",
target: "cursor",
wantSuffix: filepath.Join(".cursor", "skills"),
wantErr: false,
},
{
name: "codex target",
target: "codex",
wantSuffix: filepath.Join(".codex", "skills"),
wantErr: false,
},
{
name: "opencode target",
target: "opencode",
wantSuffix: filepath.Join(".config", "opencode", "skills"),
wantErr: false,
},
{
name: "case insensitive - QODER",
target: "QODER",
wantSuffix: filepath.Join(".qoder", "skills"),
wantErr: false,
},
{
name: "case insensitive - Claude",
target: "Claude",
wantSuffix: filepath.Join(".claude", "skills"),
wantErr: false,
},
{
name: "invalid target",
target: "invalid",
wantErr: true,
},
{
name: "empty target",
target: "",
wantErr: true,
},
{
name: "whitespace only",
target: " ",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := resolveSkillTargetPath(tt.target)
if (err != nil) != tt.wantErr {
t.Errorf("resolveSkillTargetPath() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !tt.wantErr {
expected := filepath.Join(homeDir, tt.wantSuffix)
if got != expected {
t.Errorf("resolveSkillTargetPath() = %v, want %v", got, expected)
}
}
})
}
}
func TestResolveSkillTargetPathCurrentDir(t *testing.T) {
// Test "." target returns current working directory
cwd, err := os.Getwd()
if err != nil {
t.Fatalf("failed to get working directory: %v", err)
}
got, err := resolveSkillTargetPath(".")
if err != nil {
t.Errorf("resolveSkillTargetPath(\".\") error = %v", err)
return
}
if got != cwd {
t.Errorf("resolveSkillTargetPath(\".\") = %v, want %v", got, cwd)
}
}
func TestParseDownloadSkillResponse(t *testing.T) {
tests := []struct {
name string
jsonInput string
wantSuccess bool
wantURL string
wantFile string
wantErrCode string
wantErrMsg string
}{
{
name: "successful response",
jsonInput: `{
"success": true,
"result": {
"downloadUrl": "https://example.com/skill.zip",
"fileName": "my-skill.zip"
}
}`,
wantSuccess: true,
wantURL: "https://example.com/skill.zip",
wantFile: "my-skill.zip",
},
{
name: "error response",
jsonInput: `{
"success": false,
"errorCode": "SKILL_NOT_FOUND",
"errorMsg": "The skill does not exist"
}`,
wantSuccess: false,
wantErrCode: "SKILL_NOT_FOUND",
wantErrMsg: "The skill does not exist",
},
{
name: "success with empty result",
jsonInput: `{
"success": true,
"result": null
}`,
wantSuccess: true,
wantURL: "",
wantFile: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var resp downloadSkillResponse
if err := json.Unmarshal([]byte(tt.jsonInput), &resp); err != nil {
t.Fatalf("failed to unmarshal JSON: %v", err)
}
if resp.Success != tt.wantSuccess {
t.Errorf("Success = %v, want %v", resp.Success, tt.wantSuccess)
}
if tt.wantSuccess && resp.Result != nil {
if resp.Result.DownloadURL != tt.wantURL {
t.Errorf("DownloadURL = %v, want %v", resp.Result.DownloadURL, tt.wantURL)
}
if resp.Result.FileName != tt.wantFile {
t.Errorf("FileName = %v, want %v", resp.Result.FileName, tt.wantFile)
}
}
if !tt.wantSuccess {
if resp.ErrorCode != tt.wantErrCode {
t.Errorf("ErrorCode = %v, want %v", resp.ErrorCode, tt.wantErrCode)
}
if resp.ErrorMsg != tt.wantErrMsg {
t.Errorf("ErrorMsg = %v, want %v", resp.ErrorMsg, tt.wantErrMsg)
}
}
})
}
}
func TestExtractSkillZip(t *testing.T) {
// Create a temporary zip file with test content
tempDir := t.TempDir()
zipPath := filepath.Join(tempDir, "test.zip")
destDir := filepath.Join(tempDir, "extracted")
// Create zip file with test content
zipFile, err := os.Create(zipPath)
if err != nil {
t.Fatalf("failed to create zip file: %v", err)
}
zipWriter := zip.NewWriter(zipFile)
// Add a file to the zip
fileContent := []byte("test content")
writer, err := zipWriter.Create("test-file.txt")
if err != nil {
t.Fatalf("failed to create file in zip: %v", err)
}
if _, err := writer.Write(fileContent); err != nil {
t.Fatalf("failed to write file content: %v", err)
}
// Add a subdirectory with a file
writer, err = zipWriter.Create("subdir/nested-file.txt")
if err != nil {
t.Fatalf("failed to create nested file in zip: %v", err)
}
if _, err := writer.Write([]byte("nested content")); err != nil {
t.Fatalf("failed to write nested file content: %v", err)
}
if err := zipWriter.Close(); err != nil {
t.Fatalf("failed to close zip writer: %v", err)
}
if err := zipFile.Close(); err != nil {
t.Fatalf("failed to close zip file: %v", err)
}
// Extract the zip
if err := extractSkillZip(zipPath, destDir); err != nil {
t.Fatalf("extractSkillZip() error = %v", err)
}
// Verify extracted files
extractedFile := filepath.Join(destDir, "test-file.txt")
content, err := os.ReadFile(extractedFile)
if err != nil {
t.Errorf("failed to read extracted file: %v", err)
}
if string(content) != "test content" {
t.Errorf("extracted content = %v, want %v", string(content), "test content")
}
// Verify nested file
nestedFile := filepath.Join(destDir, "subdir", "nested-file.txt")
content, err = os.ReadFile(nestedFile)
if err != nil {
t.Errorf("failed to read nested file: %v", err)
}
if string(content) != "nested content" {
t.Errorf("nested content = %v, want %v", string(content), "nested content")
}
}
func TestExtractSkillZipPreventZipSlip(t *testing.T) {
tempDir := t.TempDir()
zipPath := filepath.Join(tempDir, "malicious.zip")
destDir := filepath.Join(tempDir, "extracted")
// Create a zip file with a path traversal attempt
zipFile, err := os.Create(zipPath)
if err != nil {
t.Fatalf("failed to create zip file: %v", err)
}
zipWriter := zip.NewWriter(zipFile)
// Try to create a file with path traversal
writer, err := zipWriter.Create("../../../etc/passwd")
if err != nil {
t.Fatalf("failed to create malicious file in zip: %v", err)
}
if _, err := writer.Write([]byte("malicious content")); err != nil {
t.Fatalf("failed to write malicious content: %v", err)
}
if err := zipWriter.Close(); err != nil {
t.Fatalf("failed to close zip writer: %v", err)
}
if err := zipFile.Close(); err != nil {
t.Fatalf("failed to close zip file: %v", err)
}
// Extract should fail due to zip slip protection
err = extractSkillZip(zipPath, destDir)
if err == nil {
t.Error("extractSkillZip() should have failed for zip slip attack")
}
if !strings.Contains(err.Error(), "invalid file path") {
t.Errorf("error should mention invalid file path, got: %v", err)
}
}
func TestSkillAddCommandValidation(t *testing.T) {
tests := []struct {
name string
args []string
wantErr bool
errMsg string
}{
{
name: "missing arguments",
args: []string{"skill", "add"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
{
name: "missing target",
args: []string{"skill", "add", "skill-123"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
{
name: "too many arguments",
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs(tt.args)
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
err := cmd.Execute()
if (err != nil) != tt.wantErr {
t.Errorf("Execute() error = %v, wantErr %v", err, tt.wantErr)
}
if tt.wantErr && !strings.Contains(err.Error(), tt.errMsg) {
t.Errorf("error = %v, should contain %v", err, tt.errMsg)
}
})
}
}
func TestSkillAddInvalidTarget(t *testing.T) {
// Setup: Create config directory with valid token
tempDir := t.TempDir()
configDir := filepath.Join(tempDir, "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
// Save a valid token
err := authpkg.SaveTokenData(configDir, &authpkg.TokenData{
AccessToken: "test-token",
RefreshToken: "refresh-token",
ExpiresAt: time.Now().Add(time.Hour),
RefreshExpAt: time.Now().Add(24 * time.Hour),
})
if err != nil {
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
}
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
err = cmd.Execute()
if err == nil {
t.Error("Execute() should have failed for invalid target")
}
if !strings.Contains(err.Error(), "invalid target") {
t.Errorf("error should mention invalid target, got: %v", err)
}
}
func TestSkillAddRequiresAuth(t *testing.T) {
// Setup: Create config directory without token
tempDir := t.TempDir()
configDir := filepath.Join(tempDir, "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
// Ensure the config directory exists but has no token
if err := os.MkdirAll(configDir, 0755); err != nil {
t.Fatalf("failed to create config dir: %v", err)
}
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
err := cmd.Execute()
if err == nil {
t.Error("Execute() should have failed without auth")
}
// Check for authentication-related error (English or Chinese)
errStr := err.Error()
if !strings.Contains(errStr, "not logged in") && !strings.Contains(errStr, "token") && !strings.Contains(errStr, "未登录") && !strings.Contains(errStr, "auth") {
t.Errorf("error should mention authentication, got: %v", err)
}
}
func TestFetchSkillDownloadInfoUnauthorized(t *testing.T) {
// Create mock server that returns 401
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusUnauthorized)
}))
defer server.Close()
// We can't easily test the actual fetchSkillDownloadInfo function
// because it uses a hardcoded URL. This test verifies HTTP 401 handling pattern.
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.Get(server.URL)
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusUnauthorized {
t.Errorf("expected 401, got %d", resp.StatusCode)
}
}
func TestSupportedTargets(t *testing.T) {
targets := supportedTargets()
// Should contain all predefined targets
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode", "."}
for _, expected := range expectedTargets {
if !strings.Contains(targets, expected) {
t.Errorf("supportedTargets() should contain %s, got: %s", expected, targets)
}
}
}
func TestAgentSkillPathsCrossPlatform(t *testing.T) {
// Verify that paths use platform-appropriate separators
for target, path := range agentSkillPaths {
if runtime.GOOS == "windows" {
if strings.Contains(path, "/") && !strings.Contains(path, "\\") {
// On Windows, filepath.Join should use backslashes
// But raw map values may use forward slashes
t.Logf("Note: %s path '%s' uses forward slashes (will be converted by filepath.Join)", target, path)
}
}
// Test that resolveSkillTargetPath produces valid paths
resolved, err := resolveSkillTargetPath(target)
if err != nil {
t.Errorf("resolveSkillTargetPath(%s) failed: %v", target, err)
continue
}
// Path should be absolute
if !filepath.IsAbs(resolved) {
t.Errorf("resolveSkillTargetPath(%s) returned non-absolute path: %s", target, resolved)
}
}
}
func TestCleanupTempFile(t *testing.T) {
// Create a temporary file
tempFile, err := os.CreateTemp("", "test-cleanup-*.txt")
if err != nil {
t.Fatalf("failed to create temp file: %v", err)
}
tempPath := tempFile.Name()
tempFile.Close()
// Verify file exists
if _, err := os.Stat(tempPath); os.IsNotExist(err) {
t.Fatalf("temp file should exist before cleanup")
}
// Clean up
cleanupTempFile(tempPath)
// Verify file is deleted
if _, err := os.Stat(tempPath); !os.IsNotExist(err) {
t.Errorf("temp file should be deleted after cleanup")
}
// Cleanup should not panic on empty path
cleanupTempFile("")
// Cleanup should not panic on non-existent file
cleanupTempFile("/nonexistent/path/file.txt")
}
func TestDownloadSkillResponseJSON(t *testing.T) {
// Test JSON marshaling/unmarshaling round-trip
original := downloadSkillResponse{
Success: true,
Result: &downloadSkillResult{
DownloadURL: "https://example.com/skill.zip",
FileName: "skill.zip",
},
}
data, err := json.Marshal(original)
if err != nil {
t.Fatalf("failed to marshal: %v", err)
}
var parsed downloadSkillResponse
if err := json.Unmarshal(data, &parsed); err != nil {
t.Fatalf("failed to unmarshal: %v", err)
}
if parsed.Success != original.Success {
t.Errorf("Success mismatch: got %v, want %v", parsed.Success, original.Success)
}
if parsed.Result.DownloadURL != original.Result.DownloadURL {
t.Errorf("DownloadURL mismatch: got %v, want %v", parsed.Result.DownloadURL, original.Result.DownloadURL)
}
if parsed.Result.FileName != original.Result.FileName {
t.Errorf("FileName mismatch: got %v, want %v", parsed.Result.FileName, original.Result.FileName)
}
}
func TestSkillCommandHelp(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "--help"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
output := out.String()
// Check for the Long description which is shown in help
if !strings.Contains(output, "技能") {
t.Errorf("help should mention '技能', got: %s", output)
}
for _, subcmd := range []string{"add", "find", "get"} {
if !strings.Contains(output, subcmd) {
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
}
}
}
func TestSkillAddCommandHelp(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "--help"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
output := out.String()
// Should mention supported targets
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode"}
for _, target := range expectedTargets {
if !strings.Contains(output, target) {
t.Errorf("help should mention target '%s', got: %s", target, output)
}
}
}
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 TestSkillFindCommandValidation(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "find"})
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 TestSkillSearchHintCommand(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "search"})
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 find --context") {
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")
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/zip")
w.WriteHeader(http.StatusOK)
w.Write(expectedContent)
}))
defer server.Close()
// Download the file
ctx := context.Background()
tempPath, err := downloadSkillFile(ctx, server.URL, "test.zip")
if err != nil {
t.Fatalf("downloadSkillFile() error = %v", err)
}
defer os.Remove(tempPath)
// Verify the downloaded content
content, err := os.ReadFile(tempPath)
if err != nil {
t.Fatalf("failed to read downloaded file: %v", err)
}
if !bytes.Equal(content, expectedContent) {
t.Errorf("downloaded content mismatch: got %v, want %v", content, expectedContent)
}
}
func TestDownloadSkillFileServerError(t *testing.T) {
// Create a mock server that returns 500
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
ctx := context.Background()
_, err := downloadSkillFile(ctx, server.URL, "test.zip")
if err == nil {
t.Error("downloadSkillFile() should fail on server error")
}
}
func TestExtractSkillZipEmptyZip(t *testing.T) {
tempDir := t.TempDir()
zipPath := filepath.Join(tempDir, "empty.zip")
destDir := filepath.Join(tempDir, "extracted")
// Create an empty zip file
zipFile, err := os.Create(zipPath)
if err != nil {
t.Fatalf("failed to create zip file: %v", err)
}
zipWriter := zip.NewWriter(zipFile)
if err := zipWriter.Close(); err != nil {
t.Fatalf("failed to close zip writer: %v", err)
}
if err := zipFile.Close(); err != nil {
t.Fatalf("failed to close zip file: %v", err)
}
// Extract should succeed even for empty zip
if err := extractSkillZip(zipPath, destDir); err != nil {
t.Errorf("extractSkillZip() should not fail for empty zip: %v", err)
}
// Destination directory should be created
if _, err := os.Stat(destDir); os.IsNotExist(err) {
t.Errorf("destination directory should be created")
}
}
func TestExtractSkillZipWithDirectories(t *testing.T) {
tempDir := t.TempDir()
zipPath := filepath.Join(tempDir, "test.zip")
destDir := filepath.Join(tempDir, "extracted")
// Create zip with directory entries
zipFile, err := os.Create(zipPath)
if err != nil {
t.Fatalf("failed to create zip file: %v", err)
}
zipWriter := zip.NewWriter(zipFile)
// Add a directory entry with proper permissions
header := &zip.FileHeader{
Name: "mydir/",
Method: zip.Deflate,
}
header.SetMode(0755 | os.ModeDir)
_, err = zipWriter.CreateHeader(header)
if err != nil {
t.Fatalf("failed to create directory in zip: %v", err)
}
// Add a file in the directory
fileHeader := &zip.FileHeader{
Name: "mydir/file.txt",
Method: zip.Deflate,
}
fileHeader.SetMode(0644)
writer, err := zipWriter.CreateHeader(fileHeader)
if err != nil {
t.Fatalf("failed to create file in zip: %v", err)
}
if _, err := writer.Write([]byte("content")); err != nil {
t.Fatalf("failed to write content: %v", err)
}
if err := zipWriter.Close(); err != nil {
t.Fatalf("failed to close zip writer: %v", err)
}
if err := zipFile.Close(); err != nil {
t.Fatalf("failed to close zip file: %v", err)
}
// Extract
if err := extractSkillZip(zipPath, destDir); err != nil {
t.Fatalf("extractSkillZip() error = %v", err)
}
// Verify directory was created
dirPath := filepath.Join(destDir, "mydir")
info, err := os.Stat(dirPath)
if err != nil {
t.Errorf("directory should exist: %v", err)
} else if !info.IsDir() {
t.Errorf("mydir should be a directory")
}
// Verify file exists
filePath := filepath.Join(destDir, "mydir", "file.txt")
content, err := os.ReadFile(filePath)
if err != nil {
t.Errorf("file should exist: %v", err)
} else if string(content) != "content" {
t.Errorf("file content mismatch: got %s, want 'content'", string(content))
}
}
+56
View File
@@ -0,0 +1,56 @@
// 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 (
"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.
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
stdioMu.RLock()
defer stdioMu.RUnlock()
c, ok := stdioClients[productID]
return c, ok
}
// 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)
}
+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")
}
}
+381
View File
@@ -0,0 +1,381 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"context"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
)
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{}
// TimingEntry represents a single timing measurement.
type TimingEntry struct {
Name string
Duration time.Duration
Timestamp time.Time
Seq int // insertion order
}
// TimingCollector collects timing measurements for a single command execution.
// It is safe for concurrent use.
type TimingCollector struct {
mu sync.Mutex
start time.Time
entries []TimingEntry
seq int
}
// NewTimingCollector creates a new collector with the start time set to now.
func NewTimingCollector() *TimingCollector {
return &TimingCollector{
start: time.Now(),
entries: make([]TimingEntry, 0, 16),
}
}
// Record adds a timing entry with the given name and duration.
func (tc *TimingCollector) Record(name string, d time.Duration) {
if tc == nil {
return
}
tc.mu.Lock()
defer tc.mu.Unlock()
tc.entries = append(tc.entries, TimingEntry{
Name: name,
Duration: d,
Timestamp: time.Now(),
Seq: tc.seq,
})
tc.seq++
}
// StartTimer returns a function that, when called, records the elapsed time
// since StartTimer was called. This is convenient for defer usage:
//
// defer tc.StartTimer("operation")()
func (tc *TimingCollector) StartTimer(name string) func() {
if tc == nil {
return func() {}
}
start := time.Now()
return func() {
tc.Record(name, time.Since(start))
}
}
// Total returns the total elapsed time since the collector was created.
func (tc *TimingCollector) Total() time.Duration {
if tc == nil {
return 0
}
return time.Since(tc.start)
}
// Entries returns a copy of all recorded entries in insertion order.
func (tc *TimingCollector) Entries() []TimingEntry {
if tc == nil {
return nil
}
tc.mu.Lock()
defer tc.mu.Unlock()
result := make([]TimingEntry, len(tc.entries))
copy(result, tc.entries)
sort.Slice(result, func(i, j int) bool {
return result[i].Seq < result[j].Seq
})
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[Perf] Total: %v (no detailed entries)\n", formatDuration(total))
return
}
fmt.Fprintln(w)
fmt.Fprintln(w, "[Perf] Execution breakdown:")
for _, e := range entries {
fmt.Fprintf(w, " %-30s %v\n", e.Name, formatDuration(e.Duration))
}
fmt.Fprintf(w, " %-30s %v\n", "──────────────────────────────", "──────────")
fmt.Fprintf(w, " %-30s %v\n", "Total", formatDuration(total))
}
// PrintIfEnabled prints timing info to stderr if DWS_PERF_DEBUG is set.
func (tc *TimingCollector) PrintIfEnabled() {
if tc == nil {
return
}
if os.Getenv(PerfDebugEnv) == "" {
return
}
tc.Print(os.Stderr)
}
// WithTimingCollector returns a new context with the TimingCollector attached.
func WithTimingCollector(ctx context.Context, tc *TimingCollector) context.Context {
return context.WithValue(ctx, timingContextKey{}, tc)
}
// TimingCollectorFromContext extracts the TimingCollector from context, or nil.
func TimingCollectorFromContext(ctx context.Context) *TimingCollector {
if ctx == nil {
return nil
}
tc, _ := ctx.Value(timingContextKey{}).(*TimingCollector)
return tc
}
// RecordTiming is a convenience function to record timing to the collector in context.
func RecordTiming(ctx context.Context, name string, d time.Duration) {
if tc := TimingCollectorFromContext(ctx); tc != nil {
tc.Record(name, d)
}
}
// StartTiming is a convenience function that returns a stop function for defer usage.
// Example:
//
// defer StartTiming(ctx, "operation")()
func StartTiming(ctx context.Context, name string) func() {
tc := TimingCollectorFromContext(ctx)
if tc == nil {
return func() {}
}
return tc.StartTimer(name)
}
// 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, " ")
}
+453
View File
@@ -0,0 +1,453 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"bytes"
"context"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"time"
)
func TestTimingCollector_Basic(t *testing.T) {
tc := NewTimingCollector()
if tc == nil {
t.Fatal("NewTimingCollector returned nil")
}
// Record some timings
tc.Record("op1", 10*time.Millisecond)
tc.Record("op2", 20*time.Millisecond)
entries := tc.Entries()
if len(entries) != 2 {
t.Errorf("expected 2 entries, got %d", len(entries))
}
// Check ordering
if entries[0].Name != "op1" {
t.Errorf("expected first entry to be 'op1', got %q", entries[0].Name)
}
if entries[1].Name != "op2" {
t.Errorf("expected second entry to be 'op2', got %q", entries[1].Name)
}
}
func TestTimingCollector_StartTimer(t *testing.T) {
tc := NewTimingCollector()
stop := tc.StartTimer("timed_op")
time.Sleep(5 * time.Millisecond)
stop()
entries := tc.Entries()
if len(entries) != 1 {
t.Fatalf("expected 1 entry, got %d", len(entries))
}
if entries[0].Name != "timed_op" {
t.Errorf("expected entry name 'timed_op', got %q", entries[0].Name)
}
if entries[0].Duration < 5*time.Millisecond {
t.Errorf("expected duration >= 5ms, got %v", entries[0].Duration)
}
}
func TestTimingCollector_NilSafe(t *testing.T) {
var tc *TimingCollector
// Should not panic on nil collector
tc.Record("op", 10*time.Millisecond)
stop := tc.StartTimer("op")
stop()
_ = tc.Total()
_ = tc.Entries()
tc.Print(nil)
tc.PrintIfEnabled()
}
func TestTimingCollector_Print(t *testing.T) {
tc := NewTimingCollector()
tc.Record("auth_token", 44*time.Millisecond)
tc.Record("mcp_call", 150*time.Millisecond)
var buf bytes.Buffer
tc.Print(&buf)
output := buf.String()
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'")
}
if !strings.Contains(output, "mcp_call") {
t.Error("output should contain 'mcp_call'")
}
if !strings.Contains(output, "Total") {
t.Error("output should contain 'Total'")
}
}
func TestTimingCollector_PrintIfEnabled(t *testing.T) {
// Set environment variable
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
tc := NewTimingCollector()
tc.Record("test_op", 10*time.Millisecond)
// This should not panic and should print to stderr
tc.PrintIfEnabled()
}
func TestTimingCollector_ContextIntegration(t *testing.T) {
tc := NewTimingCollector()
ctx := WithTimingCollector(context.Background(), tc)
// Retrieve from context
retrieved := TimingCollectorFromContext(ctx)
if retrieved != tc {
t.Error("TimingCollectorFromContext should return the same collector")
}
// Use convenience functions
RecordTiming(ctx, "ctx_op", 30*time.Millisecond)
stop := StartTiming(ctx, "ctx_timed")
time.Sleep(2 * time.Millisecond)
stop()
entries := tc.Entries()
if len(entries) != 2 {
t.Errorf("expected 2 entries, got %d", len(entries))
}
}
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")
}
}
func TestTimingCollectorFromContext_NoCollector(t *testing.T) {
tc := TimingCollectorFromContext(context.Background())
if tc != nil {
t.Error("TimingCollectorFromContext with no collector should return nil")
}
}
func TestStartTiming_NoCollector(t *testing.T) {
ctx := context.Background()
stop := StartTiming(ctx, "no_collector")
// Should not panic
stop()
}
func TestIsPerfDebugEnabled(t *testing.T) {
// Clear the env var first
os.Unsetenv(PerfDebugEnv)
if IsPerfDebugEnabled() {
t.Error("IsPerfDebugEnabled should return false when env var is not set")
}
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
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 '框架开销'")
}
}
+104
View File
@@ -0,0 +1,104 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"context"
"encoding/json"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
// toolCallerAdapter bridges executor.Runner to the public edition.ToolCaller
// interface so that private overlays can invoke MCP tools without importing
// internal packages.
type toolCallerAdapter struct {
runner executor.Runner
flags *GlobalFlags
}
func newToolCallerAdapter(runner executor.Runner, flags *GlobalFlags) edition.ToolCaller {
return &toolCallerAdapter{runner: runner, flags: flags}
}
func (a *toolCallerAdapter) CallTool(ctx context.Context, productID, toolName string, args map[string]any) (*edition.ToolResult, error) {
inv := executor.NewHelperInvocation("overlay."+productID+"."+toolName, productID, toolName, args)
result, err := a.runner.Run(ctx, inv)
if err != nil {
return nil, err
}
return convertResult(result), nil
}
func (a *toolCallerAdapter) Format() string {
if a.flags != nil {
return a.flags.Format
}
return "json"
}
func (a *toolCallerAdapter) DryRun() bool {
return a.flags != nil && a.flags.DryRun
}
func convertResult(r executor.Result) *edition.ToolResult {
resp := r.Response
if resp == nil {
return &edition.ToolResult{}
}
// The runtime runner stores MCP response content under "content".
contentRaw, ok := resp["content"]
if !ok {
// Dry-run or echo mode: serialize the whole response as text.
data, _ := json.Marshal(resp)
return &edition.ToolResult{
Content: []edition.ContentBlock{{Type: "text", Text: string(data)}},
}
}
// Content may be a []any of {type, text} blocks from the MCP response,
// or a single map for mock mode.
switch v := contentRaw.(type) {
case []any:
blocks := make([]edition.ContentBlock, 0, len(v))
for _, item := range v {
if m, ok := item.(map[string]any); ok {
blocks = append(blocks, edition.ContentBlock{
Type: strVal(m, "type"),
Text: strVal(m, "text"),
})
}
}
return &edition.ToolResult{Content: blocks}
case map[string]any:
data, _ := json.Marshal(v)
return &edition.ToolResult{
Content: []edition.ContentBlock{{Type: "text", Text: string(data)}},
}
default:
data, _ := json.Marshal(contentRaw)
return &edition.ToolResult{
Content: []edition.ContentBlock{{Type: "text", Text: string(data)}},
}
}
}
func strVal(m map[string]any, key string) string {
if v, ok := m[key].(string); ok {
return v
}
return ""
}
+747
View File
@@ -0,0 +1,747 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0
package app
import (
"context"
"encoding/json"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/upgrade"
"github.com/fatih/color"
"github.com/spf13/cobra"
)
var (
ugBold = color.New(color.Bold).SprintFunc()
ugGreen = color.New(color.FgGreen).SprintFunc()
ugYellow = color.New(color.FgYellow).SprintFunc()
ugRed = color.New(color.FgRed).SprintFunc()
ugCyan = color.New(color.FgCyan).SprintFunc()
ugDim = color.New(color.Faint).SprintFunc()
ugBoldGrn = color.New(color.Bold, color.FgGreen).SprintFunc()
)
const defaultListLimit = 10
func newUpgradeCommand() *cobra.Command {
var (
flagCheck bool
flagList bool
flagVersion string
flagRollback bool
flagForce bool
flagSkipSkills bool
flagAll bool
)
cmd := &cobra.Command{
Use: "upgrade",
Short: "升级 DWS CLI 到最新版本",
Long: `检查并升级 DWS CLI 到最新版本。
自动下载匹配当前平台的二进制文件和技能包,通过 SHA256 校验后原子替换。
升级前会自动备份当前版本,可通过 --rollback 回滚。`,
Example: ` dws upgrade # 交互式升级到最新版本
dws upgrade --check # 仅检查是否有新版本
dws upgrade --list # 列出最近版本
dws upgrade --list --all # 列出所有版本
dws upgrade --version v1.0.5 # 升级到指定版本
dws upgrade --rollback # 回滚到上一版本
dws upgrade -y # 跳过确认直接升级`,
Args: cobra.NoArgs,
RunE: func(cmd *cobra.Command, args []string) error {
yes, _ := cmd.Flags().GetBool("yes")
format := resolveUpgradeFormat(cmd)
if flagList {
limit := defaultListLimit
if flagAll {
limit = 0
}
return runUpgradeList(cmd, format, limit)
}
if flagRollback {
return runUpgradeRollback(yes)
}
if flagCheck {
return runUpgradeCheck(cmd, format)
}
return runUpgrade(cmd.Context(), upgradeOptions{
targetVersion: flagVersion,
force: flagForce,
skipSkills: flagSkipSkills,
yes: yes,
})
},
}
cmd.Flags().BoolVar(&flagCheck, "check", 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, "强制重新安装当前版本")
cmd.Flags().BoolVar(&flagSkipSkills, "skip-skills", false, "跳过技能包更新")
return cmd
}
type upgradeOptions struct {
targetVersion string
force bool
skipSkills bool
yes bool
}
// --- dws upgrade --check ---
func runUpgradeCheck(cmd *cobra.Command, format string) error {
client := upgrade.NewClient()
if format != "json" {
fmt.Printf(" %s\n", ugDim("检查更新..."))
}
latest, err := client.FetchLatestRelease()
if err != nil {
return fmt.Errorf("检查更新失败: %w", err)
}
currentVer := version
needsUpgrade := upgrade.NeedsUpgrade(currentVer, latest.Version)
if format == "json" {
return writeJSON(cmd.OutOrStdout(), map[string]any{
"current_version": ensureV(currentVer),
"latest_version": "v" + latest.Version,
"needs_upgrade": needsUpgrade,
"release_date": latest.Date,
"prerelease": latest.Prerelease,
"changelog": parseChangelogEntries(latest.Changelog, 10),
"release_url": latest.HTMLURL,
})
}
if !needsUpgrade {
fmt.Printf("\n %s 已是最新版本 %s\n", ugBoldGrn("✔"), ugBold(ensureV(currentVer)))
return nil
}
fmt.Println()
fmt.Printf(" %s %s %s %s\n", ugBold("新版本可用:"), ugDim(ensureV(currentVer)), ugBold("→"), ugBoldGrn("v"+latest.Version))
if latest.Date != "" {
fmt.Printf(" %s %s\n", ugBold("发布日期: "), latest.Date)
}
if latest.Prerelease {
fmt.Printf(" %s %s\n", ugBold("通道: "), ugYellow("pre-release"))
}
if entries := parseChangelogEntries(latest.Changelog, 5); len(entries) > 0 {
fmt.Printf(" %s\n", ugBold("更新内容:"))
for _, e := range entries {
fmt.Printf(" %s %s\n", ugGreen("•"), e)
}
}
fmt.Println()
fmt.Printf(" %s\n", ugDim("运行 dws upgrade 进行升级"))
return nil
}
// --- dws upgrade --list ---
// 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" {
fmt.Printf(" %s\n", ugDim("获取版本列表..."))
}
versions, err := client.FetchAllReleases()
if err != nil {
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" {
var items []map[string]any
for _, v := range versions {
items = append(items, map[string]any{
"version": "v" + v.Version,
"date": v.Date,
"prerelease": v.Prerelease,
"installed": v.Version == currentVer,
"changelog": parseChangelogEntries(v.Changelog, 10),
})
}
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 totalCount == 0 {
fmt.Printf(" %s\n", ugYellow("未找到任何版本"))
return nil
}
fmt.Println()
fmt.Printf(" %s\n", ugBold(fmt.Sprintf("%-12s %-12s %-12s %s", "VERSION", "DATE", "TYPE", "CHANGELOG")))
fmt.Printf(" %s\n", ugDim(strings.Repeat("─", 70)))
for _, v := range versions {
releaseType := ugGreen("stable")
if v.Prerelease {
releaseType = ugYellow("pre-release")
}
versionStr := fmt.Sprintf("v%-11s", v.Version)
marker := ""
if v.Version == currentVer {
versionStr = ugBoldGrn(versionStr)
marker = ugCyan(" ← 已安装")
}
changelog := ugDim(truncateChangelogForList(v.Changelog, 40))
fmt.Printf(" %s %-12s %-23s %s%s\n", versionStr, v.Date, releaseType, changelog, marker)
}
fmt.Println()
fmt.Printf(" %s %s\n", ugBold("当前版本:"), ugBoldGrn(ensureV(version)))
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
}
// --- dws upgrade --rollback ---
func runUpgradeRollback(yes bool) error {
rm := upgrade.NewRollbackManager()
backups, err := rm.ListBackups()
if err != nil {
return fmt.Errorf("获取备份列表失败: %w", err)
}
if len(backups) == 0 {
return fmt.Errorf("没有可用的备份,无法回滚")
}
target := backups[0]
targetVer := ensureV(target.Version)
currentVer := ensureV(version)
fmt.Println()
fmt.Printf(" 当前版本: %s\n", ugBold(currentVer))
fmt.Printf(" 回滚目标: %s %s\n", ugCyan(targetVer), ugDim("("+target.CreatedAt.Format("2006-01-02 15:04")+")"))
if !yes {
fmt.Println()
fmt.Printf("是否回滚到 %s? [y/N] ", ugBold(targetVer))
var answer string
fmt.Scanln(&answer)
if answer != "y" && answer != "Y" {
fmt.Println("已取消")
return nil
}
}
fmt.Print(" 回滚中...")
if err := rm.RollbackTo(target); err != nil {
return fmt.Errorf("\n回滚失败: %w", err)
}
fmt.Printf(" %s\n", ugGreen("✓"))
fmt.Println()
fmt.Printf(" %s\n", ugDim("──────────────────────────────────────"))
fmt.Printf(" %s 已回滚 %s %s %s\n", ugBoldGrn("✔"), ugDim(currentVer), ugBold("→"), ugBoldGrn(targetVer))
fmt.Printf(" %s\n", ugDim("──────────────────────────────────────"))
fmt.Println()
fmt.Printf(" %s\n", ugDim("运行 dws version 验证当前版本"))
return nil
}
// --- dws upgrade (full) ---
//
// The upgrade flow is split into two phases for atomicity:
// Phase 1 (Prepare): download, verify, extract — all in a temp directory, zero side effects.
// Phase 2 (Apply): replace binary + install skills — only runs if Phase 1 fully succeeds.
// If anything fails in Phase 1, no files on disk are modified.
func runUpgrade(ctx context.Context, opts upgradeOptions) error {
fmt.Printf(" %s\n", ugDim("检查更新..."))
if err := upgrade.EnsureUpgradeDirectories(); err != nil {
return fmt.Errorf("初始化目录结构失败: %w", err)
}
upgrade.CleanupStaleFiles()
client := upgrade.NewClient()
var release *upgrade.ReleaseInfo
var err error
if opts.targetVersion != "" {
fmt.Printf(" 指定版本: %s\n", ugCyan("v"+opts.targetVersion))
release, err = client.FetchReleaseByTag(opts.targetVersion)
if err != nil {
return fmt.Errorf("获取版本 %s 信息失败: %w", opts.targetVersion, err)
}
} else {
release, err = client.FetchLatestRelease()
if err != nil {
return fmt.Errorf("检查更新失败: %w", err)
}
}
currentVer := version
if !opts.force && !upgrade.NeedsUpgrade(currentVer, release.Version) {
fmt.Printf("\n %s 已是最新版本 %s\n", ugBoldGrn("✔"), ugBold(ensureV(currentVer)))
return nil
}
fmt.Println()
fmt.Printf(" %s %s %s %s\n", ugBold("新版本可用:"), ugDim(ensureV(currentVer)), ugBold("→"), ugBoldGrn("v"+release.Version))
if release.Date != "" {
fmt.Printf(" %s %s\n", ugBold("发布日期: "), release.Date)
}
if release.Prerelease {
fmt.Printf(" %s %s\n", ugBold("通道: "), ugYellow("pre-release"))
}
if !opts.yes {
fmt.Println()
fmt.Printf("是否升级? [y/N] ")
var answer string
fmt.Scanln(&answer)
if answer != "y" && answer != "Y" {
fmt.Println("已取消")
return nil
}
}
binaryAsset, err := upgrade.FindBinaryAsset(release.Assets)
if err != nil {
return err
}
tmpDir, err := os.MkdirTemp(upgrade.DownloadCacheDir(), "upgrade-*")
if err != nil {
tmpDir, err = os.MkdirTemp("", "dws-upgrade-*")
if err != nil {
return fmt.Errorf("创建临时目录失败: %w", err)
}
}
defer os.RemoveAll(tmpDir)
hasSkills := upgrade.FindSkillsAsset(release.Assets) != nil && !opts.skipSkills
// Steps: 1.备份 2.下载 3.校验 4.解压验证 5.替换+安装
const totalSteps = 5
stepFmt := func(n int) string { return ugBold(fmt.Sprintf("[%d/%d]", n, totalSteps)) }
// ========================================================================
// Phase 1: Prepare (download + verify + extract — no side effects)
// ========================================================================
fmt.Println()
// --- Step 1: Backup ---
fmt.Printf(" %s 备份当前版本...", stepFmt(1))
rm := upgrade.NewRollbackManager()
_, backupErr := rm.Backup(strings.TrimPrefix(currentVer, "v"))
if backupErr != nil {
fmt.Printf(" %s %v\n", ugYellow("⚠"), backupErr)
} else {
fmt.Printf(" %s\n", ugGreen("✓"))
}
// Fetch checksums.txt (needed for strict verification of both binary and skills)
var checksumsContent string
checksumsAsset := upgrade.FindChecksumsAsset(release.Assets)
if checksumsAsset != nil {
checksumsPath := filepath.Join(tmpDir, "checksums.txt")
if _, dlErr := upgrade.Download(checksumsAsset.BrowserDownloadURL, checksumsPath); dlErr == nil {
if data, readErr := os.ReadFile(checksumsPath); readErr == nil {
checksumsContent = string(data)
}
}
}
// --- Step 2: Download (binary + skills together) ---
sl := stepFmt(2)
progressPrefix := fmt.Sprintf(" %s 下载 %s", sl, ugCyan(binaryAsset.Name))
fmt.Print(progressPrefix)
start := time.Now()
binaryArchivePath := filepath.Join(tmpDir, binaryAsset.Name)
n, err := upgrade.DownloadWithProgress(ctx, binaryAsset.BrowserDownloadURL, binaryArchivePath,
func(percent float64, downloaded, total int64) {
bar := progressBar(percent)
fmt.Printf("\r %s 下载 %s [%s] %5.1f%%", sl, ugCyan(binaryAsset.Name), ugCyan(bar), percent)
})
if err != nil {
fmt.Println()
return fmt.Errorf("下载二进制失败: %w", err)
}
elapsed := time.Since(start)
clearLine := strings.Repeat(" ", 100)
var skillsZipPath string
if hasSkills {
skillsAsset := upgrade.FindSkillsAsset(release.Assets)
fmt.Printf("\r%s\r%s %s %s\n", clearLine, progressPrefix, ugGreen("✓"), ugDim(fmt.Sprintf("(%.1fMB, %.1fs)", float64(n)/1024/1024, elapsed.Seconds())))
fmt.Printf(" 下载 %s...", ugCyan("dws-skills.zip"))
skillsZipPath = filepath.Join(tmpDir, "dws-skills.zip")
if _, dlErr := upgrade.Download(skillsAsset.BrowserDownloadURL, skillsZipPath); dlErr != nil {
fmt.Printf(" %s\n", ugRed("✗"))
return fmt.Errorf("技能包下载失败: %w", dlErr)
}
fmt.Printf(" %s\n", ugGreen("✓"))
} else {
fmt.Printf("\r%s\r%s %s %s\n", clearLine, progressPrefix, ugGreen("✓"), ugDim(fmt.Sprintf("(%.1fMB, %.1fs)", float64(n)/1024/1024, elapsed.Seconds())))
}
// --- Step 3: Verify SHA256 (binary + skills together) ---
if err := strictVerifyFile(stepFmt(3), binaryArchivePath, binaryAsset.Name, binaryAsset.Digest, checksumsContent); err != nil {
return err
}
if hasSkills {
skillsAsset := upgrade.FindSkillsAsset(release.Assets)
if err := strictVerifyFile(" ", skillsZipPath, "dws-skills.zip", skillsAsset.Digest, checksumsContent); err != nil {
return err
}
}
// --- Step 4: Extract + validate ---
fmt.Printf(" %s 解压并验证...", stepFmt(4))
extractDir := filepath.Join(tmpDir, "extracted")
if strings.HasSuffix(binaryAsset.Name, ".zip") {
if err := upgrade.ExtractZip(binaryArchivePath, extractDir); err != nil {
fmt.Println()
return fmt.Errorf("解压失败: %w", err)
}
} else {
if err := extractTarGz(binaryArchivePath, extractDir); err != nil {
fmt.Println()
return fmt.Errorf("解压失败: %w", err)
}
}
binaryPath := upgrade.FindBinaryInDir(extractDir)
if binaryPath == "" {
fmt.Println()
return fmt.Errorf("在解压目录中未找到 dws 二进制文件")
}
if err := validateNewBinary(binaryPath, release.Version); err != nil {
fmt.Println()
return fmt.Errorf("验证失败: %w", err)
}
var skillSrc string
if hasSkills {
skillsExtractDir := filepath.Join(tmpDir, "skills-extracted")
os.MkdirAll(skillsExtractDir, 0755)
if err := upgrade.ExtractZip(skillsZipPath, skillsExtractDir); err != nil {
fmt.Println()
return fmt.Errorf("技能包解压失败 (文件可能损坏,请检查网络后重试): %w", err)
}
skillSrc = upgrade.LocateSkillMD(skillsExtractDir)
if skillSrc == "" {
fmt.Println()
return fmt.Errorf("技能包结构异常 (未找到 SKILL.md),请反馈到 GitHub Issues")
}
}
fmt.Printf(" %s\n", ugGreen("✓"))
// ========================================================================
// Phase 2: Apply (all preparations succeeded — now do the actual changes)
// ========================================================================
// --- Step 5: Replace binary + install skills ---
fmt.Printf(" %s 替换并安装...", stepFmt(5))
if err := upgrade.ReplaceSelf(binaryPath); err != nil {
fmt.Printf(" %s\n", ugRed("✗"))
return fmt.Errorf("替换二进制失败: %w", err)
}
if hasSkills {
result, installErr := upgrade.UpgradeSkillLocations(skillSrc)
if installErr != nil {
fmt.Printf(" %s\n", ugRed("✗"))
return fmt.Errorf("技能包安装失败: %w", installErr)
}
failed := result.Failed()
if len(failed) > 0 {
fmt.Printf(" %s\n", ugRed("✗"))
for _, d := range failed {
fmt.Printf(" %s %s %s\n", ugRed("✗"), shortenHome(d.Dir), ugDim(d.Err.Error()))
}
return fmt.Errorf("技能包安装到 %d 个目录失败,请检查权限后手动重试: dws upgrade --force", len(failed))
}
succeeded := result.Succeeded()
fmt.Printf(" %s\n", ugGreen("✓"))
fmt.Printf(" %s %s\n", ugGreen("✓"), ugDim("二进制已替换"))
fmt.Printf(" %s %s\n", ugGreen("✓"), ugDim(fmt.Sprintf("技能包已安装 (%d 个位置)", len(succeeded))))
for _, d := range succeeded {
fmt.Printf(" %s %s\n", ugDim("→"), ugCyan(shortenHome(d.Dir)))
}
} else {
fmt.Printf(" %s\n", ugGreen("✓"))
}
// Cleanup old backups
rm.Cleanup(5)
// Summary
fmt.Println()
fmt.Printf(" %s\n", ugDim("──────────────────────────────────────"))
fmt.Printf(" %s 升级完成 %s %s %s\n", ugBoldGrn("✔"), ugDim(ensureV(currentVer)), ugBold("→"), ugBoldGrn("v"+release.Version))
fmt.Printf(" %s\n", ugDim("──────────────────────────────────────"))
fmt.Println()
fmt.Printf(" %s\n", ugDim("运行 dws version 验证当前版本"))
fmt.Printf(" %s\n", ugDim("如遇问题,运行 dws upgrade --rollback 回滚"))
return nil
}
// strictVerifyFile performs SHA256 verification with strict semantics:
// - If checksum info is available and matches → ✓
// - If checksum info is available but MISMATCHES → error (abort upgrade)
// - If no checksum info at all → skip (no data to compare against)
func strictVerifyFile(label, filePath, fileName, assetDigest, checksumsContent string) error {
fmt.Printf(" %s 校验 %s...", label, fileName)
// Source 1: checksums.txt
if checksumsContent != "" {
checksums := upgrade.ParseChecksumFile(checksumsContent)
if expectedHash, ok := checksums[fileName]; ok {
if err := upgrade.VerifySHA256(filePath, expectedHash); err != nil {
fmt.Printf(" %s\n", ugRed("✗"))
return fmt.Errorf("SHA256 校验失败 (%s): %w\n 文件可能被篡改或下载不完整,请重试升级", fileName, err)
}
fmt.Printf(" %s\n", ugGreen("✓"))
return nil
}
}
// Source 2: GitHub asset digest
if digest := upgrade.ExtractDigestSHA256(assetDigest); digest != "" {
if err := upgrade.VerifySHA256(filePath, digest); err != nil {
fmt.Printf(" %s\n", ugRed("✗"))
return fmt.Errorf("SHA256 校验失败 (%s): %w\n 文件可能被篡改或下载不完整,请重试升级", fileName, err)
}
fmt.Printf(" %s\n", ugGreen("✓"))
return nil
}
// No checksum info available at all
fmt.Printf(" %s\n", ugDim("- 跳过 (无可用校验信息)"))
return nil
}
// validateNewBinary checks the downloaded binary is valid.
func validateNewBinary(binaryPath, expectedVersion string) error {
info, err := os.Stat(binaryPath)
if err != nil {
return fmt.Errorf("文件不存在: %w", err)
}
if info.Size() == 0 {
return fmt.Errorf("文件为空")
}
if err := os.Chmod(binaryPath, 0755); err != nil {
return fmt.Errorf("设置执行权限失败: %w", err)
}
// Try running the binary
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
out, err := exec.CommandContext(ctx, binaryPath, "version").CombinedOutput()
if err != nil {
return fmt.Errorf("二进制无法执行: %w", err)
}
if !strings.Contains(string(out), expectedVersion) {
// Not fatal, version format might differ
fmt.Printf("\n 注意: 版本输出中未包含 %s", expectedVersion)
}
return nil
}
// extractTarGz extracts a .tar.gz file using the system tar command.
func extractTarGz(archivePath, destDir string) error {
os.MkdirAll(destDir, 0755)
cmd := exec.Command("tar", "xzf", archivePath, "-C", destDir)
if out, err := cmd.CombinedOutput(); err != nil {
return fmt.Errorf("tar 解压失败: %v: %s", err, string(out))
}
return nil
}
func progressBar(percent float64) string {
width := 20
filled := int(percent / 100 * float64(width))
if filled > width {
filled = width
}
return strings.Repeat("█", filled) + strings.Repeat("░", width-filled)
}
// parseChangelogEntries extracts human-readable commit messages from a
// GitHub Release body. The body typically looks like:
//
// ## Changelog
// * abcdef1234 - some commit message
// * 0123456789 Merge branch 'main' into main
//
// We strip the hash prefix and skip noisy entries (Merge branch, Merge pull request).
func parseChangelogEntries(body string, maxEntries int) []string {
var entries []string
for _, line := range strings.Split(body, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
line = strings.TrimPrefix(line, "* ")
line = strings.TrimPrefix(line, "- ")
msg := stripCommitHash(line)
if msg == "" {
continue
}
if isNoiseCommit(msg) {
continue
}
entries = append(entries, msg)
if maxEntries > 0 && len(entries) >= maxEntries {
break
}
}
return entries
}
// truncateChangelog returns a short one-line summary for the --check output.
func truncateChangelog(body string) string {
entries := parseChangelogEntries(body, 3)
if len(entries) == 0 {
return ""
}
return strings.Join(entries, "; ")
}
// truncateChangelogForList returns a compact summary for the --list table.
func truncateChangelogForList(body string, maxLen int) string {
entries := parseChangelogEntries(body, 2)
if len(entries) == 0 {
return "-"
}
summary := strings.Join(entries, "; ")
if len(summary) > maxLen {
return summary[:maxLen-3] + "..."
}
return summary
}
// stripCommitHash removes a leading Git commit hash (7-40 hex chars)
// and optional separator (" - ", " ") from a line.
func stripCommitHash(line string) string {
if len(line) < 8 {
return line
}
// Check if line starts with hex chars (commit hash)
hashEnd := 0
for hashEnd < len(line) && hashEnd < 40 {
c := line[hashEnd]
if (c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F') {
hashEnd++
} else {
break
}
}
if hashEnd < 7 {
return line
}
rest := line[hashEnd:]
rest = strings.TrimPrefix(rest, " - ")
rest = strings.TrimLeft(rest, " ")
return rest
}
func isNoiseCommit(msg string) bool {
lower := strings.ToLower(msg)
noisePatterns := []string{
"merge branch",
"merge pull request",
"merge remote-tracking",
}
for _, p := range noisePatterns {
if strings.HasPrefix(lower, p) {
return true
}
}
return false
}
// ensureV ensures a version string has a "v" prefix for display consistency.
// Non-semver values like "dev" or "unknown" are returned as-is.
func ensureV(ver string) string {
if ver == "" {
return "v0.0.0"
}
if strings.HasPrefix(ver, "v") {
return ver
}
// Only add "v" prefix for semver-like strings (starts with digit)
if len(ver) > 0 && ver[0] >= '0' && ver[0] <= '9' {
return "v" + ver
}
return ver
}
// resolveUpgradeFormat returns "json" only when the user explicitly passes -f json.
// Unlike other commands, upgrade defaults to table (human-friendly) output.
func resolveUpgradeFormat(cmd *cobra.Command) string {
pf := cmd.Root().PersistentFlags()
if pf.Changed("format") {
if f, err := pf.GetString("format"); err == nil {
return strings.ToLower(strings.TrimSpace(f))
}
}
return "table"
}
func writeJSON(w interface{ Write([]byte) (int, error) }, v any) error {
enc := json.NewEncoder(w)
enc.SetIndent("", " ")
return enc.Encode(v)
}
func shortenHome(path string) string {
homeDir, err := os.UserHomeDir()
if err != nil {
return path
}
if strings.HasPrefix(path, homeDir) {
return "~" + path[len(homeDir):]
}
return path
}
+430
View File
@@ -0,0 +1,430 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0
package app
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"os"
"path/filepath"
"strings"
"testing"
"github.com/spf13/cobra"
)
// --- ensureV ---
func TestEnsureV(t *testing.T) {
tests := []struct {
in string
want string
}{
{"1.0.6", "v1.0.6"},
{"v1.0.6", "v1.0.6"},
{"0.0.1", "v0.0.1"},
{"dev", "dev"},
{"unknown", "unknown"},
{"", "v0.0.0"},
{"v", "v"},
}
for _, tt := range tests {
got := ensureV(tt.in)
if got != tt.want {
t.Errorf("ensureV(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
// --- parseChangelogEntries ---
func TestParseChangelogEntries(t *testing.T) {
body := `## Changelog
* abcdef1234567 - fix login bug
* 0123456789abc Merge branch 'main' into main
* fedcba9876543 - add upgrade command
* deadbeef12345 Merge pull request #42
* 1234567890abc - improve error handling
`
entries := parseChangelogEntries(body, 10)
if len(entries) != 3 {
t.Fatalf("len(entries) = %d, want 3 (merge commits should be filtered)", len(entries))
}
if entries[0] != "fix login bug" {
t.Errorf("entries[0] = %q, want %q", entries[0], "fix login bug")
}
if entries[1] != "add upgrade command" {
t.Errorf("entries[1] = %q, want %q", entries[1], "add upgrade command")
}
if entries[2] != "improve error handling" {
t.Errorf("entries[2] = %q, want %q", entries[2], "improve error handling")
}
}
func TestParseChangelogEntries_MaxLimit(t *testing.T) {
body := "* abc1234 - msg1\n* def5678 - msg2\n* ghi9012 - msg3\n"
entries := parseChangelogEntries(body, 2)
if len(entries) != 2 {
t.Errorf("len = %d, want 2 (should respect maxEntries)", len(entries))
}
}
func TestParseChangelogEntries_EmptyBody(t *testing.T) {
entries := parseChangelogEntries("", 10)
if len(entries) != 0 {
t.Errorf("len = %d, want 0 for empty body", len(entries))
}
}
func TestParseChangelogEntries_OnlyHeaders(t *testing.T) {
body := "## Changelog\n## Another heading\n"
entries := parseChangelogEntries(body, 10)
if len(entries) != 0 {
t.Errorf("len = %d, want 0 for headers-only body", len(entries))
}
}
func TestParseChangelogEntries_OnlyMergeCommits(t *testing.T) {
body := "* abc1234 Merge branch 'main'\n* def5678 Merge pull request #10\n"
entries := parseChangelogEntries(body, 10)
if len(entries) != 0 {
t.Errorf("len = %d, want 0 (all merge commits should be filtered)", len(entries))
}
}
func TestParseChangelogEntries_DashPrefixedLines(t *testing.T) {
body := "- fix bug\n- add feature\n"
entries := parseChangelogEntries(body, 10)
if len(entries) != 2 {
t.Fatalf("len = %d, want 2", len(entries))
}
if entries[0] != "fix bug" {
t.Errorf("entries[0] = %q, want %q", entries[0], "fix bug")
}
}
// --- stripCommitHash ---
func TestStripCommitHash(t *testing.T) {
tests := []struct {
in string
want string
}{
{"abcdef1234567 - fix bug", "fix bug"},
{"abcdef1234567 fix bug", "fix bug"},
{"short", "short"}, // too short to be a hash
{"abc123", "abc123"}, // less than 7 hex chars
{"no hash here", "no hash here"},
{"ABCDEF1234567 - upper case hash", "upper case hash"},
}
for _, tt := range tests {
got := stripCommitHash(tt.in)
if got != tt.want {
t.Errorf("stripCommitHash(%q) = %q, want %q", tt.in, got, tt.want)
}
}
}
// --- isNoiseCommit ---
func TestIsNoiseCommit(t *testing.T) {
tests := []struct {
msg string
want bool
}{
{"Merge branch 'main'", true},
{"merge branch 'develop'", true},
{"Merge pull request #42", true},
{"Merge remote-tracking branch 'origin/main'", true},
{"fix login bug", false},
{"add new feature", false},
{"merge conflicts resolved", false}, // doesn't start with "merge branch"
}
for _, tt := range tests {
got := isNoiseCommit(tt.msg)
if got != tt.want {
t.Errorf("isNoiseCommit(%q) = %v, want %v", tt.msg, got, tt.want)
}
}
}
// --- truncateChangelog ---
func TestTruncateChangelog(t *testing.T) {
body := "## Changelog\n* abc1234 - fix A\n* def5678 - fix B\n* ghi9012 - fix C\n* jkl3456 - fix D\n"
result := truncateChangelog(body)
if result == "" {
t.Error("truncateChangelog returned empty")
}
// Should contain max 3 entries separated by "; "
parts := strings.Split(result, "; ")
if len(parts) > 3 {
t.Errorf("truncateChangelog should have at most 3 entries, got %d", len(parts))
}
}
func TestTruncateChangelog_EmptyBody(t *testing.T) {
if got := truncateChangelog(""); got != "" {
t.Errorf("truncateChangelog('') = %q, want empty", got)
}
}
// --- truncateChangelogForList ---
func TestTruncateChangelogForList(t *testing.T) {
tests := []struct {
body string
maxLen int
want string
}{
{"", 40, "-"},
{"## Changelog\n", 40, "-"},
}
for _, tt := range tests {
got := truncateChangelogForList(tt.body, tt.maxLen)
if got != tt.want {
t.Errorf("truncateChangelogForList(%q, %d) = %q, want %q", tt.body, tt.maxLen, got, tt.want)
}
}
}
func TestTruncateChangelogForList_Truncation(t *testing.T) {
body := "* abc1234 - a very long commit message that should be truncated eventually\n"
result := truncateChangelogForList(body, 20)
if len(result) > 20 {
t.Errorf("result len = %d, want <= 20", len(result))
}
if !strings.HasSuffix(result, "...") {
t.Errorf("truncated result should end with '...' , got %q", result)
}
}
// --- progressBar ---
func TestProgressBar(t *testing.T) {
tests := []struct {
percent float64
filled int
}{
{0, 0},
{50, 10},
{100, 20},
{150, 20}, // capped
}
for _, tt := range tests {
bar := progressBar(tt.percent)
if len(bar) != 20*len("█") && len(bar) != 20*len("░") {
// Since multi-byte chars, just check total rune count
runes := []rune(bar)
if len(runes) != 20 {
t.Errorf("progressBar(%v) rune count = %d, want 20", tt.percent, len(runes))
}
}
filledCount := strings.Count(bar, "█")
if filledCount != tt.filled {
t.Errorf("progressBar(%v) filled = %d, want %d", tt.percent, filledCount, tt.filled)
}
}
}
// --- shortenHome ---
func TestShortenHome(t *testing.T) {
// Non-home path should be unchanged
got := shortenHome("/tmp/somewhere")
if got != "/tmp/somewhere" {
t.Errorf("shortenHome(/tmp/somewhere) = %q", got)
}
}
// --- resolveUpgradeFormat ---
func TestResolveUpgradeFormat_Default(t *testing.T) {
root := &cobra.Command{}
root.PersistentFlags().String("format", "json", "output format")
child := &cobra.Command{}
root.AddCommand(child)
// format not changed => should default to "table" for upgrade
got := resolveUpgradeFormat(child)
if got != "table" {
t.Errorf("resolveUpgradeFormat(unchanged) = %q, want %q", got, "table")
}
}
func TestResolveUpgradeFormat_ExplicitJSON(t *testing.T) {
root := &cobra.Command{}
root.PersistentFlags().String("format", "json", "output format")
child := &cobra.Command{}
root.AddCommand(child)
// Simulate user explicitly setting format
root.PersistentFlags().Set("format", "json")
got := resolveUpgradeFormat(child)
if got != "json" {
t.Errorf("resolveUpgradeFormat(explicit json) = %q, want %q", got, "json")
}
}
func TestResolveUpgradeFormat_ExplicitTable(t *testing.T) {
root := &cobra.Command{}
root.PersistentFlags().String("format", "json", "output format")
child := &cobra.Command{}
root.AddCommand(child)
root.PersistentFlags().Set("format", "table")
got := resolveUpgradeFormat(child)
if got != "table" {
t.Errorf("resolveUpgradeFormat(explicit table) = %q, want %q", got, "table")
}
}
// --- writeJSON ---
func TestWriteJSON(t *testing.T) {
var buf bytes.Buffer
data := map[string]any{
"version": "v1.0.6",
"ok": true,
}
if err := writeJSON(&buf, data); err != nil {
t.Fatalf("writeJSON() error = %v", err)
}
output := buf.String()
if !strings.Contains(output, `"version": "v1.0.6"`) {
t.Errorf("output missing version: %s", output)
}
if !strings.Contains(output, `"ok": true`) {
t.Errorf("output missing ok: %s", output)
}
}
// --- strictVerifyFile ---
func TestStrictVerifyFile_MatchesChecksums(t *testing.T) {
dir := t.TempDir()
filePath := filepath.Join(dir, "test.tar.gz")
content := []byte("valid binary content")
os.WriteFile(filePath, content, 0644)
hash := computeTestSHA256(t, content)
checksums := hash + " test.tar.gz\n"
err := strictVerifyFile("[1/5]", filePath, "test.tar.gz", "", checksums)
if err != nil {
t.Errorf("expected success, got %v", err)
}
}
func TestStrictVerifyFile_ChecksumMismatch(t *testing.T) {
dir := t.TempDir()
filePath := filepath.Join(dir, "test.tar.gz")
os.WriteFile(filePath, []byte("tampered content"), 0644)
checksums := "0000000000000000000000000000000000000000000000000000000000000000 test.tar.gz\n"
err := strictVerifyFile("[1/5]", filePath, "test.tar.gz", "", checksums)
if err == nil {
t.Fatal("expected error for checksum mismatch")
}
if !strings.Contains(err.Error(), "校验失败") {
t.Errorf("error = %q, want to contain '校验失败'", err.Error())
}
}
func TestStrictVerifyFile_DigestMismatch(t *testing.T) {
dir := t.TempDir()
filePath := filepath.Join(dir, "test.tar.gz")
os.WriteFile(filePath, []byte("tampered"), 0644)
err := strictVerifyFile("[1/5]", filePath, "test.tar.gz",
"sha256:0000000000000000000000000000000000000000000000000000000000000000",
"")
if err == nil {
t.Fatal("expected error for digest mismatch")
}
}
func TestStrictVerifyFile_NoChecksumInfo(t *testing.T) {
dir := t.TempDir()
filePath := filepath.Join(dir, "test.tar.gz")
os.WriteFile(filePath, []byte("content"), 0644)
err := strictVerifyFile("[1/5]", filePath, "test.tar.gz", "", "")
if err != nil {
t.Errorf("no checksum info should skip, not error: %v", err)
}
}
func TestStrictVerifyFile_FileNotInChecksums_FallsToDigest(t *testing.T) {
dir := t.TempDir()
filePath := filepath.Join(dir, "skills.zip")
content := []byte("skills content")
os.WriteFile(filePath, content, 0644)
hash := computeTestSHA256(t, content)
// checksums.txt has entries but NOT skills.zip
checksums := "abcdef1234567890 other-file.tar.gz\n"
err := strictVerifyFile("[1/5]", filePath, "skills.zip", "sha256:"+hash, checksums)
if err != nil {
t.Errorf("should fall through to digest and succeed: %v", err)
}
}
func computeTestSHA256(t *testing.T, data []byte) string {
t.Helper()
h := sha256.Sum256(data)
return hex.EncodeToString(h[:])
}
// --- newUpgradeCommand ---
func TestNewUpgradeCommand_Flags(t *testing.T) {
cmd := newUpgradeCommand()
if cmd.Use != "upgrade" {
t.Errorf("Use = %q, want upgrade", cmd.Use)
}
expectedFlags := []string{"check", "list", "version", "rollback", "force", "skip-skills"}
for _, name := range expectedFlags {
if cmd.Flags().Lookup(name) == nil {
t.Errorf("missing flag: --%s", name)
}
}
}
func TestNewUpgradeCommand_NoArgs(t *testing.T) {
cmd := newUpgradeCommand()
// Simulate passing positional args - should error with cobra.NoArgs
cmd.SetArgs([]string{"rollback"})
err := cmd.Execute()
if err == nil {
t.Error("expected error for positional args (NoArgs)")
}
}
func TestNewUpgradeCommand_Help(t *testing.T) {
cmd := newUpgradeCommand()
var buf bytes.Buffer
cmd.SetOut(&buf)
cmd.SetArgs([]string{"--help"})
cmd.Execute()
help := buf.String()
if !strings.Contains(help, "upgrade") {
t.Error("help should contain 'upgrade'")
}
if !strings.Contains(help, "--check") {
t.Error("help should contain --check")
}
if !strings.Contains(help, "--rollback") {
t.Error("help should contain --rollback")
}
}
+25 -1
View File
@@ -13,7 +13,22 @@
package app
var version = "v1.0.2"
var version = "dev"
// SetVersion overrides the version, build time and git commit strings.
// Called by pkg/cli.SetVersion for overlay modules that inject their own
// version info via ldflags.
func SetVersion(v, bt, gc string) {
if v != "" {
version = v
}
if bt != "" {
buildTime = bt
}
if gc != "" {
gitCommit = gc
}
}
// Version returns the current CLI version string, including build metadata
// when injected via ldflags (buildTime, gitCommit).
@@ -23,3 +38,12 @@ func Version() string {
}
return version
}
// RawVersion returns the bare version string without build metadata.
func RawVersion() string { return version }
// BuildTime returns the build timestamp injected via ldflags.
func BuildTime() string { return buildTime }
// GitCommit returns the git commit hash injected via ldflags.
func GitCommit() string { return gitCommit }
+223
View File
@@ -0,0 +1,223 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/helpers"
)
const (
// appConfigFile is the filename for storing app credentials.
appConfigFile = "app.json"
)
// AppConfig represents the application credentials configuration.
// This is stored in ~/.dws/app.json with the client secret securely stored in keychain.
type AppConfig struct {
ClientID string `json:"clientId"`
ClientSecret SecretInput `json:"clientSecret"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt,omitempty"`
}
// Cached app config for performance (avoid repeated file reads).
var (
cachedAppConfig *AppConfig
cachedAppConfigOnce sync.Once
cachedAppConfigMu sync.RWMutex
)
// Cached resolved credentials (avoid repeated keychain access).
var (
cachedResolvedID string
cachedResolvedSecret string
cachedResolvedValid bool
cachedResolvedMu sync.RWMutex
)
// GetAppConfigPath returns the path to the app config file.
func GetAppConfigPath(configDir string) string {
return filepath.Join(configDir, appConfigFile)
}
// LoadAppConfig loads the app configuration from disk.
// Returns nil, nil if the config file does not exist.
func LoadAppConfig(configDir string) (*AppConfig, error) {
path := GetAppConfigPath(configDir)
data, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("reading app config: %w", err)
}
var config AppConfig
if err := json.Unmarshal(data, &config); err != nil {
return nil, fmt.Errorf("parsing app config: %w", err)
}
return &config, nil
}
// SaveAppConfig saves the app configuration to disk.
// If the client secret is a plain string, it will be stored in keychain
// and the config file will contain a reference to it.
func SaveAppConfig(configDir string, config *AppConfig) error {
// Store plain secret in keychain, convert to reference
if config.ClientSecret.IsPlain() && config.ClientID != "" {
storedRef, err := StoreSecret(config.ClientID, config.ClientSecret)
if err != nil {
return fmt.Errorf("storing client secret: %w", err)
}
config.ClientSecret = storedRef
}
// Update timestamps
if config.CreatedAt.IsZero() {
config.CreatedAt = time.Now()
}
config.UpdatedAt = time.Now()
data, err := json.MarshalIndent(config, "", " ")
if err != nil {
return fmt.Errorf("marshaling app config: %w", err)
}
path := GetAppConfigPath(configDir)
if err := helpers.AtomicWriteJSON(path, append(data, '\n')); err != nil {
return fmt.Errorf("writing app config: %w", err)
}
// Update cache
cachedAppConfigMu.Lock()
cachedAppConfig = config
cachedAppConfigMu.Unlock()
// Invalidate resolved credentials cache so next access re-resolves
cachedResolvedMu.Lock()
cachedResolvedValid = false
cachedResolvedID = ""
cachedResolvedSecret = ""
cachedResolvedMu.Unlock()
return nil
}
// DeleteAppConfig removes the app configuration and associated keychain secrets.
func DeleteAppConfig(configDir string) error {
// Load existing config to clean up keychain
existing, _ := LoadAppConfig(configDir)
if existing != nil {
RemoveSecretStore(existing.ClientSecret)
}
// Remove config file
path := GetAppConfigPath(configDir)
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing app config: %w", err)
}
// Clear cache
cachedAppConfigMu.Lock()
cachedAppConfig = nil
cachedAppConfigMu.Unlock()
// Clear resolved credentials cache
cachedResolvedMu.Lock()
cachedResolvedValid = false
cachedResolvedID = ""
cachedResolvedSecret = ""
cachedResolvedMu.Unlock()
return nil
}
// GetCachedAppConfig returns the cached app configuration.
// It loads from disk on first call and caches the result.
// Returns nil if no configuration exists or loading fails.
func GetCachedAppConfig(configDir string) *AppConfig {
cachedAppConfigOnce.Do(func() {
cfg, err := LoadAppConfig(configDir)
if err == nil && cfg != nil {
cachedAppConfigMu.Lock()
cachedAppConfig = cfg
cachedAppConfigMu.Unlock()
}
})
cachedAppConfigMu.RLock()
defer cachedAppConfigMu.RUnlock()
return cachedAppConfig
}
// ReloadAppConfig forces a reload of the app configuration from disk.
// This should be called after SaveAppConfig to ensure the cache is updated.
func ReloadAppConfig(configDir string) (*AppConfig, error) {
cfg, err := LoadAppConfig(configDir)
if err != nil {
return nil, err
}
cachedAppConfigMu.Lock()
cachedAppConfig = cfg
cachedAppConfigMu.Unlock()
return cfg, nil
}
// HasAppConfig returns true if an app configuration file exists.
func HasAppConfig(configDir string) bool {
path := GetAppConfigPath(configDir)
_, err := os.Stat(path)
return err == nil
}
// ResolveAppCredentials resolves the client ID and secret from the app config.
// Results are cached to avoid repeated keychain access.
// Returns empty strings if the config doesn't exist or resolution fails.
func ResolveAppCredentials(configDir string) (clientID, clientSecret string) {
// Fast path: check cache first
cachedResolvedMu.RLock()
if cachedResolvedValid {
id, secret := cachedResolvedID, cachedResolvedSecret
cachedResolvedMu.RUnlock()
return id, secret
}
cachedResolvedMu.RUnlock()
// Slow path: load and cache
cachedResolvedMu.Lock()
defer cachedResolvedMu.Unlock()
// Double-check after acquiring write lock
if cachedResolvedValid {
return cachedResolvedID, cachedResolvedSecret
}
cfg := GetCachedAppConfig(configDir)
if cfg != nil {
cachedResolvedID = cfg.ClientID
if secret, err := ResolveSecret(cfg.ClientSecret); err == nil {
cachedResolvedSecret = secret
}
}
cachedResolvedValid = true
return cachedResolvedID, cachedResolvedSecret
}
+13 -2
View File
@@ -26,6 +26,7 @@ func TestClientID_RuntimeOverride(t *testing.T) {
func TestClientID_EnvFallback(t *testing.T) {
SetClientID("")
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
t.Setenv("DWS_CLIENT_ID", "env-id")
if got := ClientID(); got != "env-id" {
t.Fatalf("expected env-id, got %s", got)
@@ -34,9 +35,17 @@ func TestClientID_EnvFallback(t *testing.T) {
func TestClientID_Default(t *testing.T) {
SetClientID("")
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
t.Setenv("DWS_CLIENT_ID", "")
if got := ClientID(); got != DefaultClientID {
t.Fatalf("expected default, got %s", got)
// When DefaultClientID is a placeholder (starts with "<"), ClientID() returns empty string
if strings.HasPrefix(DefaultClientID, "<") {
if got := ClientID(); got != "" {
t.Fatalf("expected empty string for placeholder, got %s", got)
}
} else {
if got := ClientID(); got != DefaultClientID {
t.Fatalf("expected default, got %s", got)
}
}
}
@@ -51,6 +60,7 @@ func TestClientSecret_RuntimeOverride(t *testing.T) {
func TestClientSecret_EnvFallback(t *testing.T) {
SetClientSecret("")
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
t.Setenv("DWS_CLIENT_SECRET", "env-secret")
if got := ClientSecret(); got != "env-secret" {
t.Fatalf("expected env-secret, got %s", got)
@@ -59,6 +69,7 @@ func TestClientSecret_EnvFallback(t *testing.T) {
func TestClientSecret_Default(t *testing.T) {
SetClientSecret("")
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
t.Setenv("DWS_CLIENT_SECRET", "")
if got := ClientSecret(); got != DefaultClientSecret {
t.Fatalf("expected default, got %s", got)
+840
View File
@@ -0,0 +1,840 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
)
// setupMCPConfigDir creates a temp config dir with mcp_url pointing to the
// given test server and sets DWS_CONFIG_DIR via t.Setenv.
// NOTE: tests calling this must NOT use t.Parallel().
func setupMCPConfigDir(t *testing.T, srvURL string) string {
t.Helper()
dir := t.TempDir()
os.WriteFile(filepath.Join(dir, "mcp_url"), []byte(srvURL), 0o600)
t.Setenv("DWS_CONFIG_DIR", dir)
return dir
}
// resetClientIDFromMCP clears the MCP-sourced flag (test helper).
func resetClientIDFromMCP() {
clientMu.Lock()
defer clientMu.Unlock()
clientIDFromMCP = false
}
// ---------------------------------------------------------------------------
// 1. CheckCLIAuthEnabled: interface error → fail-closed with retry
// ---------------------------------------------------------------------------
func TestCheckCLIAuthEnabled_ServerError_FailClosed(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
if err == nil {
t.Fatal("expected error from CheckCLIAuthEnabled when server returns 500, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ Server 500 → fail-closed: error=%q, attempts=%d", err, calls.Load())
}
func TestCheckCLIAuthEnabled_ConnectionRefused_FailClosed(t *testing.T) {
configDir := setupMCPConfigDir(t, "http://127.0.0.1:1")
p := &OAuthProvider{
configDir: configDir,
httpClient: &http.Client{Timeout: 2 * time.Second},
}
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
if err == nil {
t.Fatal("expected error when connection is refused, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
t.Logf("✅ Connection refused → fail-closed: error=%q", err)
}
func TestCheckCLIAuthEnabled_MalformedJSON_FailClosed(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `{this is not valid json}`)
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
if err == nil {
t.Fatal("expected error for malformed JSON, got nil")
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ Malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
}
func TestCheckCLIAuthEnabled_Timeout_FailClosed(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
time.Sleep(5 * time.Second)
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{
configDir: configDir,
httpClient: &http.Client{Timeout: 200 * time.Millisecond},
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
_, err := p.CheckCLIAuthEnabled(ctx, "fake-token")
if err == nil {
t.Fatal("expected error on request timeout, got nil")
}
t.Logf("✅ Timeout → fail-closed: error=%q", err)
}
// ---------------------------------------------------------------------------
// 2. CheckCLIAuthEnabled: transient error then recovery → succeeds
// ---------------------------------------------------------------------------
func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
n := calls.Add(1)
if n <= 2 {
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: true},
})
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
if err != nil {
t.Fatalf("expected success after transient failures, got error: %v", err)
}
if !status.Success || !status.Result.CLIAuthEnabled {
t.Fatalf("expected CLIAuthEnabled=true, got %+v", status)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
}
t.Logf("✅ Transient error then success: attempts=%d, enabled=%v", calls.Load(), status.Result.CLIAuthEnabled)
}
// ---------------------------------------------------------------------------
// 3. CheckCLIAuthEnabled: normal responses (pass-through)
// ---------------------------------------------------------------------------
func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("x-user-access-token") != "good-token" {
t.Errorf("missing access token header")
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: true},
})
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
status, err := p.CheckCLIAuthEnabled(context.Background(), "good-token")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !status.Result.CLIAuthEnabled {
t.Fatal("expected CLIAuthEnabled=true")
}
t.Logf("✅ Normal enabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
}
func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: false},
})
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status.Result.CLIAuthEnabled {
t.Fatal("expected CLIAuthEnabled=false")
}
t.Logf("✅ Normal disabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
}
// ---------------------------------------------------------------------------
// 4. OAuth callback: simulates the fail-closed logic at the /callback level
// ---------------------------------------------------------------------------
func TestOAuthCallback_CLIAuthError_ShowsNotEnabledPage(t *testing.T) {
t.Parallel()
var statusErr error = fmt.Errorf("simulated network error")
var authStatus *CLIAuthStatus
_ = authStatus
// This is the exact expression used in oauth_provider.go callback:
// cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
cliAuthEnabled := statusErr == nil // false
if cliAuthEnabled {
t.Fatal("cliAuthEnabled should be false when statusErr != nil")
}
t.Logf("✅ OAuth callback: statusErr=%v → cliAuthEnabled=%v → shows notEnabledHTML (fail-closed)", statusErr, cliAuthEnabled)
}
func TestOAuthCallback_CLIAuthEnabled_ShowsSuccessPage(t *testing.T) {
t.Parallel()
var statusErr error
authStatus := &CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: true},
}
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
if !cliAuthEnabled {
t.Fatal("cliAuthEnabled should be true when API returns enabled")
}
t.Logf("✅ OAuth callback: statusErr=nil, enabled=true → cliAuthEnabled=%v → shows successHTML", cliAuthEnabled)
}
func TestOAuthCallback_CLIAuthDisabledByServer_ShowsNotEnabledPage(t *testing.T) {
t.Parallel()
var statusErr error
authStatus := &CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: false},
}
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
if cliAuthEnabled {
t.Fatal("cliAuthEnabled should be false when server says disabled")
}
t.Logf("✅ OAuth callback: statusErr=nil, enabled=false → cliAuthEnabled=%v → shows notEnabledHTML", cliAuthEnabled)
}
// ---------------------------------------------------------------------------
// 5. Device Flow: loginOnce with broken cliAuthEnabled endpoint
// ---------------------------------------------------------------------------
func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
SetClientIDFromMCP("test-client-id")
t.Cleanup(func() {
SetClientID("")
resetClientIDFromMCP()
})
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch {
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
writeServiceResult(w, true, DeviceAuthResponse{
DeviceCode: "dc-test",
UserCode: "TEST-CODE",
VerificationURI: "https://example.com/verify",
ExpiresIn: 60,
Interval: 1,
}, "", "")
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
writeServiceResult(w, true, DeviceTokenResponse{
AuthCode: "test-auth-code",
}, "", "")
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"accessToken": "test-access-token",
"refreshToken": "test-refresh-token",
"expiresIn": 7200,
"corpId": "corp123",
})
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
default:
http.Error(w, "not found: "+r.URL.Path, http.StatusNotFound)
}
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
provider := &DeviceFlowProvider{
configDir: configDir,
clientID: "test-client-id",
scope: DefaultScopes,
baseURL: srv.URL,
logger: newDeviceFlowTestLogger(),
Output: io.Discard,
httpClient: srv.Client(),
}
_, err := provider.loginOnce(context.Background(), 1)
if err == nil {
t.Fatal("expected loginOnce to fail when CLI auth check fails, got nil")
}
if !strings.Contains(err.Error(), "检查 CLI 授权状态失败") && !strings.Contains(err.Error(), "Failed to check CLI auth status") {
t.Fatalf("unexpected error message: %s", err)
}
t.Logf("✅ Device Flow: CLI auth check error → login blocked: %s", err)
}
func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
SetClientIDFromMCP("test-client-id")
t.Cleanup(func() {
SetClientID("")
resetClientIDFromMCP()
})
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch {
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
writeServiceResult(w, true, DeviceAuthResponse{
DeviceCode: "dc-test",
UserCode: "TEST-CODE",
VerificationURI: "https://example.com/verify",
ExpiresIn: 60,
Interval: 1,
}, "", "")
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
writeServiceResult(w, true, DeviceTokenResponse{
AuthCode: "test-auth-code",
}, "", "")
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"accessToken": "test-access-token",
"refreshToken": "test-refresh-token",
"expiresIn": 7200,
"corpId": "corp123",
})
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
json.NewEncoder(w).Encode(CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: false},
})
case strings.HasSuffix(r.URL.Path, SuperAdminPath):
json.NewEncoder(w).Encode(SuperAdminResponse{
Success: true,
Result: []SuperAdmin{{StaffID: "admin1", Name: "张三"}},
})
default:
http.Error(w, "not found", http.StatusNotFound)
}
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
provider := &DeviceFlowProvider{
configDir: configDir,
clientID: "test-client-id",
scope: DefaultScopes,
baseURL: srv.URL,
logger: newDeviceFlowTestLogger(),
Output: io.Discard,
httpClient: srv.Client(),
}
_, err := provider.loginOnce(context.Background(), 1)
if err == nil {
t.Fatal("expected loginOnce to fail when CLI auth is disabled, got nil")
}
t.Logf("✅ Device Flow: CLI auth disabled by server → login blocked: %s", err)
}
func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
SetClientIDFromMCP("test-client-id")
t.Cleanup(func() {
SetClientID("")
resetClientIDFromMCP()
})
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
switch {
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
writeServiceResult(w, true, DeviceAuthResponse{
DeviceCode: "dc-test",
UserCode: "TEST-CODE",
VerificationURI: "https://example.com/verify",
ExpiresIn: 60,
Interval: 1,
}, "", "")
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
writeServiceResult(w, true, DeviceTokenResponse{
AuthCode: "test-auth-code",
}, "", "")
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"accessToken": "test-access-token",
"refreshToken": "test-refresh-token",
"expiresIn": 7200,
"corpId": "corp123",
})
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
json.NewEncoder(w).Encode(CLIAuthStatus{
Success: true,
Result: struct {
CLIAuthEnabled bool `json:"cliAuthEnabled"`
}{CLIAuthEnabled: true},
})
default:
http.Error(w, "not found", http.StatusNotFound)
}
}))
defer srv.Close()
configDir := setupMCPConfigDir(t, srv.URL)
provider := &DeviceFlowProvider{
configDir: configDir,
clientID: "test-client-id",
scope: DefaultScopes,
baseURL: srv.URL,
logger: newDeviceFlowTestLogger(),
Output: io.Discard,
httpClient: srv.Client(),
}
token, err := provider.loginOnce(context.Background(), 1)
if err != nil {
t.Fatalf("expected loginOnce to succeed, got error: %v", err)
}
if token.AccessToken != "test-access-token" {
t.Fatalf("unexpected token: %s", token.AccessToken)
}
t.Logf("✅ Device Flow: CLI auth enabled → login succeeded, token=%s", token.AccessToken)
}
// ---------------------------------------------------------------------------
// 6. FetchClientIDFromMCP: /cli/clientId error handling
// ---------------------------------------------------------------------------
func TestFetchClientIDFromMCP_ServerError_FailClosed(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
_, err := FetchClientIDFromMCP(context.Background())
if err == nil {
t.Fatal("expected error when /cli/clientId returns 500, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ /cli/clientId 500 → fail-closed with retry: error=%q, attempts=%d", err, calls.Load())
}
func TestFetchClientIDFromMCP_ConnectionRefused_FailClosed(t *testing.T) {
setupMCPConfigDir(t, "http://127.0.0.1:1")
_, err := FetchClientIDFromMCP(context.Background())
if err == nil {
t.Fatal("expected error when connection is refused, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
t.Logf("✅ /cli/clientId connection refused → fail-closed: error=%q", err)
}
func TestFetchClientIDFromMCP_MalformedJSON_FailClosed(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
w.Header().Set("Content-Type", "application/json")
fmt.Fprint(w, `not json at all`)
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
_, err := FetchClientIDFromMCP(context.Background())
if err == nil {
t.Fatal("expected error for malformed JSON, got nil")
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ /cli/clientId malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
}
func TestFetchClientIDFromMCP_BusinessError_FailClosed(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(ClientIDResponse{
Success: false,
ErrorCode: "forbidden",
ErrorMsg: "access denied",
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
_, err := FetchClientIDFromMCP(context.Background())
if err == nil {
t.Fatal("expected error when server returns success=false, got nil")
}
if !strings.Contains(err.Error(), "access denied") {
t.Fatalf("error should contain server error message, got: %s", err)
}
t.Logf("✅ /cli/clientId business error → fail-closed: error=%q", err)
}
func TestFetchClientIDFromMCP_TransientThenSuccess(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
n := calls.Add(1)
if n <= 2 {
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(ClientIDResponse{
Success: true,
Result: "recovered-client-id",
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
id, err := FetchClientIDFromMCP(context.Background())
if err != nil {
t.Fatalf("expected success after transient failures, got error: %v", err)
}
if id != "recovered-client-id" {
t.Fatalf("expected client ID 'recovered-client-id', got %q", id)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
}
t.Logf("✅ /cli/clientId transient then success: attempts=%d, id=%s", calls.Load(), id)
}
func TestFetchClientIDFromMCP_Success(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != ClientIDPath {
http.Error(w, "not found", http.StatusNotFound)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(ClientIDResponse{
Success: true,
Result: "my-client-id-123",
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
id, err := FetchClientIDFromMCP(context.Background())
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if id != "my-client-id-123" {
t.Fatalf("expected 'my-client-id-123', got %q", id)
}
t.Logf("✅ /cli/clientId normal success: id=%s", id)
}
// ---------------------------------------------------------------------------
// 7. GetSuperAdmins: /cli/superAdmin error handling
// ---------------------------------------------------------------------------
func TestGetSuperAdmins_ServerError_RetriesAndFails(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
_, err := GetSuperAdmins(context.Background(), "fake-token")
if err == nil {
t.Fatal("expected error when /cli/superAdmin returns 500, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ /cli/superAdmin 500 → retried 3 times: error=%q", err)
}
func TestGetSuperAdmins_TransientThenSuccess(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
n := calls.Add(1)
if n <= 2 {
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(SuperAdminResponse{
Success: true,
Result: []SuperAdmin{{StaffID: "a1", Name: "张三"}, {StaffID: "a2", Name: "李四"}},
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
result, err := GetSuperAdmins(context.Background(), "fake-token")
if err != nil {
t.Fatalf("expected success after transient failures, got error: %v", err)
}
if !result.Success || len(result.Result) != 2 {
t.Fatalf("unexpected result: %+v", result)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 attempts, got %d", c)
}
t.Logf("✅ /cli/superAdmin transient then success: attempts=%d, admins=%v", calls.Load(), result.Result)
}
func TestGetSuperAdmins_Success(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("x-user-access-token") != "good-token" {
t.Errorf("missing access token header")
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(SuperAdminResponse{
Success: true,
Result: []SuperAdmin{{StaffID: "admin1", Name: "王五"}},
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
result, err := GetSuperAdmins(context.Background(), "good-token")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(result.Result) != 1 || result.Result[0].Name != "王五" {
t.Fatalf("unexpected result: %+v", result)
}
t.Logf("✅ /cli/superAdmin normal success: admins=%v", result.Result)
}
// ---------------------------------------------------------------------------
// 8. SendCliAuthApply: /cli/sendCliAuthApply error handling
// ---------------------------------------------------------------------------
func TestSendCliAuthApply_ServerError_RetriesAndFails(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
_, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
if err == nil {
t.Fatal("expected error when /cli/sendCliAuthApply returns 500, got nil")
}
if !strings.Contains(err.Error(), "failed after 3 attempts") {
t.Fatalf("error should mention retry exhaustion, got: %s", err)
}
if c := calls.Load(); c != 3 {
t.Fatalf("expected 3 retry attempts, got %d", c)
}
t.Logf("✅ /cli/sendCliAuthApply 500 → retried 3 times: error=%q", err)
}
func TestSendCliAuthApply_TransientThenSuccess(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
n := calls.Add(1)
if n <= 1 {
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
return
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
result, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
if err != nil {
t.Fatalf("expected success after transient failure, got error: %v", err)
}
if !result.Success || !result.Result {
t.Fatalf("unexpected result: %+v", result)
}
if c := calls.Load(); c != 2 {
t.Fatalf("expected 2 attempts, got %d", c)
}
t.Logf("✅ /cli/sendCliAuthApply transient then success: attempts=%d", calls.Load())
}
func TestSendCliAuthApply_Success(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("x-user-access-token") != "good-token" {
t.Errorf("missing access token header")
}
if !strings.Contains(r.URL.RawQuery, "adminStaffId=admin123") {
t.Errorf("missing or wrong adminStaffId param: %s", r.URL.RawQuery)
}
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
result, err := SendCliAuthApply(context.Background(), "good-token", "admin123")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !result.Success || !result.Result {
t.Fatalf("unexpected result: %+v", result)
}
t.Logf("✅ /cli/sendCliAuthApply normal success: result=%+v", result)
}
func TestSendCliAuthApply_BusinessError(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(SendApplyResponse{
Success: false,
ErrorCode: "invalid_admin",
ErrorMsg: "admin not found",
Result: false,
})
}))
defer srv.Close()
setupMCPConfigDir(t, srv.URL)
result, err := SendCliAuthApply(context.Background(), "fake-token", "nonexistent")
if err != nil {
t.Fatalf("unexpected transport error: %v", err)
}
if result.Success {
t.Fatal("expected success=false for business error")
}
t.Logf("✅ /cli/sendCliAuthApply business error: errorCode=%s, errorMsg=%s", result.ErrorCode, result.ErrorMsg)
}
+63 -1
View File
@@ -26,8 +26,8 @@ import (
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/fatih/color"
)
@@ -96,6 +96,23 @@ type serviceResult struct {
}
func (p *DeviceFlowProvider) Login(ctx context.Context) (*TokenData, error) {
// Ensure we have a valid client ID (fetch from MCP if not available)
if p.clientID == "" {
if p.logger != nil {
p.logger.Debug("client ID not configured, fetching from MCP server")
}
mcpClientID, mcpErr := FetchClientIDFromMCP(ctx)
if mcpErr != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("获取 Client ID 失败"), mcpErr)
}
p.clientID = mcpClientID
// Mark that clientID is from MCP
SetClientIDFromMCP(mcpClientID)
if p.logger != nil {
p.logger.Debug("fetched client ID from MCP server", "clientID", mcpClientID)
}
}
const maxAttempts = 3
for attempt := 1; attempt <= maxAttempts; attempt++ {
tokenData, err := p.loginOnce(ctx, attempt)
@@ -149,9 +166,54 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
if err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
}
// Check if CLI auth is enabled for this organization (fail-closed: block on error)
dfPrintStep(p.output(), 4, i18n.T("检查组织 CLI 授权状态..."), 0)
authStatus, authErr := oauthProvider.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
if authErr != nil {
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 无法检查 CLI 数据访问权限状态")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请检查网络连接后重试。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, fmt.Errorf("%s: %w", i18n.T("检查 CLI 授权状态失败"), authErr)
} else if authStatus.Success && !authStatus.Result.CLIAuthEnabled {
// CLI auth is disabled - show detailed error with admin info
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
_, _ = fmt.Fprintln(p.output(), "")
// Try to get super admin list
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
// Show up to 3 admins
maxAdmins := 3
if len(admins.Result) < maxAdmins {
maxAdmins = len(admins.Result)
}
var adminNames []string
for i := 0; i < maxAdmins; i++ {
adminNames = append(adminNames, admins.Result[i].Name)
}
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
}
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), i18n.T(" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
}
// Save token data with associated client ID for refresh
tokenData.ClientID = p.clientID
if err := SaveTokenData(p.configDir, tokenData); err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
}
// Persist app credentials if using custom client credentials
oauthProvider.persistAppConfigIfNeeded()
return tokenData, nil
}
+4
View File
@@ -46,6 +46,10 @@ func writeServiceResult(w http.ResponseWriter, success bool, result any, errCode
func TestRequestDeviceCodeSuccess(t *testing.T) {
t.Parallel()
// Set a test client ID
SetClientID("test-client-id")
t.Cleanup(func() { SetClientID("") })
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
t.Fatalf("method = %s, want POST", r.Method)
+170 -5
View File
@@ -15,9 +15,31 @@ package auth
import (
"os"
"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"
@@ -58,16 +80,107 @@ const (
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
LogoutContinueURL = "https://login.dingtalk.com"
// MCP API endpoints for CLI authorization management.
DefaultMCPBaseURL = "https://mcp.dingtalk.com"
CLIAuthEnabledPath = "/cli/cliAuthEnabled"
SuperAdminPath = "/cli/superAdmin"
SendCliAuthApplyPath = "/cli/sendCliAuthApply"
ClientIDPath = "/cli/clientId"
// MCP OAuth endpoints (used when clientId is fetched from MCP).
MCPOAuthTokenPath = "/oauth2/getToken"
MCPRefreshTokenPath = "/oauth2/refreshToken"
MCPRevokeTokenPath = "/oauth2/revokeToken"
)
// GetMCPBaseURL returns the MCP base URL with priority:
// 1. ~/.dws/mcp_url file content (for pre-release environment)
// 2. Default value (https://mcp.dingtalk.com)
func GetMCPBaseURL() string {
mcpURLPath := filepath.Join(getDefaultConfigDir(), "mcp_url")
if data, err := os.ReadFile(mcpURLPath); err == nil {
if url := strings.TrimSpace(string(data)); url != "" {
return url
}
}
return DefaultMCPBaseURL
}
// Runtime overrides set via CLI flags (--client-id, --client-secret).
// These take highest priority over environment variables and defaults.
var (
clientMu sync.RWMutex
runtimeClientID string
runtimeClientSecret string
// clientIDFromMCP indicates whether the clientID was fetched from MCP server.
// When true, MCP OAuth endpoints should be used instead of direct DingTalk API.
clientIDFromMCP bool
)
// SetClientIDFromMCP sets the clientID fetched from MCP server and marks it as MCP-sourced.
func SetClientIDFromMCP(id string) {
clientMu.Lock()
defer clientMu.Unlock()
runtimeClientID = id
clientIDFromMCP = true
}
// IsClientIDFromMCP returns true if the current clientID was fetched from MCP server.
func IsClientIDFromMCP() bool {
clientMu.RLock()
defer clientMu.RUnlock()
return clientIDFromMCP || edition.Get().AuthClientFromMCP
}
// GetUserAccessTokenURL returns the appropriate token exchange URL.
// Uses MCP endpoint when clientID is from MCP, otherwise uses direct DingTalk API.
func GetUserAccessTokenURL() string {
if IsClientIDFromMCP() {
return GetMCPBaseURL() + MCPOAuthTokenPath
}
return UserAccessTokenURL
}
// GetRefreshTokenURL returns the appropriate token refresh URL.
// Uses MCP endpoint when clientID is from MCP, otherwise uses direct DingTalk API.
func GetRefreshTokenURL() string {
if IsClientIDFromMCP() {
return GetMCPBaseURL() + MCPRefreshTokenPath
}
return UserAccessTokenURL // DingTalk uses same endpoint for refresh
}
// GetRevokeTokenURL returns the token revocation URL (MCP only).
// Returns empty string if not using MCP mode.
func GetRevokeTokenURL() string {
if IsClientIDFromMCP() {
return GetMCPBaseURL() + MCPRevokeTokenPath
}
return "" // Direct mode doesn't have revoke endpoint
}
// resolveCredentialSource determines the source of the current credentials.
// Returns one of: "flag", "env", "app", "default".
// This is used to track where credentials came from for token refresh.
func resolveCredentialSource() string {
clientMu.RLock()
hasRuntimeOverride := runtimeClientID != "" || runtimeClientSecret != ""
clientMu.RUnlock()
if hasRuntimeOverride {
return "flag"
}
// Check if loaded from app config
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
return "app"
}
if os.Getenv("DWS_CLIENT_ID") != "" || os.Getenv("DWS_CLIENT_SECRET") != "" {
return "env"
}
return "default"
}
// SetClientID allows runtime override of the client ID (e.g., from CLI flags).
func SetClientID(id string) {
clientMu.Lock()
@@ -84,8 +197,11 @@ func SetClientSecret(secret string) {
// ClientID returns the OAuth client ID with priority:
// 1. Runtime override (CLI flag --client-id)
// 2. Environment variable (DWS_CLIENT_ID)
// 3. Default hardcoded value
// 2. Persisted app config (from previous login)
// 3. Environment variable (DWS_CLIENT_ID)
// 4. Default hardcoded value (if not a placeholder)
// Returns empty string if no valid client ID is available.
// Note: MCP server fetch (priority 4 in the full flow) is handled in OAuthProvider.Login()
func ClientID() string {
clientMu.RLock()
override := runtimeClientID
@@ -93,16 +209,28 @@ 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
}
if v := os.Getenv("DWS_CLIENT_ID"); v != "" {
return v
}
return DefaultClientID
// Only return default if it's not a placeholder
if !strings.HasPrefix(DefaultClientID, "<") {
return DefaultClientID
}
return ""
}
// ClientSecret returns the OAuth client secret with priority:
// 1. Runtime override (CLI flag --client-secret)
// 2. Environment variable (DWS_CLIENT_SECRET)
// 3. Default hardcoded value
// 2. Persisted app config (from previous login, stored in keychain)
// 3. Environment variable (DWS_CLIENT_SECRET)
// 4. Default hardcoded value
func ClientSecret() string {
clientMu.RLock()
override := runtimeClientSecret
@@ -110,8 +238,45 @@ func ClientSecret() string {
if override != "" {
return override
}
// Try loading from persisted app config (secret is in keychain)
if _, secret := ResolveAppCredentials(getDefaultConfigDir()); secret != "" {
return secret
}
if v := os.Getenv("DWS_CLIENT_SECRET"); v != "" {
return v
}
return DefaultClientSecret
}
// HasValidClientSecret returns true if a valid client secret is available.
// A valid secret is one that is not a placeholder (e.g., <YOUR_CLIENT_SECRET>).
func HasValidClientSecret() bool {
secret := ClientSecret()
return secret != "" && !strings.HasPrefix(secret, "<")
}
// getRuntimeCredentials returns the runtime-override credentials if set.
// Returns empty strings if no runtime overrides were provided.
func getRuntimeCredentials() (clientID, clientSecret string) {
clientMu.RLock()
defer clientMu.RUnlock()
return runtimeClientID, runtimeClientSecret
}
// getEnvClientID returns the environment variable client ID if set.
func getEnvClientID() string {
return os.Getenv("DWS_CLIENT_ID")
}
// getDefaultConfigDir returns the default configuration directory.
// Priority: DWS_CONFIG_DIR env var > ~/.dws
func getDefaultConfigDir() string {
if envDir := os.Getenv("DWS_CONFIG_DIR"); envDir != "" {
return envDir
}
homeDir, err := os.UserHomeDir()
if err != nil {
return ".dws"
}
return filepath.Join(homeDir, ".dws")
}
+1 -1
View File
@@ -21,7 +21,7 @@ import (
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
+7 -3
View File
@@ -25,7 +25,8 @@ import (
"os"
"path/filepath"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
const identityFile = "identity.json"
@@ -82,8 +83,11 @@ func (id *Identity) Headers() map[string]string {
if id.Source != "" {
h["x-dws-source"] = id.Source
}
// Constant headers for MCP gateway tracking
h["x-dingtalk-scenario-code"] = "com.dingtalk.cli"
scenarioCode := "com.dingtalk.cli"
if sc := edition.Get().ScenarioCode; sc != "" {
scenarioCode = sc
}
h["x-dingtalk-scenario-code"] = scenarioCode
h["x-dingtalk-source"] = "github"
return h
}
+42
View File
@@ -105,3 +105,45 @@ func EnsureMigration(configDir string, logger *slog.Logger) {
func IsMigrationDone() bool {
return migrationDone
}
// Client credential storage functions.
// These store the clientSecret associated with a specific clientId,
// allowing token refresh to work even if environment variables change.
const clientSecretPrefix = "client-secret:"
// SaveClientSecret stores the client secret for a specific client ID.
// This is called during login to snapshot the credentials used.
func SaveClientSecret(clientID, clientSecret string) error {
if clientID == "" || clientSecret == "" {
return nil // Nothing to save
}
account := clientSecretPrefix + clientID
if err := keychain.Set(keychain.Service, account, clientSecret); err != nil {
return fmt.Errorf("save client secret: %w", err)
}
return nil
}
// LoadClientSecret retrieves the stored client secret for a specific client ID.
// Returns empty string if not found.
func LoadClientSecret(clientID string) string {
if clientID == "" {
return ""
}
account := clientSecretPrefix + clientID
secret, err := keychain.Get(keychain.Service, account)
if err != nil {
return ""
}
return secret
}
// DeleteClientSecret removes the stored client secret for a specific client ID.
func DeleteClientSecret(clientID string) error {
if clientID == "" {
return nil
}
account := clientSecretPrefix + clientID
return keychain.Remove(keychain.Service, account)
}
+1 -1
View File
@@ -15,8 +15,8 @@ package auth
import (
"fmt"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"log/slog"
"os"
"path/filepath"
File diff suppressed because it is too large Load Diff
+326 -12
View File
@@ -15,6 +15,7 @@ package auth
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
@@ -22,6 +23,7 @@ import (
"net"
"net/http"
"os"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
@@ -92,6 +94,23 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
}
// Fall through: full browser OAuth flow.
// Ensure we have a valid client ID (fetch from MCP if not available)
if p.clientID == "" {
if p.logger != nil {
p.logger.Debug("client ID not configured, fetching from MCP server")
}
mcpClientID, mcpErr := FetchClientIDFromMCP(ctx)
if mcpErr != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("获取 Client ID 失败"), mcpErr)
}
p.clientID = mcpClientID
// Mark that clientID is from MCP, so we use MCP OAuth endpoints
SetClientIDFromMCP(mcpClientID)
if p.logger != nil {
p.logger.Debug("fetched client ID from MCP server", "clientID", mcpClientID)
}
}
// Find a free port for the callback server.
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
@@ -100,15 +119,79 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
port := listener.Addr().(*net.TCPAddr).Port
redirectURI := fmt.Sprintf("http://127.0.0.1:%d%s", port, CallbackPath)
codeCh := make(chan string, 1)
// Channel to pass callback result (token data or error with CLI auth status)
type callbackResult struct {
token *TokenData
err error
cliAuthDisabled bool
}
resultCh := make(chan callbackResult, 1)
errCh := make(chan error, 1)
// Shared state for API handlers (protected by mutex)
var (
callbackToken *TokenData
callbackProcessedCode string // The auth code that has been successfully processed
callbackAuthDisabled bool
callbackApplySent bool // Whether apply request was sent
callbackSelectedAdminId string // Selected admin ID for apply
callbackCodeInProgress string // Code currently being processed (to prevent concurrent exchange)
callbackTokenMu sync.Mutex
)
mux := http.NewServeMux()
mux.HandleFunc(CallbackPath, func(w http.ResponseWriter, r *http.Request) {
// Get code first to check if this is a new authorization or page refresh
code := r.URL.Query().Get("authCode")
if code == "" {
code = r.URL.Query().Get("code")
}
// Check state and handle page refresh or concurrent requests
callbackTokenMu.Lock()
processedCode := callbackProcessedCode
processedAuthDisabled := callbackAuthDisabled
codeInProgress := callbackCodeInProgress
hasToken := callbackToken != nil
// Case 1: This code was already successfully processed - show cached page
if code != "" && code == processedCode {
callbackTokenMu.Unlock()
w.Header().Set("Content-Type", "text/html; charset=utf-8")
if processedAuthDisabled {
_, _ = fmt.Fprint(w, notEnabledHTML)
} else {
_, _ = fmt.Fprint(w, successHTML)
}
return
}
// Case 2: This code is being processed by another request - show wait page
if code != "" && code == codeInProgress {
callbackTokenMu.Unlock()
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = fmt.Fprint(w, `<html><head><meta http-equiv="refresh" content="1"></head><body><p>正在处理授权,请稍候...</p></body></html>`)
return
}
// Case 3: No code but we have a processed token - show cached page
if code == "" && hasToken {
callbackTokenMu.Unlock()
w.Header().Set("Content-Type", "text/html; charset=utf-8")
if processedAuthDisabled {
_, _ = fmt.Fprint(w, notEnabledHTML)
} else {
_, _ = fmt.Fprint(w, successHTML)
}
return
}
// Case 4: New code - mark as in-progress and process
if code != "" {
callbackCodeInProgress = code
}
callbackTokenMu.Unlock()
if code == "" {
select {
case errCh <- errors.New(i18n.T("回调中未收到授权码")):
@@ -118,14 +201,149 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
_, _ = fmt.Fprint(w, i18n.T("授权失败:未收到授权码"))
return
}
select {
case codeCh <- code:
// Exchange code for token
tokenData, exchangeErr := p.exchangeCode(ctx, code)
if exchangeErr != nil {
// Clear in-progress state on error
callbackTokenMu.Lock()
if callbackCodeInProgress == code {
callbackCodeInProgress = ""
}
callbackTokenMu.Unlock()
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = fmt.Fprint(w, successHTML)
default:
// Select already exited (timeout/cancel); discard late callback.
w.WriteHeader(http.StatusGone)
_, _ = fmt.Fprintf(w, "<html><body><h1>授权失败</h1><p>%s</p></body></html>", exchangeErr.Error())
select {
case resultCh <- callbackResult{err: exchangeErr}:
default:
}
return
}
// Mark as processed immediately after successful exchange
callbackTokenMu.Lock()
previouslyProcessed := callbackProcessedCode != ""
callbackToken = tokenData
callbackProcessedCode = code // Remember this code was successfully processed
callbackCodeInProgress = "" // Clear in-progress state
// Reset apply state for new authorization (user switched org)
if previouslyProcessed {
callbackApplySent = false
callbackSelectedAdminId = ""
}
callbackTokenMu.Unlock()
// Check CLI auth enabled status (fail-closed: treat errors as disabled)
authStatus, statusErr := p.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
// Update CLI auth disabled state
callbackTokenMu.Lock()
callbackAuthDisabled = !cliAuthEnabled
callbackTokenMu.Unlock()
// Display appropriate HTML based on CLI auth status
w.Header().Set("Content-Type", "text/html; charset=utf-8")
if !cliAuthEnabled {
_, _ = fmt.Fprint(w, notEnabledHTML)
} else {
_, _ = fmt.Fprint(w, successHTML)
}
// Ensure response is flushed to client
if f, ok := w.(http.Flusher); ok {
f.Flush()
}
// Notify main goroutine with full result
select {
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled}:
default:
}
})
// API endpoint: get super admins
mux.HandleFunc("/api/superAdmin", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
callbackTokenMu.Lock()
token := callbackToken
callbackTokenMu.Unlock()
if token == nil {
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
return
}
result, err := GetSuperAdmins(ctx, token.AccessToken)
if err != nil {
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
return
}
data, _ := json.Marshal(result)
_, _ = w.Write(data)
})
// API endpoint: send CLI auth apply
mux.HandleFunc("/api/sendApply", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
adminStaffID := r.URL.Query().Get("adminStaffId")
if adminStaffID == "" {
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"缺少 adminStaffId 参数"}`))
return
}
callbackTokenMu.Lock()
token := callbackToken
callbackTokenMu.Unlock()
if token == nil {
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
return
}
result, err := SendCliAuthApply(ctx, token.AccessToken, adminStaffID)
if err != nil {
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
return
}
// Mark apply as sent and save selected admin on success
if result.Success && result.Result {
callbackTokenMu.Lock()
callbackApplySent = true
callbackSelectedAdminId = adminStaffID
callbackTokenMu.Unlock()
}
data, _ := json.Marshal(result)
_, _ = w.Write(data)
})
// API endpoint: get current status (clientId, applySent, selectedAdminId)
mux.HandleFunc("/api/status", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
callbackTokenMu.Lock()
applySent := callbackApplySent
selectedAdminId := callbackSelectedAdminId
callbackTokenMu.Unlock()
_, _ = fmt.Fprintf(w, `{"clientId":"%s","applySent":%t,"selectedAdminId":"%s"}`, p.clientID, applySent, selectedAdminId)
})
// API endpoint: check CLI auth enabled status
mux.HandleFunc("/api/cliAuthEnabled", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
callbackTokenMu.Lock()
token := callbackToken
callbackTokenMu.Unlock()
if token == nil {
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
return
}
result, err := p.CheckCLIAuthEnabled(ctx, token.AccessToken)
if err != nil {
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
return
}
data, _ := json.Marshal(result)
_, _ = w.Write(data)
})
// Success page endpoint
mux.HandleFunc("/success", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = fmt.Fprint(w, successHTML)
})
server := &http.Server{Handler: mux}
@@ -161,9 +379,9 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
timeout := time.NewTimer(5 * time.Minute)
defer timeout.Stop()
var authCode string
var result callbackResult
select {
case authCode = <-codeCh:
case result = <-resultCh:
case err := <-errCh:
return nil, err
case <-timeout.C:
@@ -172,13 +390,82 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
return nil, ctx.Err()
}
tokenData, err := p.exchangeCode(ctx, authCode)
if err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
// Handle callback errors
if result.err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), result.err)
}
// Handle CLI auth disabled - keep server running for user to apply
if result.cliAuthDisabled {
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
// Poll for CLI auth status while waiting
applyTimeout := time.NewTimer(10 * time.Minute)
defer applyTimeout.Stop()
pollTicker := time.NewTicker(5 * time.Second)
defer pollTicker.Stop()
elapsedSeconds := 0
for {
select {
case <-applyTimeout.C:
return nil, errors.New(i18n.T("操作超时,请重新登录"))
case <-ctx.Done():
return nil, ctx.Err()
case <-pollTicker.C:
elapsedSeconds += 5
// Get latest token and state (user may have switched org)
callbackTokenMu.Lock()
currentToken := callbackToken
currentAuthDisabled := callbackAuthDisabled
applySent := callbackApplySent
callbackTokenMu.Unlock()
// Check if user switched to an org with CLI auth enabled
if currentToken != nil && !currentAuthDisabled {
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
time.Sleep(2 * time.Second)
result.token = currentToken
result.cliAuthDisabled = false
goto continueLogin
}
// Check if CLI auth is now enabled (admin approved)
if currentToken != nil {
authStatus, err := p.CheckCLIAuthEnabled(ctx, currentToken.AccessToken)
if err == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled {
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
time.Sleep(2 * time.Second)
result.token = currentToken
result.cliAuthDisabled = false
goto continueLogin
}
}
// Show polling status based on apply state
if applySent {
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待管理员审批中"), elapsedSeconds)
} else {
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待提交申请中"), elapsedSeconds)
}
}
}
}
continueLogin:
tokenData := result.token
// Save token data with associated client ID for refresh
tokenData.ClientID = p.clientID
if err := SaveTokenData(p.configDir, tokenData); err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
}
// Persist app credentials if using custom client credentials
p.persistAppConfigIfNeeded()
return tokenData, nil
}
@@ -290,3 +577,30 @@ func (p *OAuthProvider) Logout() error {
func (p *OAuthProvider) Status() (*TokenData, error) {
return LoadTokenData(p.configDir)
}
// persistAppConfigIfNeeded saves app credentials if custom ones were used.
// This ensures the client secret is available for future token refreshes.
func (p *OAuthProvider) persistAppConfigIfNeeded() {
// Check if custom credentials were provided via runtime flags
clientID, clientSecret := getRuntimeCredentials()
if clientID == "" || clientSecret == "" {
return
}
// Only persist if they differ from environment/default values
envID := getEnvClientID()
if clientID == envID || clientID == DefaultClientID {
return
}
// Save app config with secret stored in keychain
config := &AppConfig{
ClientID: clientID,
ClientSecret: PlainSecret(clientSecret),
}
if err := SaveAppConfig(p.configDir, config); err != nil {
if p.logger != nil {
p.logger.Warn("failed to persist app credentials", "error", err)
}
}
}
+147
View File
@@ -0,0 +1,147 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"encoding/json"
"fmt"
"os"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
)
const (
// secretKeyPrefix is the keychain account prefix for app secrets.
secretKeyPrefix = "appsecret:"
)
// SecretRef references a secret stored externally.
type SecretRef struct {
Source string `json:"source"` // "keychain" | "file"
ID string `json:"id"` // keychain key or file path
}
// SecretInput represents a secret value: either a plain string or a SecretRef object.
type SecretInput struct {
Plain string // non-empty for plain string values
Ref *SecretRef // non-nil for SecretRef values
}
// PlainSecret creates a SecretInput from a plain string.
func PlainSecret(s string) SecretInput {
return SecretInput{Plain: s}
}
// IsZero returns true if the SecretInput has no value.
func (s SecretInput) IsZero() bool {
return s.Plain == "" && s.Ref == nil
}
// IsSecretRef returns true if this is a SecretRef object.
func (s SecretInput) IsSecretRef() bool {
return s.Ref != nil
}
// IsPlain returns true if this is a plain text string (not a SecretRef).
func (s SecretInput) IsPlain() bool {
return s.Ref == nil && s.Plain != ""
}
// MarshalJSON serializes SecretInput: plain string → JSON string, SecretRef → JSON object.
func (s SecretInput) MarshalJSON() ([]byte, error) {
if s.Ref != nil {
return json.Marshal(s.Ref)
}
return json.Marshal(s.Plain)
}
// UnmarshalJSON deserializes SecretInput from either a JSON string or a SecretRef object.
func (s *SecretInput) UnmarshalJSON(data []byte) error {
// Try string first
var plain string
if err := json.Unmarshal(data, &plain); err == nil {
s.Plain = plain
s.Ref = nil
return nil
}
// Try SecretRef object
var ref SecretRef
if err := json.Unmarshal(data, &ref); err == nil && isValidSource(ref.Source) && ref.ID != "" {
s.Ref = &ref
s.Plain = ""
return nil
}
return fmt.Errorf("clientSecret must be a string or {source, id} object")
}
// ValidSecretSources is the set of recognized SecretRef sources.
var ValidSecretSources = map[string]bool{
"file": true, "keychain": true,
}
func isValidSource(source string) bool {
return ValidSecretSources[source]
}
// secretAccountKey generates the keychain account key for an app's secret.
func secretAccountKey(clientID string) string {
return secretKeyPrefix + clientID
}
// ResolveSecret resolves a SecretInput to a plain string.
// SecretRef objects are resolved by source (file / keychain).
func ResolveSecret(input SecretInput) (string, error) {
if input.Ref == nil {
return input.Plain, nil
}
switch input.Ref.Source {
case "file":
data, err := os.ReadFile(input.Ref.ID)
if err != nil {
return "", fmt.Errorf("failed to read secret file %s: %w", input.Ref.ID, err)
}
return strings.TrimSpace(string(data)), nil
case "keychain":
val, err := keychain.Get(keychain.Service, input.Ref.ID)
if err != nil {
return "", fmt.Errorf("failed to get secret from keychain: %w", err)
}
return val, nil
default:
return "", fmt.Errorf("unknown secret source: %s", input.Ref.Source)
}
}
// StoreSecret stores a plain text secret in keychain and returns a SecretRef.
// If the input is already a SecretRef, it is returned as-is.
// Returns error if keychain is unavailable.
func StoreSecret(clientID string, input SecretInput) (SecretInput, error) {
if !input.IsPlain() {
return input, nil // SecretRef → keep as-is
}
key := secretAccountKey(clientID)
if err := keychain.Set(keychain.Service, key, input.Plain); err != nil {
return SecretInput{}, fmt.Errorf("keychain unavailable: %w\nhint: use file reference in config to bypass keychain", err)
}
return SecretInput{Ref: &SecretRef{Source: "keychain", ID: key}}, nil
}
// RemoveSecretStore cleans up keychain entries when an app is removed.
// Errors are intentionally ignored — cleanup is best-effort.
func RemoveSecretStore(input SecretInput) {
if input.IsSecretRef() && input.Ref.Source == "keychain" {
_ = keychain.Remove(keychain.Service, input.Ref.ID)
}
}
+1 -1
View File
@@ -21,8 +21,8 @@ import (
"path/filepath"
"sync"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/security"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const secureDataFile = ".data"
+117 -18
View File
@@ -14,11 +14,17 @@
package auth
import (
"bytes"
"context"
"encoding/json"
"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.
@@ -32,6 +38,7 @@ type TokenData struct {
UserID string `json:"user_id,omitempty"`
UserName string `json:"user_name,omitempty"`
CorpName string `json:"corp_name,omitempty"`
ClientID string `json:"client_id,omitempty"` // Associated app client ID for refresh
UpdatedAt string `json:"updated_at,omitempty"`
Source string `json:"source,omitempty"`
}
@@ -58,54 +65,104 @@ 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
}
return legacyErr
}
// RevokeTokenRemote calls the DingTalk logout endpoint to invalidate the access token.
// RevokeTokenRemote calls the appropriate logout/revoke endpoint to invalidate the access token.
// Uses MCP revoke endpoint when clientID is from MCP, otherwise uses DingTalk logout.
// This should be called before deleting local token data.
// The function is best-effort: errors are returned but callers may choose to ignore them.
func RevokeTokenRemote(ctx context.Context) error {
// Use MCP revoke endpoint when clientID is from MCP
if IsClientIDFromMCP() {
return revokeTokenViaMCP(ctx)
}
// Direct mode: use DingTalk logout endpoint
logoutURL, err := url.Parse(LogoutURL)
if err != nil {
return fmt.Errorf("parsing logout URL: %w", err)
@@ -142,3 +199,45 @@ func RevokeTokenRemote(ctx context.Context) error {
return nil
}
// revokeTokenViaMCP revokes token via MCP endpoint.
func revokeTokenViaMCP(ctx context.Context) error {
revokeURL := GetRevokeTokenURL()
if revokeURL == "" {
return nil // No revoke endpoint available
}
// Load current token to get accessToken
tokenData, err := LoadTokenData(getDefaultConfigDir())
if err != nil || tokenData == nil {
return nil // No token to revoke
}
body := map[string]string{
"clientId": ClientID(),
"accessToken": tokenData.AccessToken,
}
bodyBytes, err := json.Marshal(body)
if err != nil {
return fmt.Errorf("marshaling revoke request: %w", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, revokeURL, bytes.NewReader(bodyBytes))
if err != nil {
return fmt.Errorf("creating revoke request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("calling revoke endpoint: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("revoke endpoint returned status %d", resp.StatusCode)
}
return nil
}
+60 -1
View File
@@ -20,17 +20,18 @@ import (
"errors"
"fmt"
"io"
"log/slog"
"sort"
"strconv"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/convert"
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/ir"
"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/pkg/convert"
"github.com/spf13/cobra"
)
@@ -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,18 +17,90 @@ import (
"context"
"encoding/json"
"fmt"
"log/slog"
"os"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/ir"
"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"
@@ -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
+87 -2
View File
@@ -14,8 +14,10 @@
package compat
import (
"bufio"
"encoding/json"
"fmt"
"os"
"sort"
"strconv"
"strings"
@@ -24,11 +26,12 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/convert"
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/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/convert"
"github.com/spf13/cobra"
"github.com/spf13/pflag"
)
type ValueKind string
@@ -137,13 +140,31 @@ 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
}
}
if blocked, _ := params["_blocked"].(bool); blocked {
return nil
// Interactive confirmation for destructive operations (consistent with Helper commands)
fmt.Fprintln(cmd.ErrOrStderr(), "⚠️ This is a destructive operation.")
fmt.Fprint(cmd.ErrOrStderr(), "Confirm? (yes/no): ")
reader := bufio.NewReader(os.Stdin)
answer, _ := reader.ReadString('\n')
answer = strings.TrimSpace(strings.ToLower(answer))
if answer != "yes" && answer != "y" {
fmt.Fprintln(cmd.ErrOrStderr(), "Operation cancelled")
return nil
}
// User confirmed, continue execution
delete(params, "_blocked")
}
invocation := executor.NewCompatibilityInvocation(
@@ -231,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)
}
}
}
+116 -15
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"
@@ -152,29 +170,112 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
}, nil
}
const perServerDiscoveryTimeout = 5 * 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))
var wg sync.WaitGroup
for _, srv := range filtered {
wg.Add(1)
go func(server market.ServerDescriptor) {
defer wg.Done()
serverCtx, cancel := context.WithTimeout(ctx, perServerDiscoveryTimeout)
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
+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 errors
import "strings"
// ServerDiagnostics holds server-side diagnostic fields extracted from
// MCP response bodies or HTTP response headers. Fields are populated
// on a best-effort basis during error construction.
type ServerDiagnostics struct {
TraceID string `json:"trace_id,omitempty"`
ServerErrorCode string `json:"server_error_code,omitempty"`
TechnicalDetail string `json:"technical_detail,omitempty"`
ServerRetryable *bool `json:"server_retryable,omitempty"`
}
// IsEmpty returns true when no diagnostic field has been populated.
func (d ServerDiagnostics) IsEmpty() bool {
return d.TraceID == "" && d.ServerErrorCode == "" &&
d.TechnicalDetail == "" && d.ServerRetryable == nil
}
// WithServerDiag attaches server diagnostics to the error.
func WithServerDiag(diag ServerDiagnostics) Option {
if diag.IsEmpty() {
return func(*Error) {}
}
return func(e *Error) {
e.ServerDiag = diag
// Override retryable if server explicitly specified.
if diag.ServerRetryable != nil {
e.Retryable = *diag.ServerRetryable
}
}
}
// WithTraceID records the server-provided trace identifier.
// Used when only the trace ID is available (e.g. from HTTP headers)
// without a full ServerDiagnostics struct.
func WithTraceID(id string) Option {
id = strings.TrimSpace(id)
if id == "" {
return func(*Error) {}
}
return func(e *Error) {
e.ServerDiag.TraceID = id
}
}
+119 -36
View File
@@ -14,13 +14,12 @@
package errors
import (
"bytes"
"encoding/json"
stderrors "errors"
"fmt"
"io"
"strings"
"bytes"
)
// Category represents a stable error class with a documented exit code.
@@ -36,18 +35,19 @@ const (
// Error is the structured repository-local error model for the Go rewrite.
type Error struct {
Category Category
Message string
Operation string
ServerKey string
Retryable bool
Reason string
Hint string
Actions []string
Snapshot string
RPCCode int `json:"rpc_code,omitempty"`
RPCData json.RawMessage `json:"rpc_data,omitempty"`
Cause error `json:"-"`
Category Category
Message string
Operation string
ServerKey string
Retryable bool
Reason string
Hint string
Actions []string
Snapshot string
RPCCode int `json:"rpc_code,omitempty"`
RPCData json.RawMessage `json:"rpc_data,omitempty"`
ServerDiag ServerDiagnostics `json:"-"`
Cause error `json:"-"`
}
func (e *Error) Error() string {
@@ -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
}
@@ -247,6 +266,23 @@ func PrintJSON(w io.Writer, err error) error {
errorPayload["rpc_data"] = parsed
}
}
if !typed.ServerDiag.IsEmpty() {
if typed.ServerDiag.TraceID != "" {
errorPayload["trace_id"] = typed.ServerDiag.TraceID
}
if typed.ServerDiag.ServerErrorCode != "" {
errorPayload["server_error_code"] = typed.ServerDiag.ServerErrorCode
// Add user-friendly hint for specific server error codes
switch typed.ServerDiag.ServerErrorCode {
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
}
}
if typed.ServerDiag.TechnicalDetail != "" {
errorPayload["technical_detail"] = typed.ServerDiag.TechnicalDetail
}
}
if typed.Cause != nil {
errorPayload["cause"] = typed.Cause.Error()
}
@@ -263,8 +299,25 @@ func PrintJSON(w io.Writer, err error) error {
return writeErr
}
// PrintHuman writes a concise human-readable error rendering.
// Verbosity controls how much detail PrintHuman includes.
type Verbosity int
const (
// VerbosityNormal shows essential info: error, hint, actions, trace_id, server_code.
VerbosityNormal Verbosity = 0
// VerbosityVerbose adds technical_detail, snapshot, execution context.
VerbosityVerbose Verbosity = 1
// VerbosityDebug adds all internal diagnostics (category, operation, reason, rpc_code).
VerbosityDebug Verbosity = 2
)
// PrintHuman writes a concise human-readable error rendering at normal verbosity.
func PrintHuman(w io.Writer, err error) error {
return PrintHumanAt(w, err, VerbosityNormal)
}
// PrintHumanAt writes a human-readable error rendering at the given verbosity level.
func PrintHumanAt(w io.Writer, err error, v Verbosity) error {
if err == nil {
return nil
}
@@ -275,21 +328,23 @@ func PrintHuman(w io.Writer, err error) error {
return writeErr
}
// Line 1: Error summary
lines := []string{
fmt.Sprintf("Error: [%s] %s", strings.ToUpper(string(typed.Category)), typed.Message),
}
if typed.Reason != "" {
lines = append(lines, fmt.Sprintf("Reason: %s", typed.Reason))
}
if typed.Operation != "" {
lines = append(lines, fmt.Sprintf("Operation: %s", typed.Operation))
}
if typed.ServerKey != "" {
lines = append(lines, fmt.Sprintf("Server: %s", typed.ServerKey))
}
// Always shown: hint, actions, retryable
if typed.Hint != "" {
lines = append(lines, fmt.Sprintf("Hint: %s", typed.Hint))
}
// Add user-friendly hint for specific server error codes
switch typed.ServerDiag.ServerErrorCode {
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
}
if len(typed.Actions) > 0 {
for _, action := range typed.Actions {
if strings.TrimSpace(action) == "" {
@@ -298,22 +353,50 @@ func PrintHuman(w io.Writer, err error) error {
lines = append(lines, fmt.Sprintf("Action: %s", action))
}
}
if typed.Snapshot != "" {
lines = append(lines, fmt.Sprintf("Snapshot: %s", typed.Snapshot))
}
if typed.RPCCode != 0 {
lines = append(lines, fmt.Sprintf("RPC Code: %d", typed.RPCCode))
}
if len(typed.RPCData) > 0 {
lines = append(lines, fmt.Sprintf("RPC Data: %s", string(typed.RPCData)))
}
if typed.Cause != nil {
lines = append(lines, fmt.Sprintf("Cause: %s", typed.Cause.Error()))
}
if typed.Retryable {
lines = append(lines, "Retryable: true")
}
// Always shown when present: Trace ID, Server Code
if typed.ServerDiag.TraceID != "" {
lines = append(lines, fmt.Sprintf("Trace ID: %s", typed.ServerDiag.TraceID))
}
if typed.ServerDiag.ServerErrorCode != "" {
lines = append(lines, fmt.Sprintf("Server Code: %s", typed.ServerDiag.ServerErrorCode))
}
// Verbose+: technical detail, snapshot, reason, server key
if v >= VerbosityVerbose {
if typed.ServerDiag.TechnicalDetail != "" {
lines = append(lines, fmt.Sprintf("Detail: %s", typed.ServerDiag.TechnicalDetail))
}
if typed.Reason != "" {
lines = append(lines, fmt.Sprintf("Reason: %s", typed.Reason))
}
if typed.ServerKey != "" {
lines = append(lines, fmt.Sprintf("Server: %s", typed.ServerKey))
}
if typed.Snapshot != "" {
lines = append(lines, fmt.Sprintf("Snapshot: %s", typed.Snapshot))
}
if typed.Cause != nil {
lines = append(lines, fmt.Sprintf("Cause: %s", typed.Cause.Error()))
}
}
// Debug: all internal diagnostics
if v >= VerbosityDebug {
if typed.Operation != "" {
lines = append(lines, fmt.Sprintf("Operation: %s", typed.Operation))
}
if typed.RPCCode != 0 {
lines = append(lines, fmt.Sprintf("RPC Code: %d", typed.RPCCode))
}
if len(typed.RPCData) > 0 {
lines = append(lines, fmt.Sprintf("RPC Data: %s", string(typed.RPCData)))
}
}
_, writeErr := fmt.Fprintln(w, strings.Join(lines, "\n"))
return writeErr
}
+17 -6
View File
@@ -94,18 +94,29 @@ func TestPrintJSON_AllFields(t *testing.T) {
}
}
func TestPrintHuman_WithCause(t *testing.T) {
func TestPrintHuman_WithCause_Verbose(t *testing.T) {
t.Parallel()
cause := fmt.Errorf("timeout")
e := &Error{Category: CategoryDiscovery, Message: "discovery failed", Cause: cause}
var buf bytes.Buffer
PrintHumanAt(&buf, e, VerbosityVerbose)
if !strings.Contains(buf.String(), "timeout") {
t.Fatalf("expected cause in verbose human output: %s", buf.String())
}
}
func TestPrintHuman_WithCause_NormalHidesCause(t *testing.T) {
t.Parallel()
cause := fmt.Errorf("timeout")
e := &Error{Category: CategoryDiscovery, Message: "discovery failed", Cause: cause}
var buf bytes.Buffer
PrintHuman(&buf, e)
if !strings.Contains(buf.String(), "timeout") {
t.Fatalf("expected cause in human output: %s", buf.String())
if strings.Contains(buf.String(), "Cause:") {
t.Fatalf("normal mode should not show Cause: %s", buf.String())
}
}
func TestPrintHuman_AllFields(t *testing.T) {
func TestPrintHuman_AllFields_Debug(t *testing.T) {
t.Parallel()
e := NewAPI("api error",
WithOperation("initialize"),
@@ -120,11 +131,11 @@ func TestPrintHuman_AllFields(t *testing.T) {
WithCause(fmt.Errorf("network")),
)
var buf bytes.Buffer
PrintHuman(&buf, e)
PrintHumanAt(&buf, e, VerbosityDebug)
out := buf.String()
for _, expected := range []string{"API", "initialize", "connection_refused", "doc", "check network", "run again", "snap", "-32601", "network", "Retryable"} {
if !strings.Contains(out, expected) {
t.Fatalf("missing %q in human output: %s", expected, out)
t.Fatalf("missing %q in debug human output: %s", expected, out)
}
}
}
+74 -8
View File
@@ -81,7 +81,7 @@ func TestPrintHuman(t *testing.T) {
t.Parallel()
var b strings.Builder
if err := PrintHuman(&b, NewValidation(
if err := PrintHumanAt(&b, NewValidation(
"bad flag",
WithReason("missing_required_flag"),
WithOperation("calendar.list"),
@@ -90,7 +90,7 @@ func TestPrintHuman(t *testing.T) {
WithRetryable(true),
WithActions("retry command"),
WithSnapshot("/tmp/dws-recovery/snapshot.json"),
)); err != nil {
), VerbosityVerbose); err != nil {
t.Fatalf("PrintHuman() error = %v", err)
}
@@ -108,13 +108,64 @@ func TestPrintHuman(t *testing.T) {
t.Fatalf("expected action in output, got %q", got)
}
if !strings.Contains(got, "Snapshot: /tmp/dws-recovery/snapshot.json") {
t.Fatalf("expected snapshot in output, got %q", got)
t.Fatalf("expected snapshot in verbose output, got %q", got)
}
if !strings.Contains(got, "Retryable: true") {
t.Fatalf("expected retryable marker in output, got %q", got)
}
}
func TestPrintHuman_NormalMode(t *testing.T) {
t.Parallel()
var b strings.Builder
PrintHuman(&b, NewValidation(
"bad flag",
WithHint("fix it"),
WithRetryable(true),
WithActions("retry"),
WithServerDiag(ServerDiagnostics{TraceID: "trace-abc", ServerErrorCode: "PARAM_ERROR"}),
))
got := b.String()
if !strings.Contains(got, "Error: [VALIDATION] bad flag") {
t.Fatalf("expected header, got %q", got)
}
if !strings.Contains(got, "Trace ID: trace-abc") {
t.Fatalf("expected trace id in normal output, got %q", got)
}
if !strings.Contains(got, "Server Code: PARAM_ERROR") {
t.Fatalf("expected server code in normal output, got %q", got)
}
}
func TestPrintJSONIncludesServerDiag(t *testing.T) {
t.Parallel()
var b strings.Builder
if err := PrintJSON(&b, NewAPI(
"server error",
WithServerDiag(ServerDiagnostics{
TraceID: "trace-xyz",
ServerErrorCode: "TIMEOUT_ERROR",
TechnicalDetail: "deadline exceeded",
}),
)); err != nil {
t.Fatalf("PrintJSON() error = %v", err)
}
got := b.String()
if !strings.Contains(got, `"trace_id": "trace-xyz"`) {
t.Fatalf("expected trace_id in output, got %q", got)
}
if !strings.Contains(got, `"server_error_code": "TIMEOUT_ERROR"`) {
t.Fatalf("expected server_error_code in output, got %q", got)
}
if !strings.Contains(got, `"technical_detail": "deadline exceeded"`) {
t.Fatalf("expected technical_detail in output, got %q", got)
}
}
func TestPrintJSONIncludesRPCCodeAndData(t *testing.T) {
t.Parallel()
@@ -137,23 +188,38 @@ func TestPrintJSONIncludesRPCCodeAndData(t *testing.T) {
}
}
func TestPrintHumanIncludesRPCCode(t *testing.T) {
func TestPrintHumanIncludesRPCCode_Debug(t *testing.T) {
t.Parallel()
var b strings.Builder
if err := PrintHuman(&b, NewValidation(
if err := PrintHumanAt(&b, NewValidation(
"invalid params",
WithRPCCode(-32602),
WithRPCData([]byte(`"missing field"`)),
)); err != nil {
), VerbosityDebug); err != nil {
t.Fatalf("PrintHuman() error = %v", err)
}
got := b.String()
if !strings.Contains(got, "RPC Code: -32602") {
t.Fatalf("expected RPC Code in output, got %q", got)
t.Fatalf("expected RPC Code in debug output, got %q", got)
}
if !strings.Contains(got, "RPC Data:") {
t.Fatalf("expected RPC Data in output, got %q", got)
t.Fatalf("expected RPC Data in debug output, got %q", got)
}
}
func TestPrintHumanHidesRPCCode_Normal(t *testing.T) {
t.Parallel()
var b strings.Builder
PrintHuman(&b, NewValidation(
"invalid params",
WithRPCCode(-32602),
))
got := b.String()
if strings.Contains(got, "RPC Code:") {
t.Fatalf("normal mode should not show RPC Code, got %q", got)
}
}
+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"
+44 -7
View File
@@ -26,9 +26,9 @@ import (
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
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/pkg/config"
"github.com/spf13/cobra"
)
@@ -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
}
+515
View File
@@ -0,0 +1,515 @@
// 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"
"time"
"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/spf13/cobra"
)
func init() {
RegisterPublic(func() Handler {
return reportHandler{}
})
}
type reportHandler struct{}
func (reportHandler) Name() string {
return "report"
}
func (reportHandler) Command(runner executor.Runner) *cobra.Command {
root := &cobra.Command{
Use: "report",
Aliases: []string{"log"},
Short: "日志 / 模版 / 统计",
Long: `钉钉日志:模版、创建、详情、列表、统计。
子命令:
template 日志模版(list / detail)
create 创建日志
detail 获取日志详情
list 查询收到的日志列表
stats 获取日志统计数据
sent 查询已发送的日志列表`,
Args: cobra.NoArgs,
TraverseChildren: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
return cmd.Help()
},
}
template := &cobra.Command{
Use: "template",
Short: "日志模版",
Args: cobra.NoArgs,
TraverseChildren: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
return cmd.Help()
},
}
template.AddCommand(
newReportTemplateListCommand(runner),
newReportTemplateDetailCommand(runner),
)
root.AddCommand(
template,
newReportCreateCommand(runner),
newReportDetailCommand(runner),
newReportListCommand(runner),
newReportStatsCommand(runner),
newReportSentCommand(runner),
)
return root
}
// ── flexTimeLayouts: supported date formats, most specific first ──
var flexTimeLayouts = []string{
time.RFC3339, // 2006-01-02T15:04:05+08:00
"2006-01-02T15:04:05Z", // UTC Z suffix
"2006-01-02T15:04:05-07:00", // with offset but no colon
"2006-01-02T15:04:05", // no timezone
"2006-01-02 15:04:05", // space-separated
"2006-01-02T15:04", // no seconds
"2006-01-02 15:04", // no seconds, space
"2006-01-02", // date only
"2006/01/02 15:04:05", // slash + time
"2006/01/02", // slash date
"20060102", // compact YYYYMMDD
}
// parseFlexTimeToMillis parses a date string using multiple formats and returns Unix milliseconds.
// Supports 11 formats for maximum compatibility with user input.
func parseFlexTimeToMillis(flagName, value string) (int64, error) {
value = strings.TrimSpace(value)
if value == "" {
return 0, apperrors.NewValidation(fmt.Sprintf(
"--%s is required\n hint: example: 2026-03-10T14:00:00+08:00", flagName))
}
loc, _ := time.LoadLocation("Asia/Shanghai")
if loc == nil {
loc = time.Local
}
for _, layout := range flexTimeLayouts {
t, err := time.ParseInLocation(layout, value, loc)
if err == nil {
return t.UnixMilli(), nil
}
}
return 0, apperrors.NewValidation(fmt.Sprintf(
"cannot parse time for --%s (input: %q)\n hint: supported formats: 2026-03-23T14:00:00+08:00, 2026-03-23 14:00:00, 2026-03-23",
flagName, value))
}
// validateTimeRange checks that endMs is strictly after startMs.
func validateTimeRange(startMs, endMs int64) error {
if endMs <= startMs {
return apperrors.NewValidation("--end must be after --start\n hint: swap the values or adjust the time range")
}
return nil
}
// ── template list ──────────────────────────────────────────
func newReportTemplateListCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "list",
Short: "获取当前用户可用的日志模版列表",
Example: " dws report template list",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
params := map[string]any{}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
preferLegacyLeaf(cmd)
return cmd
}
// ── template detail ────────────────────────────────────────
func newReportTemplateDetailCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "detail",
Short: "获取日志模版详情",
Example: " dws report template detail --name <templateName>",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
name, _ := cmd.Flags().GetString("name")
if name == "" {
return apperrors.NewValidation("--name is required")
}
params := map[string]any{
"report_template_name": name,
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().String("name", "", "模版名称 (必填)")
preferLegacyLeaf(cmd)
return cmd
}
// ── create ─────────────────────────────────────────────────
func newReportCreateCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "create",
Short: "创建日志",
Long: `按模版创建一条日志。--contents 为 JSON 数组,每项需含 key、sort、content、contentType、type,
与远程 create_report 一致;可先通过 report template list / template detail 取得 templateId 与控件定义。`,
Example: ` dws report create --template-id TPL_ID --contents '[{"content":"完成开发","sort":"0","key":"今日完成","contentType":"markdown","type":"1"}]'
dws report create --template-id TPL_ID --contents '[...]' --to-chat --to-user-ids userId1,userId2`,
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
tplID, _ := cmd.Flags().GetString("template-id")
if tplID == "" {
return apperrors.NewValidation("--template-id is required")
}
contentsJSON, _ := cmd.Flags().GetString("contents")
if contentsJSON == "" {
return apperrors.NewValidation("--contents is required")
}
var contents []map[string]any
if err := json.Unmarshal([]byte(contentsJSON), &contents); err != nil {
return apperrors.NewValidation(fmt.Sprintf("--contents JSON parse failed: %v", err))
}
ddFrom, _ := cmd.Flags().GetString("dd-from")
if ddFrom == "" {
ddFrom = "dws"
}
toChat, _ := cmd.Flags().GetBool("to-chat")
params := map[string]any{
"templateId": tplID,
"contents": contents,
"ddFrom": ddFrom,
"toChat": toChat,
}
if v, _ := cmd.Flags().GetString("to-user-ids"); v != "" {
params["toUserIds"] = parseUserIDs(v)
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().String("template-id", "", "日志模版 ID (必填)")
cmd.Flags().String("contents", "", "日志内容 JSON 数组 (必填),每项含 key/sort/content/contentType/type")
cmd.Flags().String("dd-from", "dws", "创建来源标识")
cmd.Flags().Bool("to-chat", false, "是否发送到日志接收人单聊")
cmd.Flags().String("to-user-ids", "", "接收人 userId,逗号分隔 (可选)")
preferLegacyLeaf(cmd)
return cmd
}
// ── detail ─────────────────────────────────────────────────
func newReportDetailCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "detail",
Short: "获取日志详情",
Example: " dws report detail --report-id <reportId>",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
reportID, _ := cmd.Flags().GetString("report-id")
if reportID == "" {
return apperrors.NewValidation("--report-id is required")
}
params := map[string]any{
"report_id": reportID,
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().String("report-id", "", "日志 ID (必填)")
preferLegacyLeaf(cmd)
return cmd
}
// ── list (received reports) ────────────────────────────────
// Key fix: cursor defaults to 0, size defaults to 20, flexible date parsing
func newReportListCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "list",
Short: "查询当前人收到的日志列表",
Example: ` dws report list --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
dws report list --start "2026-03-10 00:00:00" --end "2026-03-10 23:59:59" --cursor 0 --size 20`,
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
startStr, _ := cmd.Flags().GetString("start")
endStr, _ := cmd.Flags().GetString("end")
startMs, err := parseFlexTimeToMillis("start", startStr)
if err != nil {
return err
}
endMs, err := parseFlexTimeToMillis("end", endStr)
if err != nil {
return err
}
if err := validateTimeRange(startMs, endMs); err != nil {
return err
}
// cursor defaults to 0, size defaults to 20
cursor, _ := cmd.Flags().GetInt("cursor")
size, _ := cmd.Flags().GetInt("size")
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
size = v
}
params := map[string]any{
"startTime": float64(startMs),
"endTime": float64(endMs),
"cursor": float64(cursor),
"size": float64(size),
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().String("start", "", "开始时间 ISO-8601 (如 2026-03-10T00:00:00+08:00) (必填)")
cmd.Flags().String("end", "", "结束时间 ISO-8601 (如 2026-03-10T23:59:59+08:00) (必填)")
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
cmd.Flags().Int("limit", 0, "--size 的别名")
_ = cmd.Flags().MarkHidden("limit")
preferLegacyLeaf(cmd)
return cmd
}
// ── stats ──────────────────────────────────────────────────
func newReportStatsCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "stats",
Short: "获取日志统计数据",
Example: " dws report stats --report-id <reportId>",
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
reportID, _ := cmd.Flags().GetString("report-id")
if reportID == "" {
return apperrors.NewValidation("--report-id is required")
}
params := map[string]any{
"report_id": reportID,
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().String("report-id", "", "日志 ID (必填)")
preferLegacyLeaf(cmd)
return cmd
}
// ── sent (my created reports) ──────────────────────────────
// Key fix: cursor defaults to 0, size defaults to 20, start/end default to last 30 days
func newReportSentCommand(runner executor.Runner) *cobra.Command {
cmd := &cobra.Command{
Use: "sent",
Short: "查询当前人创建的日志列表",
Example: ` dws report sent
dws report sent --cursor 0 --size 20
dws report sent --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
dws report sent --template-name "日报"`,
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
// cursor defaults to 0, size defaults to 20
cursor, _ := cmd.Flags().GetInt("cursor")
size, _ := cmd.Flags().GetInt("size")
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
size = v
}
params := map[string]any{
"cursor": float64(cursor),
"size": float64(size),
}
// Default time range: last 30 days
now := time.Now()
startDefault := now.AddDate(0, 0, -30).Truncate(24 * time.Hour).Format(time.RFC3339)
endDefault := time.Date(now.Year(), now.Month(), now.Day(), 23, 59, 59, 0, now.Location()).Format(time.RFC3339)
startStr, _ := cmd.Flags().GetString("start")
if startStr == "" {
startStr = startDefault
}
endStr, _ := cmd.Flags().GetString("end")
if endStr == "" {
endStr = endDefault
}
startMs, err := parseFlexTimeToMillis("start", startStr)
if err != nil {
return err
}
params["startTime"] = float64(startMs)
endMs, err := parseFlexTimeToMillis("end", endStr)
if err != nil {
return err
}
params["endTime"] = float64(endMs)
if err := validateTimeRange(startMs, endMs); err != nil {
return err
}
// Optional modified time filters
if v, _ := cmd.Flags().GetString("modified-start"); v != "" {
ms, err := parseFlexTimeToMillis("modified-start", v)
if err != nil {
return err
}
params["modifiedStartTime"] = float64(ms)
}
if v, _ := cmd.Flags().GetString("modified-end"); v != "" {
ms, err := parseFlexTimeToMillis("modified-end", v)
if err != nil {
return err
}
params["modifiedEndTime"] = float64(ms)
}
// Optional template name filter
if v, _ := cmd.Flags().GetString("template-name"); v != "" {
params["report_template_name"] = v
}
if commandDryRun(cmd) {
return writeCommandPayload(cmd, executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
))
}
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
))
if err != nil {
return err
}
return writeCommandPayload(cmd, result)
},
}
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
cmd.Flags().Int("limit", 0, "--size 的别名")
_ = cmd.Flags().MarkHidden("limit")
cmd.Flags().String("start", "", "创建开始时间 ISO-8601 (默认最近 30 天)")
cmd.Flags().String("end", "", "创建结束时间 ISO-8601 (默认最近 30 天)")
cmd.Flags().String("modified-start", "", "修改开始时间 ISO-8601 (可选)")
cmd.Flags().String("modified-end", "", "修改结束时间 ISO-8601 (可选)")
cmd.Flags().String("template-name", "", "日志模板名称 (可选,不传查全部)")
preferLegacyLeaf(cmd)
return cmd
}
// ── helpers ────────────────────────────────────────────────
func parseUserIDs(s string) []string {
parts := strings.Split(s, ",")
out := make([]string, 0, len(parts))
for _, p := range parts {
p = strings.TrimSpace(p)
if p != "" {
out = append(out, p)
}
}
return out
}
+149
View File
@@ -0,0 +1,149 @@
// 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 (
"testing"
"time"
)
func TestParseFlexTimeToMillis_RFC3339(t *testing.T) {
t.Parallel()
ms, err := parseFlexTimeToMillis("start", "2026-03-10T00:00:00+08:00")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_SpaceSeparated(t *testing.T) {
t.Parallel()
// This is the format that was causing the HTTP 400 error
ms, err := parseFlexTimeToMillis("start", "2026-03-01 00:00:00")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_DateOnly(t *testing.T) {
t.Parallel()
ms, err := parseFlexTimeToMillis("start", "2026-03-01")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_NoTimezone(t *testing.T) {
t.Parallel()
ms, err := parseFlexTimeToMillis("start", "2026-03-10T14:00:00")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_SlashFormat(t *testing.T) {
t.Parallel()
ms, err := parseFlexTimeToMillis("start", "2026/03/10")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_CompactFormat(t *testing.T) {
t.Parallel()
ms, err := parseFlexTimeToMillis("start", "20260310")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if ms <= 0 {
t.Errorf("expected positive milliseconds, got %d", ms)
}
}
func TestParseFlexTimeToMillis_Empty(t *testing.T) {
t.Parallel()
_, err := parseFlexTimeToMillis("start", "")
if err == nil {
t.Fatal("expected error for empty value")
}
}
func TestParseFlexTimeToMillis_Invalid(t *testing.T) {
t.Parallel()
_, err := parseFlexTimeToMillis("start", "not-a-date")
if err == nil {
t.Fatal("expected error for invalid date")
}
}
func TestValidateTimeRange_Valid(t *testing.T) {
t.Parallel()
start := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
end := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
if err := validateTimeRange(start, end); err != nil {
t.Fatalf("unexpected error: %v", err)
}
}
func TestValidateTimeRange_Invalid(t *testing.T) {
t.Parallel()
start := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
end := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
if err := validateTimeRange(start, end); err == nil {
t.Fatal("expected error when end is before start")
}
}
func TestValidateTimeRange_Equal(t *testing.T) {
t.Parallel()
ts := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
if err := validateTimeRange(ts, ts); err == nil {
t.Fatal("expected error when start equals end")
}
}
func TestParseUserIDs(t *testing.T) {
t.Parallel()
tests := []struct {
input string
want int
}{
{"user1,user2,user3", 3},
{"user1", 1},
{"user1, user2, user3", 3},
{"user1,,user2", 2},
{"", 0},
{" , , ", 0},
}
for _, tt := range tests {
got := parseUserIDs(tt.input)
if len(got) != tt.want {
t.Errorf("parseUserIDs(%q) = %d items, want %d", tt.input, len(got), tt.want)
}
}
}
+4 -34
View File
@@ -14,15 +14,14 @@
package helpers
import (
"fmt"
"strconv"
"strings"
"time"
"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/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
"github.com/spf13/cobra"
)
@@ -88,7 +87,7 @@ func newTodoTaskCreateCommand(runner executor.Runner) *cobra.Command {
Args: cobra.NoArgs,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
title := flagOrFallback(cmd, "title", "subject", "content")
title := cmdutil.FlagOrFallback(cmd, "title", "subject", "content")
if strings.TrimSpace(title) == "" {
return apperrors.NewValidation("--title is required")
}
@@ -103,7 +102,7 @@ func newTodoTaskCreateCommand(runner executor.Runner) *cobra.Command {
"executorIds": executorIds,
}
if v, _ := cmd.Flags().GetString("due"); v != "" {
ms, err := parseISOTimeToMillis("due", v)
ms, err := cmdutil.ParseISOTimeToMillis("due", v)
if err != nil {
return err
}
@@ -270,7 +269,7 @@ func newTodoTaskUpdateCommand(runner executor.Runner) *cobra.Command {
inner["subject"] = v
}
if v, _ := cmd.Flags().GetString("due"); v != "" {
ms, err := parseISOTimeToMillis("due", v)
ms, err := cmdutil.ParseISOTimeToMillis("due", v)
if err != nil {
return err
}
@@ -444,16 +443,6 @@ func newTodoTaskDeleteCommand(runner executor.Runner) *cobra.Command {
// ── helpers ────────────────────────────────────────────────
// flagOrFallback returns the first non-empty value among the given flag names.
func flagOrFallback(cmd *cobra.Command, names ...string) string {
for _, name := range names {
if v, _ := cmd.Flags().GetString(name); strings.TrimSpace(v) != "" {
return v
}
}
return ""
}
// parseExecutorIds splits "id1,id2" into []string for the MCP executorIds array.
func parseExecutorIds(s string) []string {
s = strings.TrimSpace(s)
@@ -470,25 +459,6 @@ func parseExecutorIds(s string) []string {
return ids
}
// parseISOTimeToMillis parses an ISO-8601 datetime string and returns Unix
// milliseconds. It supports timezone offsets (e.g. +08:00) and UTC "Z" suffix.
func parseISOTimeToMillis(flagName, value string) (int64, error) {
formats := []string{
time.RFC3339,
"2006-01-02T15:04:05Z07:00",
"2006-01-02T15:04:05",
"2006-01-02 15:04:05",
}
for _, layout := range formats {
if t, err := time.Parse(layout, value); err == nil {
return t.UnixMilli(), nil
}
}
return 0, apperrors.NewValidation(
fmt.Sprintf("--%s format error, use ISO-8601 e.g. 2026-03-10T18:00:00+08:00", flagName),
)
}
// ── list pagination helpers ────────────────────────────────
func normalizePage(raw string) string {
+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
+15
View File
@@ -155,5 +155,20 @@
"返回数据缺少 uploadUrl 或 fileToken": "Response data missing uploadUrl or fileToken",
"附件工作流": "Attachment workflow",
"页码 (必填)": "page number (required)",
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ Unable to verify CLI data access permission status",
" 请检查网络连接后重试。": " Please check your network connection and retry.",
"检查 CLI 授权状态失败": "Failed to check CLI auth status",
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ CLI data access is not enabled for this organization",
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " The organization admin has not enabled \"Allow members to access their personal data via CLI\".",
" 组织主管理员:": " Organization super admins: ",
" 请联系组织主管理员开启后重新登录。": " Please contact the organization super admin to enable it and re-login.",
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " Admin settings: https://open-dev.dingtalk.com/fe/old#/developerSettings",
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "CLI data access is not enabled for this organization, please contact admin to enable it",
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ CLI data access is not enabled for this organization, please submit an authorization request in the browser...",
"✅ 权限已开启,继续登录...": "✅ Permission enabled, continuing login...",
"等待管理员审批中": "Waiting for admin approval",
"等待提交申请中": "Waiting to submit request",
"操作超时,请重新登录": "Operation timed out, please re-login",
"检查组织 CLI 授权状态...": "Checking organization CLI auth status...",
"🔐 登录钉钉": "🔐 Login to DingTalk"
}
+15
View File
@@ -155,5 +155,20 @@
"返回数据缺少 uploadUrl 或 fileToken": "返回数据缺少 uploadUrl 或 fileToken",
"附件工作流": "附件工作流",
"页码 (必填)": "页码 (必填)",
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ 无法检查 CLI 数据访问权限状态",
" 请检查网络连接后重试。": " 请检查网络连接后重试。",
"检查 CLI 授权状态失败": "检查 CLI 授权状态失败",
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ 该组织尚未开启 CLI 数据访问权限",
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。",
" 组织主管理员:": " 组织主管理员:",
" 请联系组织主管理员开启后重新登录。": " 请联系组织主管理员开启后重新登录。",
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings",
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "该组织尚未开启 CLI 数据访问权限,请联系管理员开启",
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...",
"✅ 权限已开启,继续登录...": "✅ 权限已开启,继续登录...",
"等待管理员审批中": "等待管理员审批中",
"等待提交申请中": "等待提交申请中",
"操作超时,请重新登录": "操作超时,请重新登录",
"检查组织 CLI 授权状态...": "检查组织 CLI 授权状态...",
"🔐 登录钉钉": "🔐 登录钉钉"
}
+1 -1
View File
@@ -24,7 +24,7 @@ import (
"strconv"
"sync"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
+89 -2
View File
@@ -13,7 +13,14 @@
package logging
import "strings"
import (
"encoding/json"
"fmt"
"log/slog"
"net/http"
"strings"
"unicode/utf8"
)
// sensitiveKeys are header/field names whose values must be redacted in logs.
var sensitiveKeys = map[string]bool{
@@ -25,12 +32,30 @@ var sensitiveKeys = map[string]bool{
"secret": true,
"password": true,
"cookie": true,
"api_key": true,
"api-key": true,
"access_token": true,
"credential": true,
}
// sensitiveSubstrings are substrings that mark a key as sensitive.
var sensitiveSubstrings = []string{
"password", "secret", "token", "credential",
}
// IsSensitiveKey returns true if the key (case-insensitive) refers to a
// credential or secret that must not appear in log files.
func IsSensitiveKey(key string) bool {
return sensitiveKeys[strings.ToLower(key)]
lower := strings.ToLower(key)
if sensitiveKeys[lower] {
return true
}
for _, sub := range sensitiveSubstrings {
if strings.Contains(lower, sub) {
return true
}
}
return false
}
// RedactValue replaces a sensitive value with a safe placeholder.
@@ -42,3 +67,65 @@ func RedactValue(value string) string {
}
return value[:4] + "***"
}
// TruncateBody returns the body truncated to maxBytes with a UTF-8 safe
// boundary. If truncated, appends a marker showing the original size.
func TruncateBody(body []byte, maxBytes int) string {
if len(body) <= maxBytes {
return string(body)
}
safe := body[:maxBytes]
// Walk back to a valid UTF-8 boundary.
for len(safe) > 0 && !utf8.Valid(safe) {
safe = safe[:len(safe)-1]
}
return fmt.Sprintf("%s...(truncated, total=%d bytes)", string(safe), len(body))
}
// SanitizeArguments returns a JSON string of the arguments map with
// sensitive-looking values replaced by "***". Truncates to maxBytes.
func SanitizeArguments(args map[string]any, maxBytes int) string {
if len(args) == 0 {
return "{}"
}
sanitized := make(map[string]any, len(args))
for k, v := range args {
sanitized[k] = v
}
redactMapValues(sanitized)
data, err := json.Marshal(sanitized)
if err != nil {
return "{}"
}
return TruncateBody(data, maxBytes)
}
// redactMapValues replaces values of sensitive keys with "***" in-place.
func redactMapValues(m map[string]any) {
for k, v := range m {
if IsSensitiveKey(k) {
m[k] = "***"
continue
}
if nested, ok := v.(map[string]any); ok {
redactMapValues(nested)
}
}
}
// RedactHeaders returns slog attributes for HTTP headers with sensitive
// values redacted.
func RedactHeaders(headers http.Header) []slog.Attr {
if len(headers) == 0 {
return nil
}
attrs := make([]slog.Attr, 0, len(headers))
for key := range headers {
value := headers.Get(key)
if IsSensitiveKey(key) {
value = RedactValue(value)
}
attrs = append(attrs, slog.String("header."+strings.ToLower(key), value))
}
return attrs
}
+99 -1
View File
@@ -13,7 +13,11 @@
package logging
import "testing"
import (
"net/http"
"strings"
"testing"
)
func TestIsSensitiveKey(t *testing.T) {
t.Parallel()
@@ -69,3 +73,97 @@ func TestRedactValue(t *testing.T) {
})
}
}
func TestTruncateBody(t *testing.T) {
t.Parallel()
short := []byte("hello")
if got := TruncateBody(short, 100); got != "hello" {
t.Fatalf("expected no truncation, got %q", got)
}
long := []byte(strings.Repeat("a", 200))
got := TruncateBody(long, 50)
if !strings.Contains(got, "truncated") {
t.Fatalf("expected truncation marker, got %q", got)
}
if !strings.Contains(got, "total=200") {
t.Fatalf("expected total size, got %q", got)
}
}
func TestTruncateBody_Empty(t *testing.T) {
t.Parallel()
if got := TruncateBody(nil, 100); got != "" {
t.Fatalf("expected empty, got %q", got)
}
}
func TestSanitizeArguments(t *testing.T) {
t.Parallel()
args := map[string]any{
"name": "test",
"password": "secret123",
"nested": map[string]any{
"api_key": "key-value",
"safe": "ok",
},
}
got := SanitizeArguments(args, 4096)
if strings.Contains(got, "secret123") {
t.Fatalf("password should be redacted: %s", got)
}
if strings.Contains(got, "key-value") {
t.Fatalf("api_key should be redacted: %s", got)
}
if !strings.Contains(got, "test") {
t.Fatalf("non-sensitive value should remain: %s", got)
}
}
func TestSanitizeArguments_Empty(t *testing.T) {
t.Parallel()
if got := SanitizeArguments(nil, 100); got != "{}" {
t.Fatalf("expected {}, got %q", got)
}
}
func TestRedactHeaders(t *testing.T) {
t.Parallel()
headers := http.Header{
"Authorization": {"Bearer token123456"},
"Content-Type": {"application/json"},
}
attrs := RedactHeaders(headers)
if len(attrs) != 2 {
t.Fatalf("expected 2 attrs, got %d", len(attrs))
}
for _, attr := range attrs {
if attr.Key == "header.authorization" && !strings.Contains(attr.Value.String(), "***") {
t.Fatalf("authorization should be redacted: %s", attr.Value.String())
}
if attr.Key == "header.content-type" && attr.Value.String() != "application/json" {
t.Fatalf("content-type should not be redacted: %s", attr.Value.String())
}
}
}
func TestIsSensitiveKey_Substrings(t *testing.T) {
t.Parallel()
tests := []struct {
key string
want bool
}{
{"x-api-token", true},
{"user_password_hash", true},
{"my_secret_key", true},
{"x-credential-id", true},
{"safe-header", false},
}
for _, tt := range tests {
t.Run(tt.key, func(t *testing.T) {
t.Parallel()
if got := IsSensitiveKey(tt.key); got != tt.want {
t.Errorf("IsSensitiveKey(%q) = %v, want %v", tt.key, got, tt.want)
}
})
}
}
+128 -2
View File
@@ -16,9 +16,17 @@ package logging
import (
"context"
"log/slog"
"runtime"
"time"
)
const (
// maxBodyLogSize is the maximum bytes of request/response body to log.
maxBodyLogSize = 4096
// maxArgLogSize is the maximum bytes for sanitized argument summaries.
maxArgLogSize = 1024
)
// LogRequest logs a JSON-RPC request at Debug level.
func LogRequest(logger *slog.Logger, method, endpoint, executionId string, bodySize int) {
if logger == nil {
@@ -32,14 +40,28 @@ func LogRequest(logger *slog.Logger, method, endpoint, executionId string, bodyS
)
}
// 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) {
// LogRequestBody logs a truncated, redacted request body for tools/call.
func LogRequestBody(logger *slog.Logger, method, executionId string, toolName string, arguments map[string]any) {
if logger == nil || method != "tools/call" {
return
}
logger.Debug("jsonrpc_request_body",
slog.String("method", method),
slog.String("execution_id", executionId),
slog.String("tool_name", toolName),
slog.String("arguments_summary", SanitizeArguments(arguments, maxArgLogSize)),
)
}
// 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()),
@@ -52,6 +74,110 @@ func LogResponse(logger *slog.Logger, method, endpoint string, statusCode int, r
logger.LogAttrs(context.TODO(), slog.LevelDebug, "jsonrpc_response", attrs...)
}
// LogResponseBody logs a truncated response body on error paths.
func LogResponseBody(logger *slog.Logger, method, executionId string, statusCode int, body []byte, traceID string) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("method", method),
slog.String("execution_id", executionId),
slog.Int("status", statusCode),
slog.String("body", TruncateBody(body, maxBodyLogSize)),
}
if traceID != "" {
attrs = append(attrs, slog.String("trace_id", traceID))
}
level := slog.LevelDebug
if statusCode >= 400 {
level = slog.LevelWarn
}
logger.LogAttrs(context.TODO(), level, "jsonrpc_response_body", attrs...)
}
// LogRetryAttempt logs a retry attempt at Warn level.
func LogRetryAttempt(logger *slog.Logger, method, executionId string, attempt, maxRetries int, statusCode int, delay time.Duration, lastErr error) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("method", method),
slog.String("execution_id", executionId),
slog.Int("attempt", attempt+1),
slog.Int("max_attempts", maxRetries+1),
slog.Int("status", statusCode),
slog.String("delay", delay.String()),
}
if lastErr != nil {
attrs = append(attrs, slog.String("error", lastErr.Error()))
}
logger.LogAttrs(context.TODO(), slog.LevelWarn, "jsonrpc_retry", attrs...)
}
// LogErrorClassified logs the final error classification at Warn level.
func LogErrorClassified(logger *slog.Logger, method, executionId, category, reason string, httpStatus, rpcCode int, retryable bool, traceID string) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("method", method),
slog.String("execution_id", executionId),
slog.String("category", category),
slog.String("reason", reason),
slog.Bool("retryable", retryable),
}
if httpStatus != 0 {
attrs = append(attrs, slog.Int("http_status", httpStatus))
}
if rpcCode != 0 {
attrs = append(attrs, slog.Int("rpc_code", rpcCode))
}
if traceID != "" {
attrs = append(attrs, slog.String("trace_id", traceID))
}
logger.LogAttrs(context.TODO(), slog.LevelWarn, "error_classified", attrs...)
}
// LogCommandStart logs the beginning of a command execution.
func LogCommandStart(logger *slog.Logger, executionId, product, tool, endpoint, version string, authPresent bool, timeoutSec int) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("execution_id", executionId),
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.
func LogCommandEnd(logger *slog.Logger, executionId, product, tool string, success bool, duration time.Duration, errCategory, errReason string) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("execution_id", executionId),
slog.String("product", product),
slog.String("tool", tool),
slog.Bool("success", success),
slog.String("duration", duration.Truncate(time.Millisecond).String()),
}
if !success {
attrs = append(attrs, slog.String("error_category", errCategory))
attrs = append(attrs, slog.String("error_reason", errReason))
}
logger.LogAttrs(context.TODO(), slog.LevelInfo, "command_end", attrs...)
}
// redactEndpoint removes query parameters from endpoint URLs in logs.
func redactEndpoint(endpoint string) string {
for i := 0; i < len(endpoint); i++ {
+91 -3
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,7 +87,95 @@ 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", "doc", "list", "https://mcp.example.com", "1.0.0", false, 0)
LogCommandEnd(nil, "exec-1", "doc", "list", true, 0, "", "")
}
func TestLogRequestBody(t *testing.T) {
t.Parallel()
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
args := map[string]any{"name": "test", "limit": 10}
LogRequestBody(logger, "tools/call", "exec-1", "doc.list", args)
out := buf.String()
if !strings.Contains(out, "jsonrpc_request_body") {
t.Error("missing message")
}
if !strings.Contains(out, "doc.list") {
t.Error("missing tool_name")
}
if !strings.Contains(out, "exec-1") {
t.Error("missing execution_id")
}
}
func TestLogRequestBody_SkipsNonToolsCall(t *testing.T) {
t.Parallel()
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogRequestBody(logger, "initialize", "exec-1", "", nil)
if buf.Len() != 0 {
t.Error("should not log body for non-tools/call methods")
}
}
func TestLogResponseBody(t *testing.T) {
t.Parallel()
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogResponseBody(logger, "tools/call", "exec-1", 500, []byte(`{"error":"fail"}`), "trace-abc")
out := buf.String()
if !strings.Contains(out, "jsonrpc_response_body") {
t.Error("missing message")
}
if !strings.Contains(out, "trace-abc") {
t.Error("missing trace_id")
}
if !strings.Contains(out, "WARN") {
t.Error("expected WARN level for 500 status")
}
}
func TestLogRetryAttempt(t *testing.T) {
t.Parallel()
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogRetryAttempt(logger, "tools/call", "exec-1", 0, 2, 429, 10*time.Millisecond, errors.New("rate limited"))
out := buf.String()
if !strings.Contains(out, "jsonrpc_retry") {
t.Error("missing message")
}
if !strings.Contains(out, "rate limited") {
t.Error("missing error")
}
}
func TestLogErrorClassified(t *testing.T) {
t.Parallel()
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogErrorClassified(logger, "tools/call", "exec-1", "auth", "http_401", 401, 0, false, "trace-xyz")
out := buf.String()
if !strings.Contains(out, "error_classified") {
t.Error("missing message")
}
if !strings.Contains(out, "trace-xyz") {
t.Error("missing trace_id")
}
}
func TestRedactEndpoint(t *testing.T) {
+19 -17
View File
@@ -26,8 +26,8 @@ import (
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
const (
@@ -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)
}
})
}
+1 -1
View File
@@ -13,7 +13,7 @@
package output
import "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/validate"
import "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/validate"
// SanitizeForTerminal strips ANSI escape sequences, control characters, and
// dangerous Unicode from text before it is printed to a terminal.
+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) {
+2 -46
View File
@@ -17,6 +17,7 @@ import (
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
)
// ParamNameHandler performs fuzzy correction on flag names that are
@@ -101,7 +102,7 @@ func tryFuzzyMatch(arg string, known map[string]bool, candidates []string) (stri
ambiguous := false
for _, candidate := range candidates {
dist := levenshtein(bare, candidate)
dist := cmdutil.LevenshteinDist(bare, candidate)
if dist < bestDist {
bestDist = dist
bestMatch = candidate
@@ -117,48 +118,3 @@ func tryFuzzyMatch(arg string, known map[string]bool, candidates []string) (stri
return "--" + bestMatch + suffix, true
}
// levenshtein computes the edit distance between two strings using
// the standard dynamic programming approach with O(min(m,n)) space.
func levenshtein(a, b string) int {
if a == b {
return 0
}
la, lb := len(a), len(b)
if la == 0 {
return lb
}
if lb == 0 {
return la
}
// Ensure a is the shorter string for O(min) space.
if la > lb {
a, b = b, a
la, lb = lb, la
}
prev := make([]int, la+1)
curr := make([]int, la+1)
for i := range prev {
prev[i] = i
}
for j := 1; j <= lb; j++ {
curr[0] = j
for i := 1; i <= la; i++ {
cost := 1
if a[i-1] == b[j-1] {
cost = 0
}
curr[i] = min(
prev[i]+1, // deletion
curr[i-1]+1, // insertion
prev[i-1]+cost, // substitution
)
}
prev, curr = curr, prev
}
return prev[la]
}
+4 -4
View File
@@ -18,6 +18,7 @@ import (
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
)
func TestLevenshtein(t *testing.T) {
@@ -40,12 +41,11 @@ func TestLevenshtein(t *testing.T) {
for _, tt := range tests {
t.Run(tt.a+"→"+tt.b, func(t *testing.T) {
got := levenshtein(tt.a, tt.b)
got := cmdutil.LevenshteinDist(tt.a, tt.b)
if got != tt.want {
t.Errorf("levenshtein(%q, %q) = %d, want %d", tt.a, tt.b, got, tt.want)
t.Errorf("LevenshteinDist(%q, %q) = %d, want %d", tt.a, tt.b, got, tt.want)
}
// Verify symmetry.
gotRev := levenshtein(tt.b, tt.a)
gotRev := cmdutil.LevenshteinDist(tt.b, tt.a)
if gotRev != got {
t.Errorf("asymmetric: (%q,%q)=%d but (%q,%q)=%d", tt.a, tt.b, got, tt.b, tt.a, gotRev)
}
@@ -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)
}
}

Some files were not shown because too many files have changed in this diff Show More