Compare commits

...
136 Commits
Author SHA1 Message Date
383aeefaf6 fix(cli): localize plugin/help command strings via i18n (closes #118) (#134)
Wire plugin module + help command + OAuth client-id/secret flags through
the existing i18n catalog so --help is consistent Chinese under zh locale
and English under en locale, instead of mixing the two.

- internal/app/plugin_cmd.go: wrap 16 Short strings with i18n.T
- internal/app/flags.go: wrap --client-id/--client-secret descriptions
- internal/app/root_help.go: override cobra default help command with a
  localized one so "dws help --help" and the utility-commands listing
  share the same catalog
- internal/i18n/locales/{en,zh}.json: add 20 new catalog entries

Verified under DWS_LANG=zh (all Chinese), DWS_LANG=en (original English
preserved), and LANG-based auto-routing. go test ./internal/app/... pass.

Co-authored-by: 修雨 <huyizhou.hyz@alibaba-inc.com>
Co-authored-by: Claude Opus 4.7 <noreply@anthropic.com>
2026-04-20 20:51:07 +08:00
ybcstudyand猷诺 a5bede3a19 fix: -f json模式下错误JSON从stdout改为输出到stderr (#133)
修复背景:
CI测试脚本通过检查stderr中是否包含错误关键词来判断命令是否失败。
但-f json模式下,printExecutionError将错误JSON输出到stdout,
导致stderr为空,测试断言失败。

改动:
- root.go:150: apperrors.PrintJSON(stdout, err) → apperrors.PrintJSON(stderr, err)
- root_execute_test.go: 3处断言方向同步调整(stdout↔stderr)

影响范围:
仅影响-f json模式下的错误输出流向,正常输出不受影响。
符合Unix惯例:错误信息输出到stderr,正常结果输出到stdout。

Co-authored-by: 猷诺 <bicheng.ybc@alibaba-inc.com>
2026-04-20 19:38:39 +08:00
xuan bbf66e23d6 fix(pat): avoid shared PAT command state in root registration (#129)
* chore: remove workspace and bundled artifacts

* chore: clean local-only repository artifacts

* fix(pat): avoid shared command state in registration
2026-04-20 14:54:57 +08:00
shangguanxuan633-lab 5e168c92cf Codex/remove red box files (#127)
* chore: remove workspace and bundled artifacts

* chore: clean local-only repository artifacts
2026-04-20 09:46:11 +08:00
shangguanxuan633-labandshangguanxuan.sgx 725577103d feat: PAT scope error visualization and auto-retry with authorization polling (#113)
* feat: PAT scope error visualization and auto-retry with authorization polling

1. PAT error result visualization (non-JSON)
   - Human-readable error output matching lark-cli style
   - Shows error type, message, hint, and authorization command
   - JSON output also available via --format json

2. Auto-retry with polling after authorization wall
   - Detects missing_scope / insufficient_scope / permission errors
   - Polls every 5s for token update after user completes auth
   - 10 minute timeout before giving up
   - Automatically retries the original command after successful auth

Changes:
- internal/app/pat_auth_retry.go: Core PAT error handling and retry logic
- internal/app/pat_auth_retry_test.go: 11 unit tests
- internal/app/runner.go: Integration with executeInvocation error path
- internal/auth/device_flow.go: Add SetScope() method

* feat: migrate PAT (Permission Authorization) infrastructure to open-source core

Move PAT permission management code from dws-wukong overlay to the
open-source dingtalk-workspace-cli core. This enables PAT handling
in the public distribution while keeping dws-wukong fully compatible
through the edition.Hooks mechanism.

Changes:
- pkg/edition: add 6 new Hooks fields (AuthClientID, AuthClientFromMCP,
  SaveToken, LoadToken, DeleteToken, ClassifyToolResult) for overlay
  extensions to inject custom auth/token/error behaviour
- internal/errors/pat: new PATError type with ExitCode()=4 and
  RawStderr() for embedded-mode passthrough, plus ClassifyToolResultContent
  and ClassifyMCPResponseText classification functions
- internal/pat: new command group (pat chmod) with tool result handling
- internal/app/root: register PAT commands before RegisterExtraCommands hook
- internal/app/runner: integrate ClassifyToolResult hook in the
  callResult.IsError branch so overlays can intercept PAT/gateway errors
  before generic handling

The dws-wukong overlay remains unchanged — its existing pat/ package and
RegisterExtraCommands hook continue to work. Through the go.mod replace
directive and deduplicateCommands mechanism, both codebases coexist
without conflict.

* feat(pat): add PAT auth check with device flow polling and auto-retry

- Add ClassifyPatAuthCheck / AsPatAuthCheckError in pat.go for
  AGENT_CODE_NOT_EXISTS error detection
- Add handlePatAuthCheck in pat_auth_retry.go: inject clientId as
  x-robot-uid header, poll device flow endpoint, auto-retry on APPROVED
- Integrate PAT check into runner.go executeInvocation pipeline
  (edition hook + open-source fallback)
- Add patAuthRequiredCodes map for extensible auth-required code matching
- Remove test_parse.go (temporary test script)

* fix(pat): update poll path to /cli/oauth/device/poll and add x-user-access-token header

- Change DevicePollPath from /api/dingtalk-workspace-cli/oauth/device/poll
  to /cli/oauth/device/poll (aligning with backend endpoint convention)
- Add x-user-access-token header to poll requests (loaded from stored token)

* feat(auth): PAT device flow improvements and flowId empty guard

- pat_auth_retry: fix JSON field parsing (result -> data) using DevicePollResponse
- pat_auth_retry: fix elapsed time calculation based on start timestamp
- device_flow: add flowId empty guard - skip polling and show auth URL for manual handling
- device_flow: unify terminalBaseURL to GetMCPBaseURL for polling endpoint
- device_flow: remove unused PollFlowApproval dead code
- runner: integrate DWS_CLIENT_ID env and app.json persistence
- oauth_provider: add x-robot-uid header removal on logout
- client: add request/response logging for transport debugging

* feat(env): switch all service discovery URLs from pre-release to production

- endpoints: pre-login/pre-api/pre-mcp/pre-open-dev -> login/api/mcp/open-dev
- loader: DefaultMarketBaseURL -> mcp.dingtalk.com
- registry: defaultBaseURL -> mcp.dingtalk.com
- constants: DefaultTerminalBaseURL -> open-dev.dingtalk.com

* fix: remove residual merge conflict markers and duplicate Hooks fields in edition.go; unify developer-settings URL via config helper in device_flow.go

* chore: remove lark-cli reference from PrintPatAuthError comment

* test: add comprehensive unit tests for PAT auth modules

- Create internal/errors/pat_test.go (38 tests): covers ClassifyToolResultContent,
  ClassifyMCPResponseText, ClassifyPatAuthCheck, AsPatAuthCheckError, cleanPATJSON,
  stripClassFields, getDWSGatewayErrorCode, isNotLoggedInError, isBusinessError,
  IsPATError, IsPATNoPermissionCode, suggestForBusinessErrorText
- Supplement pat_auth_retry_test.go (+7 tests): IsPatRetrying context checks,
  pollPatDeviceFlow edge cases (server error fallback, SSO redirect skip),
  extractPatScopeError nil/identity extraction

* feat(pat): exchange authCode for fresh token after PAT APPROVED

Previously pollPatDeviceFlow discarded the authCode from the poll
response, so handlePatAuthCheck retried with the stale token after
APPROVED. This mirrors device_flow.go loginOnce which calls
exchangeCode → SaveTokenData.

Changes:
- pollPatDeviceFlow: return (status, authCode, error) instead of
  (status, error); extract Data.AuthCode on APPROVED
- handlePatAuthCheck APPROVED path: call ExchangeCodeForToken +
  SaveTokenData before ResetRuntimeTokenCache, with graceful
  fallback on exchange failure
- Update all 6 poll test cases for new 3-return signature; add
  authCode assertions for APPROVED/REJECTED/EXPIRED/CANCELLED/
  ServerError scenarios

* fix(pat): CR round-2 must-fix items

- scopeValueRegex: support multi-segment scopes (mail:a.b:send)
  Updated regex to ([a-zA-Z][a-zA-Z0-9_.]*(?::[a-zA-Z][a-zA-Z0-9_.]*)+)

- handlePatAuthCheck: add 3 integration tests (APPROVED/REJECTED/EmptyFlowID)
  Covers the 132-line main orchestrator with mock runner + httptest server

- Extract ParseDeviceFlowStatus + status constants to auth package
  Eliminates string literal duplication across device_flow.go and pat_auth_retry.go

* fix(pat): use direct OAuth path when clientSecret is provided

When PAT error response includes both clientId and clientSecret,
use SetClientID() (direct DingTalk API mode) instead of
SetClientIDFromMCP() (MCP proxy mode). The MCP proxy does not
hold the secret for the PAT-assigned app, causing HTTP 400
'invalidParameter.idOrSecret.notFound' on token exchange.

Now: clientSecret present → direct mode; absent → MCP proxy mode.

---------

Co-authored-by: shangguanxuan.sgx <shangguanxuan.sgx@alibaba-inc.com>
2026-04-19 19:12:34 +08:00
修雨 f762117d4e Merge pull request #126 from DingTalk-Real-AI/refactor/remove-managed-plugin-privilege
refactor(plugin): 移除 default/managed plugin 特权机制
2026-04-19 14:56:59 +08:00
修雨 750b6c04d6 Merge pull request #125 from DingTalk-Real-AI/fix/issue-119-plugin-mcp-timeout
perf(plugin): fast plugin MCP startup via cache + parallel discovery (#119)
2026-04-19 14:56:55 +08:00
修雨 59e51c348a perf(plugin): share cache.Store across discovery + relax cold timeouts
Address reviewer feedback on #125:

- Hoist a single *cache.Store above the discovery fan-out so all HTTP and
  stdio goroutines share one instance rather than each spawning its own in
  registerHTTPServer/registerStdioServer. Atomic tmp+rename per cache key
  keeps concurrent writers on distinct keys collision-free; global runtime
  registries (AppendDynamicServer, RegisterStdioClient) already guard
  themselves with a mutex. Added a comment pointing at those invariants.

- Relax cold-path timeouts to survive healthy cross-region endpoints and
  Python/Node interpreter warm-up: HTTP 500ms → 1s (plain), 700ms → 1.5s
  (auth), stdio 1s → 2s. A new DWS_PLUGIN_COLD_TIMEOUT env var overrides
  all three with a single duration, registered in configmeta so it shows
  up in `dws config --help`.

- Add internal/app/plugin_discovery_concurrency_test.go covering the
  thread-safety claims end-to-end: 16 parallel SaveTools with distinct
  keys, 32 parallel AppendDynamicServer, 32 parallel RegisterStdioClient,
  plus TestResolvePluginColdTimeouts exercising defaults, a valid
  override, an unparseable value, and a non-positive value. All pass with
  go test -race.

Verified: go build ./... clean; targeted -race tests PASS in 1.5s;
./internal/{app,cli,cache,plugin}/... all green.
2026-04-19 14:41:40 +08:00
修雨 a056a9abfb refactor(plugin): purge removed plugin settings instead of merely disabling
Review follow-up on #126: RemovePlugin used to call
setPluginEnabled(name, false), which left the plugin's key in
EnabledPlugins and never touched PluginConfigs. settings.json retained
dangling state for plugins no longer on disk.

Replace the disable with purgePluginFromSettings, which deletes both
the EnabledPlugins entry and any PluginConfigs entry for the removed
plugin, leaving unrelated plugins' state intact. A re-install defaults
to enabled=true via the existing install paths, matching user
expectations.

Covered by TestRemovePluginPurgesSettings across both user and legacy
managed layouts.
2026-04-19 14:32:22 +08:00
修雨 33ae780103 refactor(plugin): remove default/managed plugin privileged mechanism
Drop the hardcoded default-managed-plugin bootstrap that auto-fetched
DingTalk-Real-AI/* plugins on CLI startup. All third-party plugins are
now installed equally via `dws plugin install` — no privileged workspace.

Motivation (issue #124): first-run startup printed
  WARN failed to fetch remote info for default plugin
       DingTalk-Real-AI/conference: invalid character '<' looking for
       beginning of value
because the GitHub Pages registry returned HTML for the missing plugin.
Beyond the error, the whole "built-in plugin" concept contradicts the
CLI's lightweight design: plugin installation planning belongs to
agent-authored skills, not to the CLI binary.

Changes:
- Delete internal/plugin/updater.go (EnsureManaged, CheckAndUpdate,
  checkRemoteVersion, downloadAndInstall, promptUpdate, zip extraction
  — the entire managed-plugin update pipeline).
- Delete internal/plugin/updater_test.go.
- Remove DefaultManagedPlugins, OfficialPluginWorkspace, and
  PluginUpdateCheckInterval from pkg/config/constants.go.
- Remove the EnsureManaged/CheckAndUpdate bootstrap block from
  internal/app/root.go; only legacy LoadManaged() kept for back-compat
  so plugins already installed under ~/.dws/plugins/managed/ still load.
- internal/plugin/loader.go: InstallFromGit always installs under
  PluginUserDir; RemovePlugin allows removal of legacy managed-dir
  plugins instead of refusing.
- internal/app/plugin_cmd.go: drop `--type managed` flag and related
  gating from `plugin install`; disable/remove commands no longer
  differentiate managed vs user.
- Update plugin_test.go accordingly (replace TestRemoveManagedPluginBlocked
  with TestRemoveLegacyManagedPlugin, drop TestPromptUpdate).

Backward compatibility: `PluginManagedDir = "managed"` constant retained
so plugins already on disk from older CLI versions still load and are
removable. No migration required.

Verification:
- go build ./cmd                 → success
- go test ./internal/plugin/...  → PASS
- go test ./internal/app/...     → PASS (436s, matches main baseline)
- Smoke: `/tmp/dws-new plugin list` no longer prints the "Pulling
  built-in plugin" or "failed to fetch remote info" WARN.
- Pre-existing failures on main (unrelated): test/cli_compat,
  test/integration/extensions (requires auth login),
  test/scripts (requires DWS_PACKAGE_VERSION or git tag),
  test/unit TestOpenSourceTreeOmitsEmbeddedHostMarkers.

Closes #124
2026-04-19 11:18:56 +08:00
修雨 daf56514f7 perf(plugin): parallelize all plugin discovery and tighten cold timeouts
Follow-up to cache-first registration (c95ec04). The cold-cache wall
clock is now bounded by the slowest individual plugin, not the sum:

- Fan out HTTP and stdio discovery together in goroutines rather than
  running the stdio loop serially after HTTP.
- HTTP cold budget: 4s → 700ms (auth) / 500ms (plain). Honest dial
  timeouts on unreachable endpoints fail fast; healthy third-party
  endpoints respond well under the window. The outcome is still saved
  as a negative cache, so a miss this run costs the next run ~0ms.
- Stdio cold budget: 4s → 1s. A local subprocess handshake is milliseconds.

With three user plugins (one pointing at TEST-NET-1, permanently
unreachable), cold `dws --help` drops from ~3.3s to ~0.8s and warm stays
at ~80ms. The savings compound linearly with plugin count.
2026-04-18 14:22:53 +08:00
修雨 8bcbceb971 perf(plugin): serve plugin MCP tool list from disk cache on startup
Follow-up to the 4s startup cap (df01f36). Plugin command registration
now reads the tools snapshot directly from the cache store on the hot
path, and only falls through to an Initialize+ListTools RPC when no
snapshot exists.

- registerHTTPServer / registerStdioServer try cache.LoadTools first;
  on hit they build Cobra commands synchronously with zero network I/O.
- Cold cache falls back to the existing synchronous discovery (already
  bounded at 2-4s) and persists the outcome — including empty tool
  lists — as a negative cache so the next invocation is fast regardless
  of endpoint health.
- Cache entries are namespaced under "plugin:<name>:<server>" so they
  are distinct from Market cache entries in `dws cache status`, and
  refresh on-demand via `dws cache clean` / `dws cache refresh`; the
  existing 7d ToolsTTL otherwise expires entries naturally.

Warm-cache `dws --help` with three user plugins (one unreachable) now
returns in ~80ms versus 3.7s with synchronous discovery, a >40x win
when endpoints are offline.
2026-04-18 00:47:20 +08:00
修雨 df01f36442 fix(transport): cap plugin MCP startup at ~4s when endpoints are unreachable
Issue #119: an unreachable third-party MCP plugin (e.g. blocked/firewalled
endpoint) blocks `dws --help` for ~10s on every CLI invocation, because
plugin discovery happens eagerly during command-tree construction.

Root cause was a stack of timeouts that multiplied under transport failure:

* `transport.Client.Initialize` loops three supported protocol versions on
  every error — including dial timeouts and HTTP 5xx — even though those
  failure modes are protocol-version-independent. Three loops × the dial
  budget = 3× the worst-case startup cost.
* Default `DialContext.Timeout` was 10s, so a single dial against an
  unroutable address (e.g. TEST-NET-1) burned the full plugin context.
* `registerHTTPServer` granted plugins with `AuthHeaders` a 10s outer
  context (intended for slow third-party services), and `registerStdioServer`
  granted every stdio plugin 10s. Either single misbehaving plugin therefore
  stalled the entire CLI.
* Registry-side `defaultDiscoveryTimeout` (10s) and `perServerDiscoveryTimeout`
  (5s) had the same shape on the LoadCatalog path.

Changes:

* `transport.Client.Initialize`: short-circuit when the underlying
  `*CallError.Stage` is anything other than `CallStageJSONRPC`. Protocol
  version negotiation is the only justification for retrying with another
  version; transport/HTTP failures fail identically and should surrender.
* `transport.defaultTransport`: `DialContext.Timeout` 10s → 3s.
* `app.registerHTTPServer`: AuthHeaders timeout 10s → 4s.
* `app.registerStdioServer`: outer ctx timeout 10s → 4s.
* `cli.defaultDiscoveryTimeout`: 10s → 4s.
* `discovery.Service`: rename `perServerDiscoveryTimeout` 5s → 2s
  (`defaultPerServerDiscoveryTimeout`); add `Service.PerServerTimeout`
  override field for tests/callers needing a tighter or looser bound.

Tests:

* `TestInitializeShortCircuitsOnHTTPError` — asserts only ONE protocol
  version is attempted on HTTP 5xx.
* `TestInitializeShortCircuitsOnDialFailure` — asserts Initialize returns
  in <2s on a refused-connection address.
* Existing `TestInitializeNegotiatesProtocolVersion` continues to pass —
  JSON-RPC-stage errors still trigger version fallback.

End-to-end measurement with a blackhole plugin
(endpoint=192.0.2.1, AuthHeaders set):

  | Variant       | median `dws --help` |
  | ------------- | ------------------- |
  | main          | 10.2s               |
  | PR #121 alone | 9.1s                |
  | this branch   | 3.72s (-63%)        |

The `loader.go` and `discovery/service.go` timeout reductions overlap with
PR #121 by @utafrali; see PR description for attribution.

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

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

Made-with: Cursor
2026-04-08 20:03:56 +08:00
coffeeBigSir e59c4f30b8 Merge pull request #95 from DingTalk-Real-AI/coffeeBigSir-patch-1
Update SKILL.md
2026-04-08 17:44:38 +08:00
coffeeBigSir fd7ef5edc2 Update SKILL.md 2026-04-08 17:21:23 +08:00
xianfeng wang a8e1acec09 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-08 10:41:33 +08:00
coffeeBigSir 31eb10985e Merge pull request #82 from wqyenjoy/main
add command
2026-04-07 20:31:49 +08:00
coffeeBigSir 1436b62a80 Merge pull request #86 from DingTalk-Real-AI/install-yh
feat(install): align skill dirs with npm and add OpenClaw
2026-04-07 20:25:31 +08:00
tianlei.qjb ec6a27635b feat(install): align skill dirs with npm and add OpenClaw 2026-04-07 16:02:05 +08:00
玉澜 1727744691 add command 2026-04-03 20:01:05 +08:00
fantiu afdd47b5a5 Merge pull request #81 from fantiu/feat-cmd-performance
perf: optimize command timeout handling, instrumentation, and diagnostics
2026-04-03 19:52:42 +08:00
fantiu d968e8e551 perf: optimize command timeout handling, instrumentation, and diagnostics 2026-04-03 18:03:28 +08:00
fantiu c649d1a762 Merge branch 'main' of github.com:fantiu/dingtalk-workspace-cli into feat-cmd-performance
perf: optimize command timeout handling, instrumentation, and diagnostics
2026-04-03 18:02:33 +08:00
fantiu a1f5d97345 Merge branch 'DingTalk-Real-AI:main' into main 2026-04-03 18:01:52 +08:00
meng93 58062515a5 Merge pull request #78 from DingTalk-Real-AI/feat/issue-ai-table-label-change
feat: 优化标签在多维表中的展示
2026-04-03 11:14:53 +08:00
meng93 5614b508f2 feat: to #73551688 优化标签在多维表中的展示 2026-04-03 10:41:49 +08:00
github-actions[bot] 5e003a41b1 chore: update coverage badge [skip ci] 2026-04-03 02:31:09 +00:00
wxianfeng 4eaeb1dd4a fix conflict 2026-04-03 10:29:05 +08:00
github-actions[bot] 84471bd6f0 chore: update coverage badge [skip ci] 2026-04-03 02:03:16 +00:00
fantiu c8e3ac21c2 Merge pull request #76 from DingTalk-Real-AI/npm
docs: add npm install method to README
2026-04-03 09:17:55 +08:00
tianlei.qjb c38892b7cf docs: add npm install method to README 2026-04-02 22:34:42 +08:00
fantiu 1a0a5324f0 docs: note upgrade command requires v1.0.7+ 2026-04-02 20:06:55 +08:00
github-actions[bot] c1e9e9e0d6 chore: update coverage badge [skip ci] 2026-04-02 09:03:29 +00:00
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
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
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
xianfeng wang 110f887181 Merge branch 'DingTalk-Real-AI:main' into main 2026-03-30 15:10:33 +08:00
github-actions[bot] 25118d1ec7 chore: update coverage badge [skip ci] 2026-03-30 03:09:54 +00:00
153 changed files with 20606 additions and 728 deletions
+1 -1
View File
@@ -1 +1 @@
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 52.8%"><title>coverage: 52.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">52.8%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">52.8%</text></g></svg>
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 48.7%"><title>coverage: 48.7%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">48.7%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">48.7%</text></g></svg>

Before

Width:  |  Height:  |  Size: 1.1 KiB

After

Width:  |  Height:  |  Size: 1.1 KiB

+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
+2 -1
View File
@@ -33,7 +33,8 @@ jobs:
title: issue.title,
body: issue.body,
state: issue.state,
html_url: issue.html_url
html_url: issue.html_url,
labels: (issue.labels || []).map(label => label.name).join(', ') || '无标签'
}
};
+2 -1
View File
@@ -28,13 +28,14 @@ jobs:
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${content}${content.length >= 500 ? '...' : ''}\n\n[点击查看详情](${url})\n\n---\n📦 ${context.repo.owner}/${context.repo.repo}\n\n**关键词**: 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`
}
};
+3
View File
@@ -27,3 +27,6 @@ test/cli_compat/testdata/
credentials*
plans
_docs
dws.zip
*.code-workspace
/dingtalk-workspace.zip
+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
+44
View File
@@ -28,6 +28,7 @@
- [Why dws?](#why-dws)
- [Installation](#installation)
- [Upgrade](#upgrade)
- [Getting Started](#getting-started)
- [Quick Start](#quick-start)
- [Using with Agents](#using-with-agents)
@@ -65,6 +66,12 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
<details>
<summary>Other install methods</summary>
**npm** (requires Node.js (npm/npx)):
```bash
npm install -g dingtalk-workspace-cli
```
**Pre-built binary**: download from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases).
> **macOS users**: If you see "cannot be opened because Apple cannot check it for malicious software", run:
@@ -85,6 +92,43 @@ 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
```bash
+44
View File
@@ -28,6 +28,7 @@
- [为什么选择 dws?](#why-dws)
- [安装](#安装)
- [升级](#升级)
- [开始使用](#开始使用)
- [快速开始](#快速开始)
- [在 Agent 中使用](#在-agent-中使用)
@@ -65,6 +66,12 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
<details>
<summary>其他安装方式</summary>
**npm**(需要 Node.js(npm/npx)):
```bash
npm install -g dingtalk-workspace-cli
```
**预编译二进制文件**:从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 下载。
> **macOS 用户注意**:如果提示“无法打开,因为 Apple 无法检查其是否包含恶意软件”,请执行:
@@ -85,6 +92,43 @@ 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>
## 开始使用
```bash
+1
View File
@@ -52,6 +52,7 @@ __KEG_ONLY_LINE__
Pathname.new(File.join(Dir.home, ".amp/skills/dws")),
Pathname.new(File.join(Dir.home, ".kiro/skills/dws")),
Pathname.new(File.join(Dir.home, ".trae/skills/dws")),
Pathname.new(File.join(Dir.home, ".openclaw/skills/dws")),
]
targets.each_with_index do |dest, index|
+2
View File
@@ -7,6 +7,7 @@ const os = require("os");
const path = require("path");
const childProcess = require("child_process");
// Canonical list: keep scripts/install.sh, scripts/install.ps1, scripts/install-skills.sh in sync.
const AGENT_DIRS = [
".agents/skills",
".claude/skills",
@@ -20,6 +21,7 @@ const AGENT_DIRS = [
".amp/skills",
".kiro/skills",
".trae/skills",
".openclaw/skills",
];
const PLATFORM_MAP = {
+83
View File
@@ -0,0 +1,83 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"context"
"errors"
"fmt"
"io"
"log/slog"
"path/filepath"
"strings"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
// the same host compatibility hooks as MCP. It mirrors the former body of
// getCachedRuntimeToken (excluding process-level cache and timing).
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, disc)
configureOAuthProviderCompatibility(provider, configDir)
token, tokenErr := provider.GetAccessToken(ctx)
if tokenErr == nil && strings.TrimSpace(token) != "" {
return strings.TrimSpace(token), nil
}
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
return "", tokenErr
}
manager := authpkg.NewManager(configDir, nil)
configureLegacyAuthManagerCompatibility(manager)
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
return strings.TrimSpace(leg), nil
}
return "", nil
}
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
// the active edition config directory, the same process-cached path as MCP is used.
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
if t := strings.TrimSpace(explicitToken); t != "" {
return t, nil
}
if strings.TrimSpace(configDir) == "" {
return "", fmt.Errorf("config directory is empty")
}
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
return tok, nil
}
return "", noCredentialsError()
}
tok, err := resolveAccessTokenFromDir(ctx, configDir)
if err != nil {
return "", err
}
if tok != "" {
return tok, nil
}
return "", noCredentialsError()
}
func noCredentialsError() error {
if edition.Get().IsEmbedded {
return fmt.Errorf("认证信息已失效,请重新认证")
}
return fmt.Errorf("no credentials found, run: dws auth login")
}
+26
View File
@@ -0,0 +1,26 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0
package app
import (
"context"
"testing"
)
func TestResolveAuxiliaryAccessToken_explicitToken(t *testing.T) {
tok, err := ResolveAuxiliaryAccessToken(context.Background(), "/any/dir", " bearer-xyz ")
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if tok != "bearer-xyz" {
t.Fatalf("got %q, want bearer-xyz", tok)
}
}
func TestResolveAuxiliaryAccessToken_emptyConfigDir(t *testing.T) {
_, err := ResolveAuxiliaryAccessToken(context.Background(), " ", "")
if err == nil {
t.Fatal("expected error for empty config directory")
}
}
+20 -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()
@@ -203,10 +207,13 @@ func newAuthLogoutCommand() *cobra.Command {
_ = os.Remove(filepath.Join(configDir, "mcp_url"))
_ = os.Remove(filepath.Join(configDir, "token"))
_ = os.Remove(filepath.Join(configDir, "token.json"))
ResetRuntimeTokenCache()
clearCompatCache()
w := cmd.OutOrStdout()
fmt.Fprintln(w, "[OK] 已清除所有认证信息")
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
if !edition.Get().IsEmbedded {
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
}
return nil
},
}
@@ -236,6 +243,8 @@ func newAuthStatusCommand() *cobra.Command {
tokenData = updatedData
refreshed = true
}
} else if edition.Get().AutoPurgeToken {
_ = authpkg.DeleteTokenData(configDir)
}
}
if authStatusAuthenticated(tokenData) {
@@ -263,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
},
@@ -299,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()
@@ -338,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 -2
View File
@@ -14,6 +14,7 @@
package app
import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/spf13/cobra"
)
@@ -35,8 +36,8 @@ type GlobalFlags struct {
}
func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", "Override OAuth client ID (DingTalk AppKey)")
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", "Override OAuth client secret (DingTalk AppSecret)")
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", i18n.T("覆盖 OAuth 客户端 ID (钉钉 AppKey)"))
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", i18n.T("覆盖 OAuth 客户端密钥 (钉钉 AppSecret)"))
cmd.PersistentFlags().BoolVar(&flags.Debug, "debug", false, "显示调试日志")
cmd.PersistentFlags().BoolVar(&flags.DryRun, "dry-run", false, "预览操作内容,不实际执行")
cmd.PersistentFlags().StringVar(&flags.Fields, "fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
+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()
}
+36 -34
View File
@@ -16,7 +16,6 @@ package app
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"net"
"net/http"
@@ -30,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...)
}
@@ -47,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;
@@ -60,13 +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 {
totalStart := time.Now()
defer func() {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] loadDynamicCommands total: %v\n", time.Since(totalStart))
}
}()
store := cacheStoreFromEnv()
partition := config.DefaultPartition
@@ -79,18 +100,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
// --- Cache-first server registry ---
cacheLoadStart := time.Now()
snapshot, freshness, cacheErr := store.LoadRegistry(partition)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache load: %v (err=%v)\n", time.Since(cacheLoadStart), cacheErr)
}
RecordTiming(ctx, "registry_cache", time.Since(cacheLoadStart))
var servers []market.ServerDescriptor
now := store.Now().UTC()
usingCachedRegistry := useCache && cacheErr == nil && len(snapshot.Servers) > 0
if usingCachedRegistry {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] using cached registry: servers=%d, freshness=%s\n", len(snapshot.Servers), freshness)
}
servers = snapshot.Servers
// Only trigger async revalidation in production (no URL override).
// Tests set discoveryBaseURLOverride and control cache expiry directly,
@@ -106,15 +122,10 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
if discoveryBaseURLOverride != "" {
baseURL = discoveryBaseURLOverride
}
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] fetching from market API: %s\n", baseURL)
}
fetchStart := time.Now()
client := market.NewClient(baseURL, ipv4OnlyHTTPClient())
resp, fetchErr := client.FetchServers(ctx, config.DefaultFetchServersLimit)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] market API fetch: %v (err=%v)\n", time.Since(fetchStart), fetchErr)
}
RecordTiming(ctx, "market_fetch", time.Since(fetchStart))
if fetchErr != nil {
slog.Debug("loadDynamicCommands: market API fetch failed", "error", fetchErr)
// Degrade to stale cache if available (production only).
@@ -126,18 +137,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
}
} else {
servers = market.NormalizeServers(resp, "market")
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] normalized servers: %d\n", len(servers))
}
// Persist fresh data (only in non-test mode).
if useCache {
saveStart := time.Now()
if saveErr := store.SaveRegistry(partition, cache.RegistrySnapshot{Servers: servers}); saveErr != nil {
slog.Debug("loadDynamicCommands: failed to save registry cache", "error", saveErr)
}
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache save: %v\n", time.Since(saveStart))
}
RecordTiming(ctx, "cache_save", time.Since(saveStart))
}
}
}
@@ -150,15 +156,11 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
detailStart := time.Now()
detailsByID := loadCachedDetailsFast(store, servers)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] load details: %v\n", time.Since(detailStart))
}
RecordTiming(ctx, "tool_metadata", time.Since(detailStart))
buildStart := time.Now()
cmds := compat.BuildDynamicCommands(servers, runner, detailsByID)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] build commands: %v (count=%d)\n", time.Since(buildStart), len(cmds))
}
RecordTiming(ctx, "build_commands", time.Since(buildStart))
return cmds
}
+598
View File
@@ -0,0 +1,598 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"context"
"encoding/json"
stderrors "errors"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"os"
"os/exec"
"regexp"
"runtime"
"strings"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/fatih/color"
)
const (
// PatAuthRetryTimeout is the maximum time to wait for user authorization
// when a PAT scope error is detected.
PatAuthRetryTimeout = 10 * time.Minute
// PatAuthPollInterval is how often we poll to check if the user has
// completed authorization.
PatAuthPollInterval = 5 * time.Second
)
// PatScopeError holds information about a missing PAT scope.
type PatScopeError struct {
OriginalError string
Identity string
ErrorType string
Message string
Hint string
MissingScope string
}
func (e *PatScopeError) Error() string {
return e.OriginalError
}
// patScopeRegex matches PAT-protocol scope error patterns from the API.
// Only matches explicit scope-related keywords; generic "permission denied" or
// "forbidden" are intentionally excluded to avoid false positives on business
// authorization errors (e.g. mailbox access denied, 403 Forbidden).
var patScopeRegex = regexp.MustCompile(`(?i)(missing_scope|insufficient_scope|scope.*required)`)
// scopeValueRegex extracts a scope identifier (e.g. "calendar:read",
// "mail:user_mailbox.message:send") from an error message.
// Supports multi-segment scopes with multiple colons (resource:sub:action).
var scopeValueRegex = regexp.MustCompile(`([a-zA-Z][a-zA-Z0-9_.]*(?::[a-zA-Z][a-zA-Z0-9_.]*)+)`)
// identityValueRegex extracts an identity label from an error message.
var identityValueRegex = regexp.MustCompile(`(?i)identity["\s:]+([a-zA-Z_]+)`)
// isPatScopeError checks if an error looks like a PAT scope/permission error
// that can be resolved by re-authorizing with additional scopes.
func isPatScopeError(err error) bool {
if err == nil {
return false
}
msg := strings.ToLower(err.Error())
// Check for missing_scope pattern in error message or hint
if patScopeRegex.MatchString(msg) {
return true
}
var typed *apperrors.Error
if stderrors.As(err, &typed) {
// Check message, reason, and hint for scope-related patterns
fullText := strings.ToLower(typed.Message + " " + typed.Reason + " " + typed.Hint)
if typed.Category == apperrors.CategoryAuth {
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") ||
(strings.Contains(fullText, "scope") && strings.Contains(fullText, "required")) {
return true
}
}
// Any category with scope/permission hints
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") {
return true
}
}
return false
}
// extractPatScopeError parses an error to extract PAT scope details.
func extractPatScopeError(err error) *PatScopeError {
if err == nil {
return nil
}
msg := err.Error()
scope := ""
var typed *apperrors.Error
if stderrors.As(err, &typed) {
msg = typed.Message
if typed.Reason != "" {
msg += " (" + typed.Reason + ")"
}
}
// Try to extract scope value (e.g. "calendar:read") from error message.
scopeMatch := scopeValueRegex.FindStringSubmatch(msg)
if len(scopeMatch) > 1 {
scope = scopeMatch[1]
}
// Try to extract identity from error message.
identity := "user"
identityMatch := identityValueRegex.FindStringSubmatch(msg)
if len(identityMatch) > 1 {
identity = identityMatch[1]
}
return &PatScopeError{
OriginalError: err.Error(),
Identity: identity,
ErrorType: "missing_scope",
Message: msg,
Hint: fmt.Sprintf("run `dws auth login --scope %q` to authorize the missing scope", scope),
MissingScope: scope,
}
}
// PrintPatAuthError prints a human-readable PAT authorization error.
func PrintPatAuthError(w io.Writer, scopeErr *PatScopeError) {
bold := color.New(color.Bold).SprintFunc()
cyan := color.New(color.FgCyan).SprintFunc()
dim := color.New(color.Faint).SprintFunc()
green := color.New(color.FgGreen).SprintFunc()
fmt.Fprintln(w)
fmt.Fprintf(w, "{\n")
fmt.Fprintf(w, " %s: %s,\n", bold("\"ok\""), "false")
fmt.Fprintf(w, " %s: %q,\n", bold("\"identity\""), scopeErr.Identity)
fmt.Fprintf(w, " %s: {\n", bold("\"error\""))
fmt.Fprintf(w, " %s: %q,\n", bold("\"type\""), scopeErr.ErrorType)
fmt.Fprintf(w, " %s: %q,\n", bold("\"message\""), scopeErr.Message)
fmt.Fprintf(w, " %s: %q\n", bold("\"hint\""), scopeErr.Hint)
fmt.Fprintf(w, " }\n")
fmt.Fprintf(w, "}\n")
fmt.Fprintln(w)
// Print authorization instructions
fmt.Fprintf(w, "%s %s\n", green("▶"), bold("需要额外授权"))
fmt.Fprintln(w)
fmt.Fprintf(w, " %s %s\n", dim("#"), dim("运行以下命令完成授权"))
if scopeErr.MissingScope != "" {
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan(fmt.Sprintf("dws auth login --scope %q", scopeErr.MissingScope)))
} else {
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan("dws auth login"))
}
fmt.Fprintln(w)
fmt.Fprintf(w, " %s 在浏览器中打开授权链接,完成授权后重新执行命令\n", dim("ℹ"))
fmt.Fprintln(w)
}
// PrintPatAuthJSON prints a machine-readable PAT authorization error.
func PrintPatAuthJSON(w io.Writer, scopeErr *PatScopeError) {
payload := map[string]any{
"ok": false,
"identity": scopeErr.Identity,
"error": map[string]any{
"type": scopeErr.ErrorType,
"message": scopeErr.Message,
"hint": scopeErr.Hint,
},
}
if scopeErr.MissingScope != "" {
payload["missing_scope"] = scopeErr.MissingScope
}
data, _ := json.MarshalIndent(payload, "", " ")
fmt.Fprintln(w, string(data))
}
// WaitForPatAuthorization polls until the user completes authorization or timeout.
// It returns true if authorization was completed, false if timed out or cancelled.
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
bold := color.New(color.Bold).SprintFunc()
yellow := color.New(color.FgYellow).SprintFunc()
green := color.New(color.FgGreen).SprintFunc()
red := color.New(color.FgRed).SprintFunc()
dim := color.New(color.Faint).SprintFunc()
timeout := PatAuthRetryTimeout
deadline := time.Now().Add(timeout)
pollTicker := time.NewTicker(PatAuthPollInterval)
defer pollTicker.Stop()
start := time.Now()
fmt.Fprintln(output)
fmt.Fprintf(output, "%s %s\n", yellow("⏳"), bold("等待用户授权..."))
fmt.Fprintf(output, " %s 请在另一个终端完成 dws auth login 授权\n", dim("ℹ"))
fmt.Fprintf(output, " %s 超时时间: %s\n", dim("⏱"), timeout)
fmt.Fprintln(output)
pollCount := 0
for {
select {
case <-ctx.Done():
fmt.Fprintf(output, "%s 操作已取消\n", red("✗"))
return false
case <-time.After(time.Until(deadline)):
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", red("✗"), timeout)
fmt.Fprintf(output, " %s 请重新执行命令\n", dim("ℹ"))
return false
case <-pollTicker.C:
pollCount++
elapsed := time.Since(start).Truncate(time.Second)
remaining := time.Until(deadline).Truncate(time.Second)
// Check if token is now valid
tokenData, err := authpkg.LoadTokenData(configDir)
if err == nil && tokenData != nil {
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
green("✓"), bold("授权成功!"), elapsed, remaining)
fmt.Fprintln(output)
return true
}
}
// Show polling status
fmt.Fprintf(output, "\r%s [%d] 等待授权中... (%s 已用, %s 剩余) ",
dim("⟳"), pollCount, elapsed, remaining)
}
}
}
// retryWithPatAuthRetry wraps an invocation that failed with a PAT scope error.
// It waits for the user to complete authorization and then retries the invocation.
func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocation executor.Invocation, scopeErr *PatScopeError, configDir string, output io.Writer) (executor.Result, error) {
// Print the PAT error in human-readable format
PrintPatAuthError(output, scopeErr)
// Wait for user to complete authorization
authorized := WaitForPatAuthorization(ctx, configDir, output)
if !authorized {
return executor.Result{}, apperrors.NewAuth(
"等待用户授权超时",
apperrors.WithReason("pat_auth_timeout"),
apperrors.WithHint(fmt.Sprintf("授权超时 (%s),请重新执行命令", PatAuthRetryTimeout)),
apperrors.WithActions("dws auth login"),
)
}
// Clear the token cache so the new token is loaded
ResetRuntimeTokenCache()
// Retry the invocation
fmt.Fprintln(output)
fmt.Fprintf(output, "%s %s\n", color.New(color.FgGreen).SprintFunc()("▶"),
color.New(color.Bold).SprintFunc()("授权完成,正在重试..."))
fmt.Fprintln(output)
return runner.Run(ctx, invocation)
}
// loadMCPClientIDIfNeeded ensures we have a client ID for device flow.
// Priority: in-memory runtime value → DWS_CLIENT_ID env → MCP remote fetch.
func loadMCPClientIDIfNeeded(ctx context.Context, configDir string) string {
clientID := authpkg.ClientID()
if clientID != "" {
return clientID
}
// Fallback: read from environment variable (set by previous PAT auth or caller).
if envID := os.Getenv("DWS_CLIENT_ID"); envID != "" {
authpkg.SetClientIDFromMCP(envID)
return envID
}
// Last resort: fetch from MCP server.
mcpClientID, err := authpkg.FetchClientIDFromMCP(ctx)
if err == nil && mcpClientID != "" {
authpkg.SetClientIDFromMCP(mcpClientID)
return mcpClientID
}
return ""
}
// ---- handlePatAuthCheck (runner.go entry point) -----------------------------
const (
// patPollInterval is how often we poll the device flow status endpoint.
patPollInterval = 2 * time.Second
// patPollTimeout is the maximum time to wait for user authorization via device flow.
patPollTimeout = 10 * time.Minute
)
// patRetryingKey is a context key to prevent recursive PAT auth checks.
// After APPROVED, the retry should not trigger another PAT flow.
type patRetryingKeyType struct{}
var patRetryingKey = patRetryingKeyType{}
// IsPatRetrying returns true if the current context is already in a PAT retry.
func IsPatRetrying(ctx context.Context) bool {
v, _ := ctx.Value(patRetryingKey).(bool)
return v
}
// handlePatAuthCheck is called by runner.executeInvocation when a PAT
// authorization error is detected. It injects the server-assigned clientId
// as x-robot-uid header, prints authorization details, opens the browser,
// polls the device flow endpoint until the user authorizes, and retries the
// original invocation on success.
func handlePatAuthCheck(
ctx context.Context,
r *runtimeRunner,
invocation executor.Invocation,
patErr *apperrors.PATError,
configDir string,
output io.Writer,
) (executor.Result, error) {
// Parse authorization details from PATError.RawJSON.
var patData struct {
Code string `json:"code"`
Data struct {
Desc string `json:"desc"`
FlowID string `json:"flowId"`
URI string `json:"uri"`
ClientID string `json:"clientId"`
ClientSecret string `json:"clientSecret"`
} `json:"data"`
}
if err := json.Unmarshal([]byte(patErr.RawJSON), &patData); err != nil {
return executor.Result{}, patErr
}
slog.Debug("PAT auth check",
"clientId", patData.Data.ClientID,
"flowId", patData.Data.FlowID,
"hasSecret", patData.Data.ClientSecret != "",
)
// Inject clientId/clientSecret from PAT response as runtime credentials
// so that subsequent device flow auth uses the server-assigned app identity.
if patData.Data.ClientID != "" {
if patData.Data.ClientSecret != "" {
// When both clientId and clientSecret are provided, use direct mode
// (DingTalk API) rather than MCP proxy — the MCP proxy does not hold
// the secret for this particular app.
authpkg.SetClientID(patData.Data.ClientID)
authpkg.SetClientSecret(patData.Data.ClientSecret)
} else {
// No clientSecret — rely on MCP proxy to manage the secret server-side.
authpkg.SetClientIDFromMCP(patData.Data.ClientID)
}
// Persist clientId (and optionally secret) to ~/.dws/app.json so that
// future process invocations can load it at startup and populate
// DWS_CLIENT_ID env before the first MCP request.
appCfg := &authpkg.AppConfig{
ClientID: patData.Data.ClientID,
}
if patData.Data.ClientSecret != "" {
appCfg.ClientSecret = authpkg.PlainSecret(patData.Data.ClientSecret)
}
if err := authpkg.SaveAppConfig(configDir, appCfg); err != nil {
slog.Warn("failed to persist app config from PAT", "error", err)
fmt.Fprintf(output, " \u26a0 保存应用配置失败: %v (下次启动可能需要重新授权)\n", err)
}
}
bold := color.New(color.Bold).SprintFunc()
cyan := color.New(color.FgCyan).SprintFunc()
greenFn := color.New(color.FgGreen).SprintFunc()
yellowFn := color.New(color.FgYellow).SprintFunc()
redFn := color.New(color.FgRed).SprintFunc()
dim := color.New(color.Faint).SprintFunc()
fmt.Fprintln(output)
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("需要 PAT 授权"))
if patData.Data.Desc != "" {
fmt.Fprintf(output, " %s %s\n", dim("ℹ"), patData.Data.Desc)
}
if patData.Data.URI != "" {
fmt.Fprintf(output, " %s %s\n\n", dim("🔗"), cyan(patData.Data.URI))
// Best-effort browser open.
_ = tryOpenBrowser(patData.Data.URI)
}
// If no flowId, we can't poll — fall back to returning PATError for host-app.
if patData.Data.FlowID == "" {
fmt.Fprintln(output)
return executor.Result{}, patErr
}
// Poll the device flow status until user authorizes, rejects, or timeout.
fmt.Fprintf(output, "%s %s\n", yellowFn("⏳"), bold("等待用户授权..."))
fmt.Fprintf(output, " %s 请在浏览器中完成授权,超时时间: %s\n", dim("ℹ"), patPollTimeout)
fmt.Fprintln(output)
pollCtx, cancel := context.WithTimeout(ctx, patPollTimeout)
defer cancel()
status, authCode, err := pollPatDeviceFlow(pollCtx, patData.Data.FlowID, configDir, output)
if err != nil {
fmt.Fprintf(output, "%s 轮询授权状态失败: %v\n", redFn("✗"), err)
return executor.Result{}, patErr
}
switch status {
case authpkg.StatusApproved:
fmt.Fprintf(output, "%s %s\n", greenFn("✓"), bold("授权成功!"))
fmt.Fprintln(output)
// Exchange authCode for a fresh access token (mirrors device_flow loginOnce).
if authCode != "" {
slog.Debug("PAT retry: exchanging authCode for token", "hasCode", true)
tokenData, exchErr := authpkg.ExchangeCodeForToken(ctx, configDir, authCode)
if exchErr != nil {
slog.Warn("PAT retry: exchangeCode failed, retrying with existing token", "error", exchErr)
fmt.Fprintf(output, " %s 换取新 token 失败: %v (将使用现有凭证重试)\n", yellowFn("⚠"), exchErr)
} else {
if err := authpkg.SaveTokenData(configDir, tokenData); err != nil {
slog.Warn("PAT retry: failed to save new token", "error", err)
fmt.Fprintf(output, " %s 保存新 token 失败: %v\n", yellowFn("⚠"), err)
} else {
slog.Debug("PAT retry: token refreshed and saved")
}
}
}
// Clear token cache so the new credentials take effect.
ResetRuntimeTokenCache()
// Workaround: brief delay to let server-side authorization state propagate
// before retrying. Without this the retry may use stale credentials.
slog.Debug("PAT retry: waiting for server-side state propagation", "delay", "1s")
time.Sleep(1 * time.Second)
// Retry the original invocation with pat-retrying flag to prevent recursion.
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("授权完成,正在重试..."))
fmt.Fprintln(output)
slog.Debug("PAT retry: identity env check",
"DWS_CLIENT_ID", os.Getenv("DWS_CLIENT_ID"),
)
retryCtx := context.WithValue(ctx, patRetryingKey, true)
return r.Run(retryCtx, invocation)
case authpkg.StatusRejected:
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("用户已拒绝授权"))
return executor.Result{}, apperrors.NewAuth(
"用户已拒绝授权",
apperrors.WithReason("pat_auth_rejected"),
apperrors.WithHint("用户在浏览器中拒绝了授权请求,请重新执行命令。"),
)
case authpkg.StatusExpired:
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("授权超时"))
return executor.Result{}, apperrors.NewAuth(
"授权超时",
apperrors.WithReason("pat_auth_expired"),
apperrors.WithHint("授权链接已过期,请重新执行命令。"),
)
case authpkg.StatusCancelled:
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("操作已取消"))
return executor.Result{}, apperrors.NewAuth(
"操作已取消",
apperrors.WithReason("pat_auth_cancelled"),
apperrors.WithHint("用户取消了授权操作。"),
)
default:
fmt.Fprintf(output, "%s 未知授权状态: %s\n", redFn("✗"), status)
return executor.Result{}, patErr
}
}
// pollPatDeviceFlow polls the PAT device flow status endpoint until a terminal
// state (APPROVED/REJECTED/EXPIRED) is reached or the context is cancelled.
// Returns the final status string and the authCode (non-empty only on APPROVED).
func pollPatDeviceFlow(ctx context.Context, flowID string, configDir string, output io.Writer) (string, string, error) {
pollURL := fmt.Sprintf("%s%s?flowId=%s",
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
// Load user access token for the poll request header.
var accessToken string
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
accessToken = tokenData.AccessToken
}
// Use a client that does NOT follow redirects, so we can detect SSO 302.
noRedirectClient := &http.Client{
CheckRedirect: func(req *http.Request, via []*http.Request) error {
return http.ErrUseLastResponse
},
}
ticker := time.NewTicker(patPollInterval)
defer ticker.Stop()
dim := color.New(color.Faint).SprintFunc()
pollCount := 0
for {
select {
case <-ctx.Done():
if ctx.Err() == context.Canceled {
return authpkg.StatusCancelled, "", nil
}
return authpkg.StatusExpired, "", nil
case <-ticker.C:
pollCount++
fmt.Fprintf(output, "\r%s [%d] 等待授权中... ", dim("⟳"), pollCount)
req, err := http.NewRequestWithContext(ctx, http.MethodGet, pollURL, nil)
if err != nil {
slog.Debug("PAT poll: failed to create request", "error", err)
continue
}
if accessToken != "" {
req.Header.Set("x-user-access-token", accessToken)
}
resp, err := noRedirectClient.Do(req)
if err != nil {
slog.Debug("PAT poll: request failed", "error", err)
continue // transient network error, keep polling
}
bodyBytes, _ := io.ReadAll(resp.Body)
resp.Body.Close()
// If we got a redirect (302/301), SSO gateway intercepted — skip JSON parse.
if resp.StatusCode == http.StatusFound || resp.StatusCode == http.StatusMovedPermanently {
continue
}
var pollResp authpkg.DevicePollResponse
if err := json.Unmarshal(bodyBytes, &pollResp); err != nil {
slog.Debug("PAT poll: failed to parse response", "error", err, "body", string(bodyBytes))
continue
}
status := authpkg.ParseDeviceFlowStatus(pollResp.Data.Status, pollResp.Success)
switch status {
case authpkg.StatusApproved:
fmt.Fprintln(output) // clear the polling line
return status, pollResp.Data.AuthCode, nil
case authpkg.StatusRejected, authpkg.StatusExpired:
fmt.Fprintln(output) // clear the polling line
return status, "", nil
case authpkg.StatusPending:
// keep polling
default:
// ParseDeviceFlowStatus normalizes empty+!success to EXPIRED,
// so this branch handles truly unknown statuses.
fmt.Fprintln(output)
return status, "", nil
}
}
}
}
// tryOpenBrowser opens url in the default browser; errors are silently ignored.
func tryOpenBrowser(url string) error {
var cmd *exec.Cmd
switch runtime.GOOS {
case "darwin":
cmd = exec.Command("open", url)
case "linux":
cmd = exec.Command("xdg-open", url)
case "windows":
cmd = exec.Command("cmd", "/c", "start", url)
default:
return nil
}
return cmd.Start()
}
+606
View File
@@ -0,0 +1,606 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync/atomic"
"testing"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
)
func TestIsPatScopeError_MissingScope(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
if !isPatScopeError(err) {
t.Fatal("expected missing_scope error to be detected")
}
}
func TestIsPatScopeError_PlainString(t *testing.T) {
t.Parallel()
err := &PatScopeError{
OriginalError: "missing_scope: user lacks required scope",
ErrorType: "missing_scope",
Message: "user lacks required scope",
}
if !isPatScopeError(err) {
t.Fatal("expected plain string with missing_scope to be detected")
}
}
func TestIsPatScopeError_NotScopeError(t *testing.T) {
t.Parallel()
err := apperrors.NewValidation("invalid parameter")
if isPatScopeError(err) {
t.Fatal("expected validation error NOT to be detected as scope error")
}
}
func TestIsPatScopeError_Nil(t *testing.T) {
t.Parallel()
if isPatScopeError(nil) {
t.Fatal("nil error should not be detected as scope error")
}
}
func TestIsPatScopeError_WithReason(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth("API error",
apperrors.WithReason("missing_scope"),
)
if !isPatScopeError(err) {
t.Fatal("expected error with missing_scope reason to be detected")
}
}
func TestIsPatScopeError_InsufficientScope(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth("insufficient_scope for resource",
apperrors.WithReason("insufficient_scope"),
)
if !isPatScopeError(err) {
t.Fatal("expected insufficient_scope error to be detected")
}
}
func TestExtractPatScopeError_MissingScope(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
scopeErr := extractPatScopeError(err)
if scopeErr == nil {
t.Fatal("expected non-nil PatScopeError")
}
if scopeErr.ErrorType != "missing_scope" {
t.Errorf("expected error type 'missing_scope', got %q", scopeErr.ErrorType)
}
if !strings.Contains(scopeErr.Hint, "dws auth login") {
t.Errorf("expected hint to contain 'dws auth login', got %q", scopeErr.Hint)
}
}
func TestExtractPatScopeError_ExtractsScope(t *testing.T) {
t.Parallel()
err := &PatScopeError{
OriginalError: "missing_scope: user needs calendar:read",
ErrorType: "missing_scope",
Message: "user needs calendar:read",
}
scopeErr := extractPatScopeError(err)
if scopeErr == nil {
t.Fatal("expected non-nil PatScopeError")
}
if scopeErr.MissingScope != "calendar:read" {
t.Errorf("expected MissingScope 'calendar:read', got %q", scopeErr.MissingScope)
}
}
func TestPrintPatAuthError_HumanReadable(t *testing.T) {
t.Parallel()
var buf strings.Builder
scopeErr := &PatScopeError{
Identity: "user",
ErrorType: "missing_scope",
Message: "missing required scope(s): mail:user_mailbox.message:send",
Hint: "run `dws auth login --scope \"mail:user_mailbox.message:send\"` to authorize",
MissingScope: "mail:user_mailbox.message:send",
}
PrintPatAuthError(&buf, scopeErr)
output := buf.String()
if !strings.Contains(output, "missing_scope") {
t.Errorf("expected output to contain 'missing_scope', got: %s", output)
}
if !strings.Contains(output, "dws auth login") {
t.Errorf("expected output to contain 'dws auth login', got: %s", output)
}
if !strings.Contains(output, "需要额外授权") {
t.Errorf("expected output to contain Chinese auth prompt, got: %s", output)
}
}
func TestPrintPatAuthJSON_MachineReadable(t *testing.T) {
t.Parallel()
var buf strings.Builder
scopeErr := &PatScopeError{
Identity: "user",
ErrorType: "missing_scope",
Message: "missing required scope(s): mail:send",
Hint: "run dws auth login --scope mail:send",
MissingScope: "mail:send",
}
PrintPatAuthJSON(&buf, scopeErr)
output := buf.String()
if !strings.Contains(output, `"ok": false`) {
t.Errorf("expected JSON to contain ok: false, got: %s", output)
}
if !strings.Contains(output, `"missing_scope": "mail:send"`) {
t.Errorf("expected JSON to contain missing_scope, got: %s", output)
}
}
func TestIsPatScopeError_BusinessPermissionDenied(t *testing.T) {
t.Parallel()
// Generic business "permission denied" should NOT trigger PAT re-auth.
err := apperrors.NewAuth("User has no permission to access this mailbox, permission denied")
if isPatScopeError(err) {
t.Fatal("generic 'permission denied' should not be detected as PAT scope error")
}
}
func TestIsPatScopeError_GenericForbidden(t *testing.T) {
t.Parallel()
// HTTP 403 Forbidden should NOT trigger PAT re-auth.
err := apperrors.NewAuth("403 Forbidden")
if isPatScopeError(err) {
t.Fatal("'403 Forbidden' should not be detected as PAT scope error")
}
}
func TestExtractPatScopeError_ComplexScope(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
scopeErr := extractPatScopeError(err)
if scopeErr == nil {
t.Fatal("expected non-nil PatScopeError")
}
if scopeErr.MissingScope != "mail:user_mailbox.message:send" {
t.Errorf("expected MissingScope 'mail:user_mailbox.message:send', got %q", scopeErr.MissingScope)
}
}
func TestPatScopeError_Error(t *testing.T) {
t.Parallel()
err := &PatScopeError{
OriginalError: "test error message",
}
if err.Error() != "test error message" {
t.Errorf("expected Error() to return OriginalError, got %q", err.Error())
}
}
// ---------------------------------------------------------------------------
// pollPatDeviceFlow integration tests — httptest mock covering four terminal
// states: APPROVED, REJECTED, EXPIRED, CANCELLED (ctx cancel).
// ---------------------------------------------------------------------------
// setupPollServer creates an httptest server that responds to
// /cli/oauth/device/poll?flowId=<fid> with the given status sequence.
// It also writes the server URL into a temp DWS_CONFIG_DIR/mcp_url so that
// GetMCPBaseURL() returns the test server address.
func setupPollServer(t *testing.T, statuses []authpkg.DevicePollResponse) (*httptest.Server, string) {
t.Helper()
var callCount atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
idx := int(callCount.Add(1)) - 1
if idx >= len(statuses) {
idx = len(statuses) - 1
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(statuses[idx])
}))
// Write mcp_url so GetMCPBaseURL picks up the test server.
tmpDir := t.TempDir()
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
t.Setenv("DWS_CONFIG_DIR", tmpDir)
return server, tmpDir
}
func TestPollPatDeviceFlow_Approved(t *testing.T) {
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
{Success: true, Data: authpkg.DevicePollData{Status: "APPROVED", AuthCode: "code123"}},
})
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var buf bytes.Buffer
status, authCode, err := pollPatDeviceFlow(ctx, "flow-1", configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "APPROVED" {
t.Errorf("expected APPROVED, got %q", status)
}
if authCode != "code123" {
t.Errorf("expected authCode 'code123', got %q", authCode)
}
}
func TestPollPatDeviceFlow_Rejected(t *testing.T) {
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
{Success: false, Data: authpkg.DevicePollData{Status: "REJECTED"}},
})
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var buf bytes.Buffer
status, authCode, err := pollPatDeviceFlow(ctx, "flow-2", configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "REJECTED" {
t.Errorf("expected REJECTED, got %q", status)
}
if authCode != "" {
t.Errorf("expected empty authCode for REJECTED, got %q", authCode)
}
}
func TestPollPatDeviceFlow_Expired(t *testing.T) {
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
{Success: false, Data: authpkg.DevicePollData{Status: "EXPIRED"}},
})
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var buf bytes.Buffer
status, authCode, err := pollPatDeviceFlow(ctx, "flow-3", configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "EXPIRED" {
t.Errorf("expected EXPIRED, got %q", status)
}
if authCode != "" {
t.Errorf("expected empty authCode for EXPIRED, got %q", authCode)
}
}
func TestPollPatDeviceFlow_Cancelled(t *testing.T) {
// Server always returns PENDING so context cancellation is the only exit.
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
})
defer server.Close()
ctx, cancel := context.WithCancel(context.Background())
// Cancel immediately after first poll tick.
go func() {
time.Sleep(500 * time.Millisecond)
cancel()
}()
var buf bytes.Buffer
status, authCode, err := pollPatDeviceFlow(ctx, "flow-4", configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "CANCELLED" {
t.Errorf("expected CANCELLED, got %q", status)
}
if authCode != "" {
t.Errorf("expected empty authCode for CANCELLED, got %q", authCode)
}
}
// ---------------------------------------------------------------------------
// IsPatRetrying tests
// ---------------------------------------------------------------------------
func TestIsPatRetrying_Default(t *testing.T) {
t.Parallel()
ctx := context.Background()
if IsPatRetrying(ctx) {
t.Fatal("expected false for plain context")
}
}
func TestIsPatRetrying_WithValue(t *testing.T) {
t.Parallel()
ctx := context.WithValue(context.Background(), patRetryingKey, true)
if !IsPatRetrying(ctx) {
t.Fatal("expected true when pat retry key is set")
}
}
// ---------------------------------------------------------------------------
// pollPatDeviceFlow edge cases
// ---------------------------------------------------------------------------
func TestPollPatDeviceFlow_ServerErrorFallback(t *testing.T) {
// When server returns success=false with empty status, should treat as EXPIRED.
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
{Success: false, Data: authpkg.DevicePollData{Status: ""}},
})
defer server.Close()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var buf bytes.Buffer
status, authCode, err := pollPatDeviceFlow(ctx, "flow-err", configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "EXPIRED" {
t.Errorf("expected EXPIRED for server error fallback, got %q", status)
}
if authCode != "" {
t.Errorf("expected empty authCode for server error, got %q", authCode)
}
}
func TestPollPatDeviceFlow_RedirectSkipped(t *testing.T) {
// When server returns 302 (SSO redirect), poll should continue until real response.
var callCount int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
callCount++
if callCount <= 1 {
// First call: simulate SSO redirect
w.Header().Set("Location", "https://sso.example.com")
w.WriteHeader(http.StatusFound)
return
}
// Second call: return APPROVED
w.Header().Set("Content-Type", "application/json")
resp := authpkg.DevicePollResponse{
Success: true,
Data: authpkg.DevicePollData{Status: "APPROVED"},
}
_ = json.NewEncoder(w).Encode(resp)
}))
defer server.Close()
tmpDir := t.TempDir()
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
t.Setenv("DWS_CONFIG_DIR", tmpDir)
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var buf bytes.Buffer
status, _, err := pollPatDeviceFlow(ctx, "flow-redirect", tmpDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if status != "APPROVED" {
t.Errorf("expected APPROVED after redirect, got %q", status)
}
}
// ---------------------------------------------------------------------------
// extractPatScopeError edge cases
// ---------------------------------------------------------------------------
func TestExtractPatScopeError_Nil(t *testing.T) {
t.Parallel()
if got := extractPatScopeError(nil); got != nil {
t.Fatalf("expected nil for nil error, got %+v", got)
}
}
func TestExtractPatScopeError_WithIdentity(t *testing.T) {
t.Parallel()
err := apperrors.NewAuth(`insufficient_scope: identity "app_user" needs calendar:write`)
scopeErr := extractPatScopeError(err)
if scopeErr == nil {
t.Fatal("expected non-nil PatScopeError")
}
if scopeErr.Identity != "app_user" {
t.Errorf("expected Identity 'app_user', got %q", scopeErr.Identity)
}
if scopeErr.MissingScope != "calendar:write" {
t.Errorf("expected MissingScope 'calendar:write', got %q", scopeErr.MissingScope)
}
}
// ---------------------------------------------------------------------------
// handlePatAuthCheck integration tests — cover the main orchestrator with
// mock runner + httptest poll server for APPROVED, REJECTED, EmptyFlowID.
// ---------------------------------------------------------------------------
// mockRunner is a simple executor.Runner for testing handlePatAuthCheck.
type mockRunner struct {
runFunc func(ctx context.Context, inv executor.Invocation) (executor.Result, error)
}
func (m *mockRunner) Run(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
return m.runFunc(ctx, inv)
}
// setupHandlePATServer creates an httptest server for handlePatAuthCheck tests.
// It responds to device poll requests with the given status after the first poll.
func setupHandlePATServer(t *testing.T, terminalStatus string, authCode string) (*httptest.Server, string) {
t.Helper()
var pollCount atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.Contains(r.URL.Path, "/cli/oauth/device/poll") {
idx := int(pollCount.Add(1)) - 1
var resp authpkg.DevicePollResponse
if idx == 0 {
resp = authpkg.DevicePollResponse{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}}
} else {
resp = authpkg.DevicePollResponse{
Success: terminalStatus == "APPROVED",
Data: authpkg.DevicePollData{Status: terminalStatus, AuthCode: authCode},
}
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(resp)
return
}
http.NotFound(w, r)
}))
tmpDir := t.TempDir()
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
t.Setenv("DWS_CONFIG_DIR", tmpDir)
return server, tmpDir
}
func makePATErrorJSON(flowID, clientID string) string {
type patData struct {
Desc string `json:"desc"`
FlowID string `json:"flowId"`
URI string `json:"uri"`
ClientID string `json:"clientId"`
}
payload := struct {
Code string `json:"code"`
Data patData `json:"data"`
}{
Code: "AGENT_CODE_NOT_EXISTS",
Data: patData{
Desc: "test auth",
FlowID: flowID,
URI: "", // empty to avoid opening browser in test
ClientID: clientID,
},
}
data, _ := json.Marshal(payload)
return string(data)
}
func TestHandlePatAuthCheck_Approved(t *testing.T) {
server, configDir := setupHandlePATServer(t, "APPROVED", "test-auth-code")
defer server.Close()
var retryCalled bool
var retryHasKey bool
mock := &mockRunner{
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
retryCalled = true
retryHasKey = IsPatRetrying(ctx)
return executor.Result{Response: map[string]any{"ok": true}}, nil
},
}
runner := &runtimeRunner{fallback: mock}
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-approved", "test-client-id")}
ctx := context.Background()
var buf bytes.Buffer
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
CanonicalProduct: "test",
Tool: "test_tool",
}, patErr, configDir, &buf)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if !retryCalled {
t.Fatal("expected mock runner to be called for retry")
}
if !retryHasKey {
t.Fatal("expected retry context to have patRetryingKey")
}
// Verify SetClientIDFromMCP was called with the PAT response clientId.
if cid := authpkg.ClientID(); cid != "test-client-id" {
t.Errorf("expected ClientID 'test-client-id', got %q", cid)
}
}
func TestHandlePatAuthCheck_Rejected(t *testing.T) {
server, configDir := setupHandlePATServer(t, "REJECTED", "")
defer server.Close()
mock := &mockRunner{
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
t.Fatal("runner should not be called on REJECTED")
return executor.Result{}, nil
},
}
runner := &runtimeRunner{fallback: mock}
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-rejected", "test-client-id")}
ctx := context.Background()
var buf bytes.Buffer
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
CanonicalProduct: "test",
Tool: "test_tool",
}, patErr, configDir, &buf)
if err == nil {
t.Fatal("expected error for REJECTED")
}
if !strings.Contains(err.Error(), "用户已拒绝授权") {
t.Errorf("expected rejection error, got: %v", err)
}
}
func TestHandlePatAuthCheck_EmptyFlowID_FallsBackToPATError(t *testing.T) {
// No poll server needed — empty flowId means no polling, return PATError directly.
tmpDir := t.TempDir()
t.Setenv("DWS_CONFIG_DIR", tmpDir)
mock := &mockRunner{
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
t.Fatal("runner should not be called when flowId is empty")
return executor.Result{}, nil
},
}
runner := &runtimeRunner{fallback: mock}
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("", "test-client-id")}
ctx := context.Background()
var buf bytes.Buffer
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
CanonicalProduct: "test",
Tool: "test_tool",
}, patErr, tmpDir, &buf)
if err == nil {
t.Fatal("expected PATError when flowId is empty")
}
// Should return the original PATError.
if _, ok := err.(*apperrors.PATError); !ok {
t.Errorf("expected *PATError, got %T: %v", err, err)
}
}
+664
View File
@@ -0,0 +1,664 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"fmt"
"os"
"path/filepath"
"strings"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
"github.com/spf13/cobra"
)
func newPluginCommand() *cobra.Command {
pluginCmd := newPlaceholderParent("plugin", i18n.T("插件管理"))
pluginCmd.AddCommand(
newPluginListCommand(),
newPluginInstallCommand(),
newPluginInfoCommand(),
newPluginEnableCommand(),
newPluginDisableCommand(),
newPluginRemoveCommand(),
newPluginValidateCommand(),
newPluginCreateCommand(),
newPluginDevCommand(),
newPluginConfigCommand(),
newPluginBuildCommand(),
)
return pluginCmd
}
func newPluginListCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "list",
Short: i18n.T("列出已安装的插件"),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
loader := plugin.NewLoader(RawVersion())
plugins := loader.ListInstalled()
wantJSON, _ := cmd.Flags().GetBool("json")
if wantJSON {
return output.WriteJSON(cmd.OutOrStdout(), plugins)
}
if len(plugins) == 0 {
fmt.Fprintln(cmd.OutOrStdout(), "No plugins installed.")
return nil
}
w := cmd.OutOrStdout()
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
"NAME", "VERSION", "TYPE", "STATUS", "DESCRIPTION")
fmt.Fprintln(w, strings.Repeat("-", 85))
for _, p := range plugins {
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
p.Name, p.Version, p.Type, statusStr(p.Enabled), p.Description)
}
return nil
},
}
cmd.Flags().Bool("json", false, "Output in JSON format")
return cmd
}
func newPluginInstallCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "install",
Short: i18n.T("安装插件"),
Example: ` dws plugin install --dir ./conference
dws plugin install --git https://github.com/DingTalk-Real-AI/conference.git`,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
dirPath, _ := cmd.Flags().GetString("dir")
gitURL, _ := cmd.Flags().GetString("git")
if dirPath == "" && gitURL == "" {
return apperrors.NewValidation("specify install source: --dir <path> or --git <url>")
}
loader := plugin.NewLoader(RawVersion())
if gitURL != "" {
p, err := loader.InstallFromGit(gitURL)
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
}
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
return nil
}
p, err := loader.InstallFromDir(dirPath)
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
}
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
return nil
},
}
cmd.Flags().String("dir", "", "Install from a local directory")
cmd.Flags().String("git", "", "Install from a Git repository")
return cmd
}
func newPluginInfoCommand() *cobra.Command {
return &cobra.Command{
Use: "info <name>",
Short: i18n.T("查看插件详情"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
name := args[0]
loader := plugin.NewLoader(RawVersion())
plugins := loader.ListInstalled()
for _, p := range plugins {
if p.Name == name {
w := cmd.OutOrStdout()
fmt.Fprintf(w, "Name: %s\n", p.Name)
fmt.Fprintf(w, "Version: %s\n", p.Version)
fmt.Fprintf(w, "Type: %s\n", p.Type)
fmt.Fprintf(w, "Status: %s\n", statusStr(p.Enabled))
fmt.Fprintf(w, "Path: %s\n", p.Path)
if p.Description != "" {
fmt.Fprintf(w, "Description: %s\n", p.Description)
}
return nil
}
}
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found", name))
},
}
}
func newPluginEnableCommand() *cobra.Command {
return &cobra.Command{
Use: "enable <name>",
Short: i18n.T("启用插件"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
loader := plugin.NewLoader(RawVersion())
if err := loader.SetEnabled(args[0], true); err != nil {
return apperrors.NewValidation(err.Error())
}
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s enabled.\n", args[0])
return nil
},
}
}
func newPluginDisableCommand() *cobra.Command {
return &cobra.Command{
Use: "disable <name>",
Short: i18n.T("禁用插件"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
loader := plugin.NewLoader(RawVersion())
if err := loader.SetEnabled(args[0], false); err != nil {
return apperrors.NewValidation(err.Error())
}
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s disabled.\n", args[0])
return nil
},
}
}
func newPluginRemoveCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "remove <name>",
Short: i18n.T("卸载已安装的插件"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
// Stop stdio clients before removing to release file locks
StopStdioClientsByPlugin(args[0])
keepData, _ := cmd.Flags().GetBool("keep-data")
loader := plugin.NewLoader(RawVersion())
if err := loader.RemovePlugin(args[0], keepData); err != nil {
return apperrors.NewValidation(err.Error())
}
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s removed.\n", args[0])
return nil
},
}
cmd.Flags().Bool("keep-data", false, "Keep plugin data directory")
return cmd
}
func newPluginValidateCommand() *cobra.Command {
return &cobra.Command{
Use: "validate <dir>",
Short: i18n.T("校验 plugin.json"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
dir := args[0]
m, err := plugin.ParseManifest(dir + "/plugin.json")
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("parse failed: %v", err))
}
if err := m.Validate(RawVersion()); err != nil {
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
}
fmt.Fprintf(cmd.OutOrStdout(), "Valid: %s (%s)\n", m.Name, m.Version)
return nil
},
}
}
func newPluginCreateCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "create <name>",
Short: i18n.T("脚手架生成新插件目录"),
Example: ` dws plugin create my-tool
dws plugin create my-tool --description "My awesome tool"`,
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
name := args[0]
desc, _ := cmd.Flags().GetString("description")
pluginType := "user"
// Validate name format
m := &plugin.Manifest{Name: name, Version: "0.1.0", Type: pluginType}
if err := m.Validate(""); err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid plugin name: %v", err))
}
dir := filepath.Join(".", name)
if _, err := os.Stat(dir); err == nil {
return apperrors.NewValidation(fmt.Sprintf("directory %q already exists", dir))
}
// Create directory structure
dirs := []string{
dir,
filepath.Join(dir, "skills", name),
filepath.Join(dir, "hooks"),
}
for _, d := range dirs {
if err := os.MkdirAll(d, 0o755); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
}
}
// Write plugin.json
pluginJSON := fmt.Sprintf(`{
"name": %q,
"version": "0.1.0",
"description": %q,
"type": %q,
"minCLIVersion": %q,
"mcpServers": {
%q: {
"type": "stdio",
"command": "${DWS_PLUGIN_ROOT}/bin/server",
"args": []
}
},
"build": {
"command": "echo 'TODO: replace with your build command, e.g.: bun build --compile src/server.ts --outfile bin/server'",
"output": "bin/server"
},
"skills": "./skills/",
"hooks": "./hooks/hooks.json"
}
`, name, desc, pluginType, RawVersion(), name)
if err := os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(pluginJSON), 0o644); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to write plugin.json: %v", err))
}
// Write SKILL.md template
skillMD := fmt.Sprintf(`---
name: %s
description: %s
cli_version: ">=%s"
---
# %s
## Intent Recognition
Use this skill when the user mentions:
- TODO: add your intent keywords here
## Command Decision Tree
| User Intent | Command | Required Parameters |
|-------------|---------|---------------------|
| TODO | `+"`dws %s <sub-command>`"+` | `+"`--param`"+` |
## Parameter Rules
### TODO: parameter type
- Format description
- Conversion rules
`, name, desc, RawVersion(), name, name)
if err := os.WriteFile(filepath.Join(dir, "skills", name, "SKILL.md"), []byte(skillMD), 0o644); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to write SKILL.md: %v", err))
}
// Write hooks.json template
hooksJSON := `{
"hooks": []
}
`
if err := os.WriteFile(filepath.Join(dir, "hooks", "hooks.json"), []byte(hooksJSON), 0o644); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to write hooks.json: %v", err))
}
w := cmd.OutOrStdout()
fmt.Fprintf(w, "Created plugin scaffold at ./%s/\n", name)
fmt.Fprintf(w, " %s/\n", name)
fmt.Fprintf(w, " ├── plugin.json\n")
fmt.Fprintf(w, " ├── skills/%s/SKILL.md\n", name)
fmt.Fprintf(w, " └── hooks/hooks.json\n")
fmt.Fprintln(w)
fmt.Fprintf(w, "Next steps:\n")
fmt.Fprintf(w, " 1. Edit plugin.json to configure your MCP servers\n")
fmt.Fprintf(w, " 2. Edit skills/%s/SKILL.md to describe your commands\n", name)
fmt.Fprintf(w, " 3. Run: dws plugin validate ./%s\n", name)
fmt.Fprintf(w, " 4. Run: dws plugin dev ./%s\n", name)
return nil
},
}
cmd.Flags().String("description", "", "Plugin description")
return cmd
}
func newPluginDevCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "dev <dir>",
Short: i18n.T("将本地目录注册为开发态插件"),
Long: `Registers a plugin from a local source directory for development.
The plugin is loaded directly from the source directory on next CLI invocation,
without copying files to ~/.dws/plugins/. Use 'dws plugin dev --off <name>'
to unregister.`,
Example: ` dws plugin dev ./my-tool
dws plugin dev --off my-tool`,
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
off, _ := cmd.Flags().GetBool("off")
loader := plugin.NewLoader(RawVersion())
if off {
// Unregister dev plugin
name := args[0]
if err := loader.UnregisterDevPlugin(name); err != nil {
return apperrors.NewValidation(err.Error())
}
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q unregistered.\n", name)
return nil
}
// Register dev plugin
dir := args[0]
absDir, err := filepath.Abs(dir)
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
}
// Validate the plugin first
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
}
if err := m.Validate(RawVersion()); err != nil {
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
}
if err := loader.RegisterDevPlugin(m.Name, absDir); err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to register: %v", err))
}
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q registered from %s\n", m.Name, absDir)
fmt.Fprintf(cmd.OutOrStdout(), "It will be loaded on next dws invocation.\n")
fmt.Fprintf(cmd.OutOrStdout(), "To unregister: dws plugin dev --off %s\n", m.Name)
return nil
},
}
cmd.Flags().Bool("off", false, "Unregister a dev plugin")
return cmd
}
func newPluginConfigCommand() *cobra.Command {
configCmd := newPlaceholderParent("config", i18n.T("管理插件配置"))
configCmd.AddCommand(
newPluginConfigSetCommand(),
newPluginConfigGetCommand(),
newPluginConfigListCommand(),
newPluginConfigUnsetCommand(),
)
return configCmd
}
func newPluginConfigSetCommand() *cobra.Command {
return &cobra.Command{
Use: "set <plugin-name> <key> <value>",
Short: i18n.T("设置插件配置项"),
Long: `Persistently set a configuration value for a plugin.
The value is stored in ~/.dws/settings.json and automatically injected
as an environment variable when the plugin is loaded.
Environment variables set by the user (e.g. via export) take precedence
over values stored in settings.json.`,
Example: ` dws plugin config set demo-devtool DASHSCOPE_API_KEY sk-xxx
dws plugin config set my-plugin API_ENDPOINT https://api.example.com`,
Args: cobra.ExactArgs(3),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
pluginName, key, value := args[0], args[1], args[2]
loader := plugin.NewLoader(RawVersion())
// Validate that the plugin exists.
plugins := loader.ListInstalled()
found := false
for _, p := range plugins {
if p.Name == pluginName {
found = true
break
}
}
if !found {
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found; use 'dws plugin list' to see installed plugins", pluginName))
}
loader.SetPluginConfig(pluginName, key, value)
fmt.Fprintf(cmd.OutOrStdout(), "Config saved: %s.%s\n", pluginName, key)
return nil
},
}
}
func newPluginConfigGetCommand() *cobra.Command {
return &cobra.Command{
Use: "get <plugin-name> <key>",
Short: i18n.T("读取插件配置项"),
Args: cobra.ExactArgs(2),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
pluginName, key := args[0], args[1]
loader := plugin.NewLoader(RawVersion())
val, ok := loader.GetPluginConfig(pluginName, key)
if !ok {
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
}
fmt.Fprintln(cmd.OutOrStdout(), val)
return nil
},
}
}
func newPluginConfigListCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "list <plugin-name>",
Short: i18n.T("列出插件所有配置项"),
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
pluginName := args[0]
loader := plugin.NewLoader(RawVersion())
wantJSON, _ := cmd.Flags().GetBool("json")
configs := loader.ListPluginConfig(pluginName)
// Also load the plugin manifest to show declared userConfig keys.
declaredKeys := loadDeclaredUserConfig(loader, pluginName)
if wantJSON {
result := make(map[string]any)
for k, v := range configs {
sensitive := false
if ci, ok := declaredKeys[k]; ok {
sensitive = ci.Sensitive
}
if sensitive {
result[k] = maskSensitiveValue(v)
} else {
result[k] = v
}
}
// Include declared but unset keys.
for k, ci := range declaredKeys {
if _, set := configs[k]; !set {
entry := map[string]any{
"value": nil,
"description": ci.Description,
"required": ci.Default == "",
}
result[k] = entry
}
}
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
"kind": "plugin_config",
"plugin": pluginName,
"config": result,
})
}
w := cmd.OutOrStdout()
if len(configs) == 0 && len(declaredKeys) == 0 {
fmt.Fprintf(w, "No configuration for plugin %q.\n", pluginName)
return nil
}
fmt.Fprintf(w, "Configuration for %s:\n\n", pluginName)
// Show set values.
for k, v := range configs {
sensitive := false
if ci, ok := declaredKeys[k]; ok {
sensitive = ci.Sensitive
}
displayVal := v
if sensitive {
displayVal = maskSensitiveValue(v)
}
fmt.Fprintf(w, " %s = %s\n", k, displayVal)
}
// Show declared but unset keys.
for k, ci := range declaredKeys {
if _, set := configs[k]; !set {
desc := ""
if ci.Description != "" {
desc = " # " + ci.Description
}
fmt.Fprintf(w, " %s = (not set)%s\n", k, desc)
}
}
return nil
},
}
cmd.Flags().Bool("json", false, "Output in JSON format")
return cmd
}
func newPluginConfigUnsetCommand() *cobra.Command {
return &cobra.Command{
Use: "unset <plugin-name> <key>",
Short: i18n.T("删除插件配置项"),
Args: cobra.ExactArgs(2),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
pluginName, key := args[0], args[1]
loader := plugin.NewLoader(RawVersion())
if !loader.UnsetPluginConfig(pluginName, key) {
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
}
fmt.Fprintf(cmd.OutOrStdout(), "Config removed: %s.%s\n", pluginName, key)
return nil
},
}
}
// loadDeclaredUserConfig loads the userConfig section from a plugin's manifest.
func loadDeclaredUserConfig(loader *plugin.Loader, pluginName string) map[string]plugin.ConfigItem {
plugins := loader.ListInstalled()
for _, p := range plugins {
if p.Name == pluginName {
m, err := plugin.ParseManifest(filepath.Join(p.Path, "plugin.json"))
if err != nil {
return nil
}
return m.UserConfig
}
}
return nil
}
// maskSensitiveValue masks a sensitive value, showing only the first 4
// and last 2 characters for values longer than 8 characters.
func maskSensitiveValue(value string) string {
if len(value) <= 8 {
return strings.Repeat("*", len(value))
}
return value[:4] + strings.Repeat("*", len(value)-6) + value[len(value)-2:]
}
func newPluginBuildCommand() *cobra.Command {
return &cobra.Command{
Use: "build <dir>",
Short: i18n.T("将插件 stdio server 编译为原生二进制"),
Long: `Runs the build command declared in plugin.json to compile the
plugin's server into a single executable. This ensures plugin users
don't need any language runtime (Node.js, Python, etc.) installed.
The build configuration is read from the "build" field in plugin.json:
{
"build": {
"command": "bun build --compile src/server.ts --outfile bin/server",
"output": "bin/server"
}
}`,
Example: ` dws plugin build ./my-plugin
dws plugin build .`,
Args: cobra.ExactArgs(1),
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
dir := args[0]
absDir, err := filepath.Abs(dir)
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
}
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
if err != nil {
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
}
if m.Build == nil {
return apperrors.NewValidation(fmt.Sprintf(
"plugin %q has no \"build\" field in plugin.json.\n"+
"Add a build config, e.g.:\n\n"+
" \"build\": {\n"+
" \"command\": \"bun build --compile src/server.js --outfile bin/server\",\n"+
" \"output\": \"bin/server\"\n"+
" }", m.Name))
}
if err := plugin.BuildPlugin(absDir); err != nil {
return apperrors.NewInternal(err.Error())
}
fmt.Fprintf(cmd.OutOrStdout(), "Build succeeded: %s\n", m.Build.Output)
return nil
},
}
}
func statusStr(enabled bool) string {
if enabled {
return "enabled"
}
return "disabled"
}
@@ -0,0 +1,198 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"fmt"
"sync"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
)
// TestSharedCacheStoreConcurrentSaveTools verifies that a single *cache.Store
// instance is safe for goroutines saving tool snapshots concurrently, as long
// as each goroutine targets a distinct (partition, serverKey). This mirrors
// the real plugin discovery path where each goroutine owns one plugin/server.
//
// Each call serializes to its own "<key>.json.tmp" file followed by a
// rename(2) to the final path, so concurrent writers targeting distinct keys
// never collide. The invariant asserted here: after N parallel writes, the
// Store returns each written snapshot intact under LoadTools.
func TestSharedCacheStoreConcurrentSaveTools(t *testing.T) {
const (
partition = "default/default"
writers = 16
)
store := cache.NewStore(t.TempDir())
var wg sync.WaitGroup
for i := 0; i < writers; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
key := fmt.Sprintf("plugin:concurrent:%d", idx)
if err := store.SaveTools(partition, key, cache.ToolsSnapshot{
ServerKey: key,
}); err != nil {
t.Errorf("SaveTools(%s): %v", key, err)
}
}(i)
}
wg.Wait()
for i := 0; i < writers; i++ {
key := fmt.Sprintf("plugin:concurrent:%d", i)
snapshot, _, err := store.LoadTools(partition, key)
if err != nil {
t.Fatalf("LoadTools(%s): %v", key, err)
}
if snapshot.ServerKey != key {
t.Errorf("LoadTools(%s) returned ServerKey %q", key, snapshot.ServerKey)
}
}
}
// TestAppendDynamicServerConcurrent exercises the dynamicMu mutex on the
// write path by spraying distinct server descriptors in parallel. Afterwards
// every injected product ID must be resolvable — a missing entry would
// indicate a lost write through an un-synchronized map update.
func TestAppendDynamicServerConcurrent(t *testing.T) {
dynamicMu.Lock()
prev := struct {
endpoints map[string]string
products map[string]bool
aliases map[string]string
toolEndpoints map[string]string
}{dynamicEndpoints, dynamicProducts, dynamicAliases, dynamicToolEndpoints}
dynamicEndpoints = nil
dynamicProducts = nil
dynamicAliases = nil
dynamicToolEndpoints = nil
dynamicMu.Unlock()
t.Cleanup(func() {
dynamicMu.Lock()
dynamicEndpoints = prev.endpoints
dynamicProducts = prev.products
dynamicAliases = prev.aliases
dynamicToolEndpoints = prev.toolEndpoints
dynamicMu.Unlock()
})
const n = 32
var wg sync.WaitGroup
for i := 0; i < n; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
id := fmt.Sprintf("plugin-id-%d", idx)
endpoint := fmt.Sprintf("https://example.test/%d", idx)
AppendDynamicServer(market.ServerDescriptor{
Endpoint: endpoint,
CLI: market.CLIOverlay{
ID: id,
Command: id,
},
})
}(i)
}
wg.Wait()
for i := 0; i < n; i++ {
id := fmt.Sprintf("plugin-id-%d", i)
if endpoint, ok := directRuntimeEndpoint(id, ""); !ok || endpoint == "" {
t.Errorf("directRuntimeEndpoint(%q) = (%q, %v), want non-empty", id, endpoint, ok)
}
}
}
// TestRegisterStdioClientConcurrent verifies the stdioMu-protected registry
// survives concurrent writers — every registered client must be looked up
// afterwards. Uses nil client pointers since LookupStdioClient only compares
// keys, not values.
func TestRegisterStdioClientConcurrent(t *testing.T) {
stdioMu.Lock()
prev := stdioClients
stdioClients = make(map[string]*transport.StdioClient)
stdioMu.Unlock()
t.Cleanup(func() {
stdioMu.Lock()
stdioClients = prev
stdioMu.Unlock()
})
const n = 32
var wg sync.WaitGroup
for i := 0; i < n; i++ {
wg.Add(1)
go func(idx int) {
defer wg.Done()
RegisterStdioClient(fmt.Sprintf("plugin/%d", idx), nil)
}(i)
}
wg.Wait()
for i := 0; i < n; i++ {
key := fmt.Sprintf("plugin/%d", i)
if _, ok := LookupStdioClient(key); !ok {
t.Errorf("LookupStdioClient(%q) missing after concurrent registration", key)
}
}
}
// TestResolvePluginColdTimeouts covers the three code paths of the env
// parser: unset (defaults), valid duration (applied to all three slots),
// and invalid duration (logged and ignored, defaults returned).
func TestResolvePluginColdTimeouts(t *testing.T) {
t.Run("defaults when env unset", func(t *testing.T) {
t.Setenv(cli.PluginColdTimeoutEnv, "")
got := resolvePluginColdTimeouts()
if got.httpNoAuth != 1*time.Second {
t.Errorf("httpNoAuth = %v, want 1s", got.httpNoAuth)
}
if got.httpAuth != 1500*time.Millisecond {
t.Errorf("httpAuth = %v, want 1.5s", got.httpAuth)
}
if got.stdio != 2*time.Second {
t.Errorf("stdio = %v, want 2s", got.stdio)
}
})
t.Run("env override applies to all slots", func(t *testing.T) {
t.Setenv(cli.PluginColdTimeoutEnv, "3500ms")
got := resolvePluginColdTimeouts()
want := 3500 * time.Millisecond
if got.httpNoAuth != want || got.httpAuth != want || got.stdio != want {
t.Errorf("override not propagated: %+v", got)
}
})
t.Run("invalid env falls back to defaults", func(t *testing.T) {
t.Setenv(cli.PluginColdTimeoutEnv, "not-a-duration")
got := resolvePluginColdTimeouts()
if got.httpNoAuth != 1*time.Second || got.stdio != 2*time.Second {
t.Errorf("invalid env should not override defaults: %+v", got)
}
})
t.Run("non-positive env falls back to defaults", func(t *testing.T) {
t.Setenv(cli.PluginColdTimeoutEnv, "0")
got := resolvePluginColdTimeouts()
if got.httpNoAuth != 1*time.Second {
t.Errorf("zero duration should not override defaults, got %v", got.httpNoAuth)
}
})
}
+764 -36
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"
@@ -37,10 +41,14 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pat"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline/handlers"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
"github.com/spf13/pflag"
)
@@ -50,14 +58,19 @@ type outputFileContextKey struct{}
const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
// Execute runs the root command and returns the process exit code.
func Execute() int {
totalStart := time.Now()
func Execute() (exitCode int) {
defer func() {
if r := recover(); r != nil {
fmt.Fprintf(os.Stderr, "Error: internal panic: %v\n", r)
exitCode = 5
}
}()
timing := NewTimingCollector()
defer func() {
StopAllStdioClients() // Ensure child processes are terminated on exit
timing.PrintIfEnabled()
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] Execute total: %v\n", time.Since(totalStart))
}
timing.WriteReportIfEnabled(RawVersion(), SanitizeCommand(os.Args))
}()
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
@@ -70,24 +83,14 @@ func Execute() int {
recovery.ResetRuntimeState()
engine := newPipelineEngine()
root := NewRootCommandWithEngine(ctx, engine)
initDuration := time.Since(initStart)
timing.Record("cmd_init", initDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] command init: %v\n", initDuration)
}
timing.Record("cmd_init", time.Since(initStart))
// Run PreParse handlers on raw argv before Cobra parses flags.
// This corrects model-generated errors like --userId → --user-id
// and --limit100 → --limit 100.
pipeline.RunPreParse(root, engine)
execStart := time.Now()
executed, err := root.ExecuteC()
execDuration := time.Since(execStart)
timing.Record("cobra_exec", execDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cobra ExecuteC: %v\n", execDuration)
}
if err != nil {
if executed == nil {
executed = root
@@ -138,8 +141,13 @@ func flagErrorWithSuggestions(cmd *cobra.Command, err error) error {
}
func printExecutionError(root *cobra.Command, stdout, stderr io.Writer, err error) error {
var raw apperrors.RawStderrError
if stderrors.As(err, &raw) {
_, writeErr := fmt.Fprintln(stderr, raw.RawStderr())
return writeErr
}
if wantsJSONErrors(root) {
return apperrors.PrintJSON(stdout, err)
return apperrors.PrintJSON(stderr, err)
}
return apperrors.PrintHumanAt(stderr, err, resolveVerbosity(root))
}
@@ -230,6 +238,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
AuthTokenFunc: func(ctx context.Context) string {
return resolveRuntimeAuthToken(ctx, "")
},
LoggerFunc: FileLoggerInstance,
}
runner := newCommandRunnerWithFlags(loader, flags)
@@ -256,9 +265,16 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
// Configure global slog level based on --debug / --verbose flags.
configureLogLevel(flags)
return configureOutputSink(cmd)
if err := configureOutputSink(cmd); err != nil {
return err
}
if fn := edition.Get().AfterPersistentPreRun; fn != nil {
return fn(cmd, args)
}
return nil
},
PersistentPostRunE: func(cmd *cobra.Command, args []string) error {
StopAllStdioClients()
CloseFileLogger()
return closeOutputSink(cmd)
},
@@ -276,17 +292,40 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
newAuthCommand(),
newSkillCommand(),
newCacheCommand(),
newConfigCommand(),
newDoctorCommand(),
newCompletionCommand(root),
newRecoveryCommand(rootCtx, loader, flags),
newUpgradeCommand(),
newVersionCommand(),
newPluginCommand(),
schemaCmd,
genSkillsCmd,
mcpCmd,
}
root.AddCommand(utilityCommands...)
root.AddCommand(newLegacyPublicCommands(rootCtx, runner)...)
root.AddCommand(newLegacyHiddenCommands(runner)...)
// --- Plugin loading: runs AFTER legacy commands so that
// AppendDynamicServer adds plugin endpoints on top of Market
// endpoints (SetDynamicServers is called inside loadDynamicCommands).
pluginCmds := loadPlugins(engine, runner)
if len(pluginCmds) > 0 {
addPluginCommandsSafe(root, pluginCmds)
}
// PAT authorization commands (open-source core)
patCaller := newToolCallerAdapter(runner, flags)
pat.RegisterCommands(root, patCaller)
if fn := edition.Get().RegisterExtraCommands; fn != nil {
caller := newToolCallerAdapter(runner, flags)
fn(root, caller)
deduplicateCommands(root)
}
hideNonDirectRuntimeCommands(root)
configureRootHelp(root)
// Set custom flag error handler for better UX
@@ -465,24 +504,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
},
}
}
@@ -579,17 +645,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()
@@ -606,11 +689,126 @@ 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)
}
// pluginColdTimeouts holds the cold-path discovery budget for plugin MCP
// servers. Timeouts only apply to the *first* discovery for a given
// plugin/server; subsequent startups take the warm cache path and bypass
// the network entirely.
type pluginColdTimeouts struct {
httpNoAuth time.Duration
httpAuth time.Duration
stdio time.Duration
}
// resolvePluginColdTimeouts returns the cold-discovery budget for plugin MCP
// servers, applying the DWS_PLUGIN_COLD_TIMEOUT override when set. Defaults
// are tuned so healthy cross-region HTTP endpoints succeed on a cold start
// and Python/Node-based stdio plugins have headroom for interpreter load,
// while an unreachable host still surrenders in bounded time.
func resolvePluginColdTimeouts() pluginColdTimeouts {
t := pluginColdTimeouts{
httpNoAuth: 1 * time.Second,
httpAuth: 1500 * time.Millisecond,
stdio: 2 * time.Second,
}
raw := strings.TrimSpace(os.Getenv(cli.PluginColdTimeoutEnv))
if raw == "" {
return t
}
d, err := time.ParseDuration(raw)
if err != nil || d <= 0 {
slog.Warn("plugin: ignoring invalid DWS_PLUGIN_COLD_TIMEOUT",
"value", raw, "error", err)
return t
}
t.httpNoAuth = d
t.httpAuth = d
t.stdio = d
return t
}
func configureOutputSink(cmd *cobra.Command) error {
if local := cmd.LocalFlags().Lookup("output"); local != nil {
return nil
@@ -867,11 +1065,535 @@ func CloseFileLogger() {
}
}
// loadPlugins scans plugin directories, injects their MCP servers into
// the dynamic server registry, and registers their pipeline hooks.
// This runs before legacy command construction so that plugin servers
// are available for EnvironmentLoader.Load().
func loadPlugins(engine *pipeline.Engine, runner executor.Runner) []*cobra.Command {
pluginLoader := plugin.NewLoader(RawVersion())
// 0a. Inject plugin config values from settings.json as environment
// variables so that expandPluginVars can resolve ${KEY} references
// in plugin.json headers, endpoints, etc. User-set env vars take
// precedence (InjectPluginConfigEnv skips already-set keys).
pluginLoader.InjectPluginConfigEnv()
// Load TokenData once; reused for stdio injection below.
tokenData, _ := authpkg.LoadTokenData(defaultConfigDir())
var userCtx *plugin.UserContext
if tokenData != nil {
// Inject user context if either UserID or CorpID is present.
if tokenData.UserID != "" || tokenData.CorpID != "" {
userCtx = &plugin.UserContext{
UserID: tokenData.UserID,
CorpID: tokenData.CorpID,
}
}
}
// 1. Load plugins from the legacy managed/ directory (backward compat
// for plugins installed by older CLI builds).
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})
}
}
}
// Collect all stdio clients up front so HTTP + stdio discovery can run
// concurrently — the slowest plugin (typically an unreachable HTTP
// endpoint hitting its dial timeout) dominates the parallel wall-clock,
// not the sum of every plugin's cold timeout.
type stdioEntry struct {
plugin *plugin.Plugin
sc plugin.StdioServerClient
}
var stdioEntries []stdioEntry
for _, p := range allPlugins {
for _, sc := range p.StdioClients(userCtx) {
// Use background context so the subprocess lives for the CLI
// process lifetime (not killed by a short timeout).
if err := sc.Client.Start(context.Background()); err != nil {
slog.Warn("plugin: failed to start stdio server",
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
continue
}
stdioEntries = append(stdioEntries, stdioEntry{plugin: p, sc: sc})
}
}
// Share one cache.Store across all discovery goroutines. Each goroutine
// writes to a distinct serverKey path ("tools/<plugin>_<server>.json") with
// atomic tmp+rename, so concurrent writes to different keys never collide
// on the filesystem. Global in-process registries (AppendDynamicServer,
// RegisterStdioClient) carry their own sync.Mutex; see direct_runtime.go
// and stdio_registry.go.
sharedStore := cacheStoreFromEnv()
coldTimeouts := resolvePluginColdTimeouts()
// Fan out HTTP and stdio discovery in parallel. Each goroutine resolves
// its cache hit locally (no network) or runs a bounded cold-path probe.
// Wall-clock cost ≈ max(individual plugin latencies), not the sum.
httpResults := make([][]*cobra.Command, len(httpServers))
stdioResults := make([][]*cobra.Command, len(stdioEntries))
var wg sync.WaitGroup
for i, ps := range httpServers {
wg.Add(1)
go func(idx int, ps pluginServer) {
defer wg.Done()
httpResults[idx] = registerHTTPServer(ps.plugin, ps.srv, tc, runner, sharedStore, coldTimeouts)
}(i, ps)
}
for i, e := range stdioEntries {
wg.Add(1)
go func(idx int, e stdioEntry) {
defer wg.Done()
stdioResults[idx] = registerStdioServer(e.plugin, e.sc, runner, sharedStore, coldTimeouts)
}(i, e)
}
wg.Wait()
for _, cmds := range httpResults {
pluginCmds = append(pluginCmds, cmds...)
}
for _, cmds := range stdioResults {
pluginCmds = append(pluginCmds, cmds...)
}
// 5. Register plugin hooks into pipeline engine
if engine != nil {
for _, p := range allPlugins {
hooksCfg, err := p.LoadHooks()
if err != nil {
slog.Warn("plugin: failed to load hooks",
"plugin", p.Manifest.Name, "error", err)
continue
}
if hooksCfg == nil {
continue
}
for _, entry := range hooksCfg.Hooks {
engine.Register(plugin.NewHookAdapter(p.Manifest.Name, entry))
}
}
}
// 7. Sync plugin skills to agent directories
plugin.SyncSkills(allPlugins)
if len(allPlugins) > 0 {
slog.Debug("plugins loaded",
"managed", len(managedPlugins),
"user", len(userPlugins),
"dev", len(devPlugins),
)
}
return pluginCmds
}
// pluginCacheKey derives the cache key used to persist a plugin MCP server's
// tool list. Prefixed with "plugin:" so entries are namespaced apart from the
// Market-derived cache, and visible distinctly via `dws cache status`.
func pluginCacheKey(pluginName, serverKey string) string {
return "plugin:" + pluginName + ":" + serverKey
}
// registerHTTPServer discovers tools from a streamable-http MCP server and
// builds CLI commands. Used for plugin-owned HTTP servers that provide CLI metadata.
//
// Startup-latency strategy (issue #119):
// - Warm cache: build commands from the persisted tools snapshot
// synchronously — no network I/O. `dws --help` returns in ms even when
// the plugin endpoint is unreachable.
// - Cold cache: synchronous discovery (Initialize + ListTools) with a tight
// timeout. The outcome — success or failure — is persisted so the next
// invocation hits the warm path. Refresh on demand via `dws cache clean`
// / `dws cache refresh`; the cache TTL (7d) otherwise expires naturally.
//
// When the server descriptor carries AuthHeaders (from plugin.json "headers"),
// a dedicated transport.Client is created with the plugin's Bearer token and
// trusted domains so that third-party MCP servers requiring independent
// authentication (e.g. Alibaba Cloud Bailian) can be discovered at startup.
func registerHTTPServer(p *plugin.Plugin, srv market.ServerDescriptor, tc *transport.Client, runner executor.Runner, store *cache.Store, timeouts pluginColdTimeouts) []*cobra.Command {
partition := config.DefaultPartition
cacheKey := pluginCacheKey(p.Manifest.Name, srv.Key)
if snapshot, freshness, err := store.LoadTools(partition, cacheKey); err == nil {
slog.Debug("plugin: http server served from cache",
"plugin", p.Manifest.Name, "server", srv.Key,
"tools", len(snapshot.Tools), "freshness", string(freshness))
return buildHTTPCommandsFromTools(srv, snapshot.Tools, runner)
}
// Cold cache: synchronous discovery. Persist the outcome even on failure
// (empty tools == negative cache) so the next invocation takes the fast
// path regardless of endpoint health.
tools := discoverHTTPTools(p, srv, tc, timeouts)
_ = store.SaveTools(partition, cacheKey, cache.ToolsSnapshot{
ServerKey: cacheKey,
Tools: tools,
})
return buildHTTPCommandsFromTools(srv, tools, runner)
}
// discoverHTTPTools performs the blocking Initialize + ListTools handshake
// for an HTTP MCP server and returns the discovered tools. Returns nil on
// any transport/protocol error; errors are logged at Debug level.
func discoverHTTPTools(p *plugin.Plugin, srv market.ServerDescriptor, tc *transport.Client, timeouts pluginColdTimeouts) []transport.ToolDescriptor {
// Cold-path budget. An unreachable endpoint will burn the full window
// via the TCP dial timeout; a healthy localhost/third-party endpoint
// typically responds in <200 ms. Third-party servers with auth get a
// slightly larger window to accommodate TLS + auth RTT. Operators with
// cross-region endpoints can relax the window via DWS_PLUGIN_COLD_TIMEOUT.
// The outcome is persisted as a negative cache so subsequent startups
// (80 ms warm) are unaffected. See issue #119.
timeout := timeouts.httpNoAuth
if len(srv.AuthHeaders) > 0 {
timeout = timeouts.httpAuth
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
discoveryClient := tc
if len(srv.AuthHeaders) > 0 {
discoveryClient = buildPluginAuthClient(tc, srv)
}
if _, err := discoveryClient.Initialize(ctx, srv.Endpoint); err != nil {
slog.Debug("plugin: http server offline, skipping tool discovery",
"plugin", p.Manifest.Name, "server", srv.Key)
return nil
}
toolsResult, err := discoveryClient.ListTools(ctx, srv.Endpoint)
if err != nil {
slog.Debug("plugin: http ListTools failed",
"plugin", p.Manifest.Name, "server", srv.Key, "error", err)
return nil
}
return toolsResult.Tools
}
// buildHTTPCommandsFromTools converts a tool list into Cobra commands via
// the BuildDynamicCommands path. Returns nil for an empty tool list.
func buildHTTPCommandsFromTools(srv market.ServerDescriptor, tools []transport.ToolDescriptor, runner executor.Runner) []*cobra.Command {
if len(tools) == 0 {
return nil
}
detailsByID := make(map[string][]market.DetailTool)
var detailTools []market.DetailTool
for _, tool := range tools {
schemaJSON := ""
if tool.InputSchema != nil {
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
schemaJSON = string(data)
}
}
detailTools = append(detailTools, market.DetailTool{
ToolName: tool.Name,
ToolTitle: tool.Title,
ToolDesc: tool.Description,
IsSensitive: tool.Sensitive,
ToolRequest: schemaJSON,
})
}
detailsByID[strings.TrimSpace(srv.CLI.ID)] = detailTools
// If the server has no ToolOverrides (e.g. third-party MCP servers that
// only declare cli.id and cli.command), auto-generate one override per
// discovered tool so BuildDynamicCommands can create leaf commands.
if len(srv.CLI.ToolOverrides) == 0 {
srv.CLI.ToolOverrides = make(map[string]market.CLIToolOverride, len(tools))
for _, tool := range tools {
srv.CLI.ToolOverrides[tool.Name] = market.CLIToolOverride{
CLIName: deriveToolCLIName(tool.Name),
}
}
}
return compat.BuildDynamicCommands(
[]market.ServerDescriptor{srv}, runner, detailsByID)
}
// deriveToolCLIName converts an MCP tool name (e.g. "web_search" or
// "maps.search_poi") into a kebab-case CLI command name ("search" or
// "search-poi"). It strips common prefixes and replaces underscores/dots
// with hyphens.
func deriveToolCLIName(toolName string) string {
// Use the last segment after "." as the base name.
if idx := strings.LastIndex(toolName, "."); idx >= 0 {
toolName = toolName[idx+1:]
}
// Replace underscores with hyphens for kebab-case.
return strings.ReplaceAll(toolName, "_", "-")
}
// buildPluginAuthClient creates a transport.Client copy with the plugin's
// Bearer token and trusted domains injected. This allows third-party MCP
// servers that require independent authentication to be discovered at startup.
func buildPluginAuthClient(base *transport.Client, srv market.ServerDescriptor) *transport.Client {
authToken := ""
extraHeaders := make(map[string]string)
for key, value := range srv.AuthHeaders {
if strings.EqualFold(key, "Authorization") {
authToken = strings.TrimPrefix(value, "Bearer ")
authToken = strings.TrimSpace(authToken)
} else {
extraHeaders[key] = value
}
}
if authToken == "" {
return base
}
client := base.WithAuth(authToken, extraHeaders)
// Trust the endpoint's hostname so the token is actually sent.
if parsed, err := url.Parse(srv.Endpoint); err == nil {
host := parsed.Hostname()
client.TrustedDomains = []string{host, "*." + host}
}
return client
}
// registerPluginAuthFromHeaders extracts authentication credentials from
// a server descriptor's AuthHeaders and registers them in the global
// PluginAuth registry. The runner uses this registry at execution time
// to inject the correct Bearer token for third-party MCP servers.
func registerPluginAuthFromHeaders(srv market.ServerDescriptor) {
authToken := ""
extraHeaders := make(map[string]string)
for key, value := range srv.AuthHeaders {
if strings.EqualFold(key, "Authorization") {
authToken = strings.TrimPrefix(value, "Bearer ")
authToken = strings.TrimSpace(authToken)
} else {
extraHeaders[key] = value
}
}
if authToken == "" {
return
}
var trustedDomains []string
if parsed, err := url.Parse(srv.Endpoint); err == nil {
host := parsed.Hostname()
trustedDomains = []string{host, "*." + host}
}
productID := strings.TrimSpace(srv.CLI.ID)
if productID == "" {
productID = srv.Key
}
RegisterPluginAuth(productID, &PluginAuth{
Token: authToken,
ExtraHeaders: extraHeaders,
TrustedDomains: trustedDomains,
})
}
// registerStdioServer initializes a stdio MCP server, discovers its tools
// via ListTools, builds CLI commands, and registers the StdioClient for
// runtime dispatch. Returns generated cobra commands.
//
// Warm-cache fast path (issue #119): when a tools snapshot is already cached
// for this plugin/server, skip the Initialize + ListTools RPC round-trip and
// rebuild commands directly from the snapshot. Cold cache falls back to
// synchronous discovery with a 4s cap and persists the outcome.
func registerStdioServer(p *plugin.Plugin, sc plugin.StdioServerClient, runner executor.Runner, store *cache.Store, timeouts pluginColdTimeouts) []*cobra.Command {
partition := config.DefaultPartition
cacheKey := pluginCacheKey(p.Manifest.Name, sc.Key)
if snapshot, freshness, err := store.LoadTools(partition, cacheKey); err == nil {
slog.Debug("plugin: stdio server served from cache",
"plugin", p.Manifest.Name, "server", sc.Key,
"tools", len(snapshot.Tools), "freshness", string(freshness))
return buildStdioCommands(p, sc, snapshot.Tools, runner)
}
tools := discoverStdioTools(p, sc, timeouts)
_ = store.SaveTools(partition, cacheKey, cache.ToolsSnapshot{
ServerKey: cacheKey,
Tools: tools,
})
return buildStdioCommands(p, sc, tools, runner)
}
// discoverStdioTools performs the blocking Initialize + ListTools handshake
// on a stdio MCP subprocess. Returns nil on any error (logged at Warn level).
// The default 2s budget comfortably accommodates Python/Node runtimes whose
// interpreter + dependency load dominates the first response. Operators with
// heavier startup chains can relax further via DWS_PLUGIN_COLD_TIMEOUT.
func discoverStdioTools(p *plugin.Plugin, sc plugin.StdioServerClient, timeouts pluginColdTimeouts) []transport.ToolDescriptor {
ctx, cancel := context.WithTimeout(context.Background(), timeouts.stdio)
defer cancel()
if _, err := sc.Client.Initialize(ctx); err != nil {
slog.Warn("plugin: stdio initialize failed",
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
return nil
}
toolsResult, err := sc.Client.ListTools(ctx)
if err != nil {
slog.Warn("plugin: stdio ListTools failed",
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
return nil
}
return toolsResult.Tools
}
// buildStdioCommands constructs Cobra commands from a tool list and
// registers the runtime dispatch state (StdioClient + dynamic server).
// Returns nil for an empty tool list.
func buildStdioCommands(p *plugin.Plugin, sc plugin.StdioServerClient, tools []transport.ToolDescriptor, runner executor.Runner) []*cobra.Command {
if len(tools) == 0 {
slog.Debug("plugin: stdio server has no tools",
"plugin", p.Manifest.Name, "server", sc.Key)
return nil
}
// Build CLIOverlay: use manifest CLI metadata if present, else auto-generate.
serverID := sc.Key
overlay := market.CLIOverlay{
ID: serverID,
Command: serverID,
}
if srv, ok := p.Manifest.MCPServers[sc.Key]; ok && len(srv.CLI) > 0 {
cliData := srv.CLI
// If cli is a JSON string, treat it as a relative file path to an overlay file.
if len(cliData) > 0 && cliData[0] == '"' {
var cliPath string
if err := json.Unmarshal(cliData, &cliPath); err == nil && cliPath != "" {
absPath := filepath.Join(p.Root, cliPath)
if fileData, readErr := os.ReadFile(absPath); readErr == nil {
cliData = fileData
} else {
slog.Warn("plugin: failed to read CLI overlay file",
"plugin", p.Manifest.Name, "path", absPath, "error", readErr)
}
}
}
if err := json.Unmarshal(cliData, &overlay); err != nil {
slog.Warn("plugin: failed to parse CLI overlay for stdio server",
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
}
if overlay.ID == "" {
overlay.ID = serverID
}
if overlay.Command == "" {
overlay.Command = serverID
}
}
// Auto-generate ToolOverrides from discovered tools when not provided.
if len(overlay.ToolOverrides) == 0 {
overlay.ToolOverrides = make(map[string]market.CLIToolOverride)
if len(overlay.Prefixes) == 0 {
overlay.Prefixes = []string{serverID}
}
for _, tool := range tools {
overlay.ToolOverrides[tool.Name] = market.CLIToolOverride{
IsSensitive: tool.Sensitive,
}
}
}
// Construct virtual endpoint and server descriptor.
endpoint := StdioEndpoint(p.Manifest.Name, sc.Key)
source := "plugin"
if p.IsManaged {
source = "plugin-managed"
}
descriptor := market.ServerDescriptor{
Key: sc.Key,
DisplayName: p.Manifest.Name + "/" + sc.Key,
Description: p.Manifest.Description,
Endpoint: endpoint,
Source: source,
CLI: overlay,
HasCLIMeta: true,
}
AppendDynamicServer(descriptor)
// Register with pluginName/serverKey format for cleanup by plugin name
RegisterStdioClient(p.Manifest.Name+"/"+serverID, sc.Client)
// Convert tool descriptors to DetailTool entries for flag generation.
detailsByID := make(map[string][]market.DetailTool)
var detailTools []market.DetailTool
for _, tool := range tools {
schemaJSON := ""
if tool.InputSchema != nil {
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
schemaJSON = string(data)
}
}
detailTools = append(detailTools, market.DetailTool{
ToolName: tool.Name,
ToolTitle: tool.Title,
ToolDesc: tool.Description,
IsSensitive: tool.Sensitive,
ToolRequest: schemaJSON,
})
}
detailsByID[serverID] = detailTools
cmds := compat.BuildDynamicCommands(
[]market.ServerDescriptor{descriptor}, runner, detailsByID)
slog.Debug("plugin: stdio server registered",
"plugin", p.Manifest.Name, "server", sc.Key,
"tools", len(tools), "commands", len(cmds))
return cmds
}
// newPipelineEngine creates and configures the pipeline engine with
// the standard set of handlers for model input correction.
// handlers for all five pipeline phases. The phases execute in order:
// Register → PreParse → PostParse → PreRequest → PostResponse.
//
// Phases are invoked at their respective integration points:
// - Register: during command tree construction (newMCPCommand)
// - PreParse: before Cobra parses raw argv (RunPreParse)
// - PostParse: after Cobra parsing, before validation (canonical RunE)
// - PreRequest: after validation, before JSON-RPC dispatch (canonical RunE)
// - PostResponse: after transport returns, before stdout (canonical RunE)
func newPipelineEngine() *pipeline.Engine {
engine := pipeline.NewEngine()
engine.RegisterAll(
// Register handler runs during command tree building.
handlers.RegisterHandler{},
// PreParse handlers run in order: alias → sticky → paramname.
// Alias normalises case first (--userId → --user-id), then
// sticky splits glued values (--limit100 → --limit 100), then
@@ -882,6 +1604,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
}
+112 -16
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()
@@ -43,11 +51,11 @@ func TestPrintExecutionErrorDefaultsToJSON(t *testing.T) {
if err != nil {
t.Fatalf("printExecutionError() error = %v", err)
}
if stderr.Len() != 0 {
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
if stdout.Len() != 0 {
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
}
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
}
}
@@ -65,11 +73,11 @@ func TestPrintExecutionErrorUsesJSONWhenFormatIsJSON(t *testing.T) {
if err != nil {
t.Fatalf("printExecutionError() error = %v", err)
}
if stderr.Len() != 0 {
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
if stdout.Len() != 0 {
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
}
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
}
}
@@ -97,11 +105,11 @@ func TestPrintExecutionErrorUsesJSONWhenCommandSetsJSONFlag(t *testing.T) {
if err != nil {
t.Fatalf("printExecutionError() error = %v", err)
}
if stderr.Len() != 0 {
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
if stdout.Len() != 0 {
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
}
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
}
}
@@ -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,86 @@ func TestNestedShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
t.Fatalf("nested short help output missing command title:\n%s", out.String())
}
}
func TestPrintExecutionError_RawStderrError_writes_raw_JSON_to_stderr(t *testing.T) {
t.Parallel()
rawJSON := `{"success":false,"code":"PAT_LOW_RISK_NO_PERMISSION","data":{}}`
err := &patLikeError{raw: rawJSON}
root := NewRootCommand()
var stdout, stderr bytes.Buffer
writeErr := printExecutionError(root, &stdout, &stderr, err)
if writeErr != nil {
t.Fatalf("printExecutionError() error = %v", writeErr)
}
if stdout.Len() != 0 {
t.Fatalf("stdout = %q, want empty for RawStderrError", stdout.String())
}
got := strings.TrimSpace(stderr.String())
if got != rawJSON {
t.Fatalf("stderr = %q, want raw JSON %q", got, rawJSON)
}
}
func TestPrintExecutionError_RawStderrError_exit_code_is_4(t *testing.T) {
t.Parallel()
err := &patLikeError{raw: `{"code":"PAT_MEDIUM_RISK_NO_PERMISSION"}`}
exitCode := apperrors.ExitCode(err)
if exitCode != 4 {
t.Fatalf("apperrors.ExitCode(patLikeError) = %d, want 4", exitCode)
}
}
func TestPrintExecutionError_RawStderrError_takes_precedence_over_JSON_mode(t *testing.T) {
t.Parallel()
rawJSON := `{"success":false,"code":"PAT_HIGH_RISK_NO_PERMISSION"}`
err := &patLikeError{raw: rawJSON}
root := NewRootCommand()
_ = root.PersistentFlags().Set("format", "json")
var stdout, stderr bytes.Buffer
writeErr := printExecutionError(root, &stdout, &stderr, err)
if writeErr != nil {
t.Fatalf("printExecutionError() error = %v", writeErr)
}
if stdout.Len() != 0 {
t.Fatalf("stdout = %q, want empty — RawStderrError should bypass JSON mode", stdout.String())
}
if !strings.Contains(stderr.String(), "PAT_HIGH_RISK_NO_PERMISSION") {
t.Fatalf("stderr = %q, want raw PAT JSON", stderr.String())
}
}
// simulateExecuteWithPanic mirrors the recovery pattern in Execute():
// named return + defer recover → exitCode = 5 on panic.
func simulateExecuteWithPanic(doPanic bool) (exitCode int) {
defer func() {
if r := recover(); r != nil {
exitCode = 5
}
}()
if doPanic {
panic("test panic")
}
return 0
}
func TestExecute_panic_recovery_returns_exit_5(t *testing.T) {
t.Parallel()
code := simulateExecuteWithPanic(true)
if code != 5 {
t.Fatalf("panic recovery exitCode = %d, want 5", code)
}
}
func TestExecute_no_panic_returns_0(t *testing.T) {
t.Parallel()
code := simulateExecuteWithPanic(false)
if code != 0 {
t.Fatalf("no-panic exitCode = %d, want 0", code)
}
}
+73 -2
View File
@@ -5,6 +5,8 @@ import (
"strings"
"text/tabwriter"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
@@ -13,6 +15,26 @@ func configureRootHelp(root *cobra.Command) {
return
}
// Replace the cobra-default English help command with a localized one so
// that both its listing short (shown in `dws --help`) and its own
// `dws help --help` long text follow the active locale.
root.SetHelpCommand(&cobra.Command{
Use: "help [command]",
Short: i18n.T("查看任意命令的帮助信息"),
Long: i18n.T("显示任意命令的帮助文案。\n" +
"用法:dws help [命令路径] 查看完整说明。"),
DisableAutoGenTag: true,
Run: func(c *cobra.Command, args []string) {
target, _, err := c.Root().Find(args)
if target == nil || err != nil {
c.Root().HelpFunc()(c.Root(), args)
return
}
target.InitDefaultHelpFlag()
_ = target.Help()
},
})
defaultHelpFunc := root.HelpFunc()
root.SetHelpFunc(func(cmd *cobra.Command, args []string) {
if cmd != root {
@@ -25,6 +47,7 @@ func configureRootHelp(root *cobra.Command) {
func renderRootHelp(root *cobra.Command) {
services := visibleMCPRootCommands(root)
utilities := visibleUtilityRootCommands(root)
w := root.OutOrStdout()
if len(services) == 0 {
@@ -44,8 +67,21 @@ func renderRootHelp(root *cobra.Command) {
_, _ = fmt.Fprintln(w, "Usage:")
_, _ = fmt.Fprintln(w, " dws <service> [command] [flags]")
if len(utilities) > 0 {
_, _ = fmt.Fprintln(w, " dws <command> [flags]")
}
_, _ = fmt.Fprintln(w)
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service.`)
if len(utilities) > 0 {
_, _ = fmt.Fprintln(w, "Utility Commands:")
_, _ = fmt.Fprintln(w)
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
for _, utility := range utilities {
_, _ = fmt.Fprintf(tw, " %s\t%s\n", utility.Name(), strings.TrimSpace(utility.Short))
}
_ = tw.Flush()
_, _ = fmt.Fprintln(w)
}
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service or "dws <command> --help" for utility commands.`)
}
func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
@@ -53,7 +89,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 +115,29 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
}
return commands
}
func visibleUtilityRootCommands(root *cobra.Command) []*cobra.Command {
if root == nil {
return nil
}
productCommands := DirectRuntimeProductIDs()
if fn := edition.Get().VisibleProducts; fn != nil {
productCommands = make(map[string]bool, len(fn()))
for _, product := range fn() {
productCommands[product] = true
}
}
commands := make([]*cobra.Command, 0)
for _, cmd := range root.Commands() {
if cmd == nil || cmd.Hidden {
continue
}
if productCommands[cmd.Name()] {
continue
}
commands = append(commands, cmd)
}
return commands
}
+280 -46
View File
@@ -15,9 +15,10 @@ package app
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"os"
@@ -29,10 +30,54 @@ import (
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/safety"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
func init() {
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_RUNTIME_CONTENT_SCAN",
Category: configmeta.CategoryRuntime,
Description: "启用 MCP 响应内容安全扫描",
Example: "true",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_RUNTIME_CONTENT_SCAN_ENFORCE",
Category: configmeta.CategoryRuntime,
Description: "内容安全扫描发现问题时阻断响应",
Example: "true",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_RUNTIME_CONTENT_SCAN_REPORT",
Category: configmeta.CategoryRuntime,
Description: "在 JSON 输出中包含安全扫描报告",
Example: "true",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DINGTALK_AGENT",
Category: configmeta.CategoryExternal,
Description: "MCP 请求 x-dingtalk-agent 头",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DINGTALK_TRACE_ID",
Category: configmeta.CategoryExternal,
Description: "MCP 请求 x-dingtalk-trace-id 头",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DINGTALK_SESSION_ID",
Category: configmeta.CategoryExternal,
Description: "MCP 请求 x-dingtalk-session-id 头",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DINGTALK_MESSAGE_ID",
Category: configmeta.CategoryExternal,
Description: "MCP 请求 x-dingtalk-message-id 头",
})
}
const (
runtimeContentScanEnv = "DWS_RUNTIME_CONTENT_SCAN"
runtimeContentScanEnforceEnv = "DWS_RUNTIME_CONTENT_SCAN_ENFORCE"
@@ -43,9 +88,21 @@ const (
envDingtalkTraceID = "DINGTALK_TRACE_ID"
envDingtalkSessionID = "DINGTALK_SESSION_ID"
envDingtalkMessageID = "DINGTALK_MESSAGE_ID"
// Environment variables for third-party channel integration
envDWSChannel = "DWS_CHANNEL"
)
func newCommandRunnerWithFlags(loader cli.CatalogLoader, flags *GlobalFlags) executor.Runner {
// Ensure DWS_CLIENT_ID env is populated from persisted config before
// resolveIdentityHeaders reads it. This covers fresh-process cold starts
// where no env var has been inherited from a parent process.
if os.Getenv("DWS_CLIENT_ID") == "" {
if cid := authpkg.ClientID(); cid != "" {
_ = os.Setenv("DWS_CLIENT_ID", cid)
}
}
var httpClient *http.Client
if flags != nil && flags.Timeout > 0 {
httpClient = &http.Client{Timeout: time.Duration(flags.Timeout) * time.Second}
@@ -75,13 +132,6 @@ type runtimeRunner struct {
}
func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
totalStart := time.Now()
defer func() {
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] runtimeRunner.Run total: %v\n", time.Since(totalStart))
}
}()
if r.loader == nil || r.transport == nil {
return r.fallback.Run(ctx, invocation)
}
@@ -95,6 +145,11 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
return r.executeInvocation(ctx, endpoint, invocation)
}
// Prefetch the Keychain token in the background. Keychain access costs
// ~70ms on macOS; starting it here lets the load overlap with endpoint
// resolution and catalog loading below.
go getCachedRuntimeToken(ctx)
if shouldUseDirectRuntime(invocation) {
if endpoint, ok := directRuntimeEndpoint(invocation.CanonicalProduct, invocation.Tool); ok {
return r.executeInvocation(ctx, endpoint, invocation)
@@ -105,7 +160,10 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
catalog, err := r.loader.Load(ctx)
RecordTiming(ctx, "catalog_load", time.Since(catalogStart))
if err != nil {
return executor.Result{}, err
var degraded *cli.CatalogDegraded
if !errors.As(err, &degraded) {
return executor.Result{}, err
}
}
product, ok := catalog.FindProduct(invocation.CanonicalProduct)
@@ -126,21 +184,60 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
return r.executeInvocation(ctx, endpoint, invocation)
}
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (result executor.Result, retErr error) {
// Route stdio:// endpoints to the local StdioClient — no HTTP, no auth.
if IsStdioEndpoint(endpoint) {
return r.executeStdioInvocation(ctx, invocation)
}
invokeStart := time.Now()
execID := generateExecutionID()
r.transport.ExecutionId = execID
// Lazy bind FileLogger: it may be nil at construction time because
// configureLogLevel runs later in PersistentPreRunE.
if r.transport.FileLogger == nil {
r.transport.FileLogger = FileLoggerInstance()
}
authStart := time.Now()
authToken := r.resolveAuthToken(ctx)
authDuration := time.Since(authStart)
RecordTiming(ctx, "auth_token", authDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] resolveAuthToken: %v\n", authDuration)
fl := r.transport.FileLogger
defer func() {
var errCat, errReason string
if retErr != nil {
var typed *apperrors.Error
if errors.As(retErr, &typed) {
errCat = string(typed.Category)
errReason = typed.Reason
} else {
errCat = "unknown"
errReason = retErr.Error()
}
}
logging.LogCommandEnd(fl, execID,
invocation.CanonicalProduct, invocation.Tool,
retErr == nil, time.Since(invokeStart), errCat, errReason)
}()
// Check if this product has plugin-level auth credentials registered.
// If so, use the plugin's token instead of the default DingTalk OAuth token.
// This allows third-party MCP servers (e.g. Bailian) to use their own API keys.
pluginAuth, hasPluginAuth := LookupPluginAuth(invocation.CanonicalProduct)
authToken := ""
if hasPluginAuth {
authToken = pluginAuth.Token
} else {
authToken = r.resolveAuthToken(ctx)
}
var timeoutSec int
if r.globalFlags != nil {
timeoutSec = r.globalFlags.Timeout
}
logging.LogCommandStart(fl, execID,
invocation.CanonicalProduct, invocation.Tool, endpoint, version, authToken != "", timeoutSec)
if invocation.DryRun {
return executor.Result{
Invocation: invocation,
@@ -181,23 +278,79 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
)
}
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
var tc *transport.Client
if hasPluginAuth {
// Use plugin-level auth: inject the plugin's token and trust its domains.
tc = r.transport.WithAuth(authToken, pluginAuth.ExtraHeaders)
tc.TrustedDomains = pluginAuth.TrustedDomains
} else {
// Default path: use DingTalk OAuth token with identity headers.
tc = r.transport.WithAuth(authToken, resolveIdentityHeaders())
}
callCtx := ctx
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
var cancel context.CancelFunc
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
defer cancel()
}
callStart := time.Now()
callResult, err := tc.CallTool(ctx, endpoint, invocation.Tool, invocation.Params)
callDuration := time.Since(callStart)
RecordTiming(ctx, "mcp_call", callDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] MCP CallTool: %v\n", callDuration)
}
callResult, err := tc.CallTool(callCtx, endpoint, invocation.Tool, invocation.Params)
RecordTiming(ctx, "mcp_call", time.Since(callStart))
if err != nil {
if isAuthError(err) {
if fn := edition.Get().OnAuthError; fn != nil {
if overrideErr := fn(defaultConfigDir(), err); overrideErr != nil {
captureRuntimeFailure(invocation, err, overrideErr)
return executor.Result{}, overrideErr
}
}
}
// PAT scope error: offer human-readable output and retry after authorization
if isPatScopeError(err) {
scopeErr := extractPatScopeError(err)
captureRuntimeFailure(invocation, err, err)
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
}
captureRuntimeFailure(invocation, err, err)
return executor.Result{}, err
}
// ---- Edition hook gets first dibs (preserves overlay PATError passthrough) ----
if fn := edition.Get().ClassifyToolResult; fn != nil {
if editionErr := fn(callResult.Content); editionErr != nil {
if patCheck := apperrors.AsPatAuthCheckError(editionErr); patCheck != nil {
if IsPatRetrying(ctx) {
return executor.Result{}, patCheck // already retried once, don't loop
}
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
}
return executor.Result{}, editionErr
}
}
// ---- Structured PAT auth check (open-source fallback) ----
if patCheck := apperrors.ClassifyPatAuthCheck(callResult.Content); patCheck != nil {
if IsPatRetrying(ctx) {
return executor.Result{}, patCheck // already retried once, don't loop
}
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
}
if callResult.IsError {
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
// ClassifyToolResult hook: let the overlay intercept known error
// patterns (PAT permission, gateway-auth) before generic handling.
if classify := edition.Get().ClassifyToolResult; classify != nil {
if hookErr := classify(callResult.Content); hookErr != nil {
captureRuntimeFailure(invocation, hookErr, hookErr)
return executor.Result{}, hookErr
}
}
mcpErr := apperrors.NewAPI(
extractMCPErrorMessage(callResult),
apperrors.WithOperation("tools/call"),
@@ -206,6 +359,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
apperrors.WithHint("MCP tool returned a business error; check tool parameters and refer to skill documentation."),
apperrors.WithServerDiag(diag),
)
// PAT scope error in business response: offer human-readable output and retry
if isPatScopeError(mcpErr) {
scopeErr := extractPatScopeError(mcpErr)
captureRuntimeFailure(invocation, mcpErr, mcpErr)
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
}
captureRuntimeFailure(invocation, mcpErr, mcpErr)
return executor.Result{}, mcpErr
}
@@ -238,12 +397,78 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
return executor.Result{Invocation: invocation, Response: response}, nil
}
// executeStdioInvocation dispatches a tool call through a local StdioClient
// subprocess instead of the HTTP transport. This is used for plugin stdio
// servers whose endpoints use the stdio:// scheme.
func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
if invocation.DryRun {
return executor.Result{
Invocation: invocation,
Response: map[string]any{
"dry_run": true,
"transport": "stdio",
"request": executor.ToolCallRequest(invocation.Tool, invocation.Params),
"note": "execution skipped by --dry-run",
},
}, nil
}
client, ok := LookupStdioClient(invocation.CanonicalProduct)
if !ok {
return executor.Result{}, apperrors.NewInternal(
fmt.Sprintf("stdio client not found for %q", invocation.CanonicalProduct))
}
callCtx := ctx
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
var cancel context.CancelFunc
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
defer cancel()
}
callResult, err := client.CallTool(callCtx, invocation.Tool, invocation.Params)
if err != nil {
return executor.Result{}, apperrors.NewAPI(
fmt.Sprintf("stdio call failed: %v", err),
apperrors.WithOperation("tools/call"),
apperrors.WithReason("stdio_error"),
)
}
if callResult.IsError {
return executor.Result{}, apperrors.NewAPI(
extractMCPErrorMessage(callResult),
apperrors.WithOperation("tools/call"),
apperrors.WithReason("mcp_tool_error"),
apperrors.WithServerKey(invocation.CanonicalProduct),
)
}
invocation.Implemented = true
return executor.Result{
Invocation: invocation,
Response: map[string]any{
"transport": "stdio",
"content": callResult.Content,
},
}, nil
}
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
explicitToken := ""
if r != nil && r.globalFlags != nil {
explicitToken = r.globalFlags.Token
}
return resolveRuntimeAuthToken(ctx, explicitToken)
if token := strings.TrimSpace(explicitToken); token != "" {
return token
}
if tp := edition.Get().TokenProvider; tp != nil {
token, _ := tp(ctx, func() (string, error) {
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
})
return token
}
return getCachedRuntimeToken(ctx)
}
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
@@ -265,38 +490,30 @@ var (
func getCachedRuntimeToken(ctx context.Context) string {
cachedRuntimeTokenOnce.Do(func() {
loadStart := time.Now()
defer func() {
loadDuration := time.Since(loadStart)
RecordTiming(ctx, "keychain_load", loadDuration)
if os.Getenv("DWS_PERF_DEBUG") != "" {
_, _ = fmt.Fprintf(os.Stderr, "[PERF] getCachedRuntimeToken (first load): %v\n", loadDuration)
}
}()
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
configDir := defaultConfigDir()
provider := authpkg.NewOAuthProvider(configDir, slog.New(slog.NewTextHandler(io.Discard, nil)))
configureOAuthProviderCompatibility(provider, configDir)
token, tokenErr := provider.GetAccessToken(ctx)
if tokenErr == nil && strings.TrimSpace(token) != "" {
cachedRuntimeToken = strings.TrimSpace(token)
return
}
// If the error is a decryption failure (corrupted data), log and bail out
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
slog.Error(tokenErr.Error())
return
}
// Try legacy manager as fallback
manager := authpkg.NewManager(configDir, nil)
configureLegacyAuthManagerCompatibility(manager)
if token, _, err := manager.GetToken(); err == nil && strings.TrimSpace(token) != "" {
cachedRuntimeToken = strings.TrimSpace(token)
return
if token != "" {
cachedRuntimeToken = token
}
})
return cachedRuntimeToken
}
// generateExecutionID returns a random 16-char hex string used to correlate
// all log entries (command_start, jsonrpc_request, command_end, etc.) belonging
// to a single command invocation.
func generateExecutionID() string {
b := make([]byte, 8)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
// This should be called after login/logout operations.
func ResetRuntimeTokenCache() {
@@ -335,6 +552,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))
@@ -353,7 +578,7 @@ func resolveIdentityHeaders() map[string]string {
headers = make(map[string]string)
}
// Inject environment variable based headers for MCP gateway tracking
// Inject environment variable based headers for MCP gateway tracking.
envHeaders := map[string]string{
"x-dingtalk-agent": os.Getenv(envDingtalkAgent),
"x-dingtalk-trace-id": os.Getenv(envDingtalkTraceID),
@@ -365,6 +590,15 @@ func resolveIdentityHeaders() map[string]string {
headers[k] = v
}
}
// Inject third-party channel headers
if v := os.Getenv(envDWSChannel); v != "" {
headers["x-dws-channel"] = v
}
if fn := edition.Get().MergeHeaders; fn != nil {
headers = fn(headers)
}
return headers
}
+139
View File
@@ -17,6 +17,7 @@ import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"os"
@@ -25,12 +26,61 @@ import (
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
)
func setupRuntimeCommandTest(t *testing.T) {
t.Helper()
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
discoverySrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
_ = json.NewEncoder(w).Encode(contactDiscoveryResponse())
}))
t.Cleanup(func() { discoverySrv.Close() })
SetDiscoveryBaseURL(discoverySrv.URL)
t.Cleanup(func() { SetDiscoveryBaseURL("") })
}
func contactDiscoveryResponse() map[string]any {
return map[string]any{
"metadata": map[string]any{"count": 1, "nextCursor": ""},
"servers": []any{
map[string]any{
"server": map[string]any{
"name": "Contact",
"description": "通讯录",
"remotes": []any{
map[string]any{
"type": "streamable-http",
"url": "https://mcp.dingtalk.com/contact/v1",
},
},
},
"_meta": map[string]any{
"com.dingtalk.mcp.registry/metadata": map[string]any{
"status": "active", "isLatest": true,
},
"com.dingtalk.mcp.registry/cli": map[string]any{
"id": "contact",
"command": "contact",
"groups": map[string]any{
"user": map[string]any{
"description": "用户管理",
},
},
"toolOverrides": map[string]any{
"get_current_user_profile": map[string]any{
"cliName": "get-self",
"group": "user",
"flags": map[string]any{},
},
},
},
},
},
},
}
}
func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
@@ -596,6 +646,95 @@ func contentScanServer() *mockmcp.Server {
return mockmcp.MustNewServer(fixture)
}
func TestClassifyToolResultHookPreemptsBusinessError(t *testing.T) {
setupRuntimeCommandTest(t)
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
t.Setenv("DWS_TRUSTED_DOMAINS", "*")
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
http.Error(w, "bad request", http.StatusBadRequest)
return
}
method, _ := req["method"].(string)
switch method {
case "initialize":
_ = json.NewEncoder(w).Encode(map[string]any{
"jsonrpc": "2.0",
"id": req["id"],
"result": map[string]any{
"protocolVersion": "2025-03-26",
"capabilities": map[string]any{"tools": map[string]any{"listChanged": false}},
"serverInfo": map[string]any{"name": "doc", "version": "1.0.0"},
},
})
case "notifications/initialized":
w.WriteHeader(http.StatusNoContent)
case "tools/list":
_ = json.NewEncoder(w).Encode(map[string]any{
"jsonrpc": "2.0",
"id": req["id"],
"result": map[string]any{
"tools": []map[string]any{{
"name": "search_documents",
"title": "Search",
"description": "Search documents",
"inputSchema": map[string]any{"type": "object"},
}},
},
})
case "tools/call":
_ = json.NewEncoder(w).Encode(map[string]any{
"jsonrpc": "2.0",
"id": req["id"],
"result": map[string]any{
"content": map[string]any{
"success": false,
"code": "PAT_LOW_RISK_NO_PERMISSION",
"data": map[string]any{"requiredScopes": []any{}},
},
},
})
}
}))
defer server.Close()
hookCalled := false
sentinelMsg := "hook-intercepted-PAT"
edition.Override(&edition.Hooks{
ClassifyToolResult: func(content map[string]any) error {
if code, ok := content["code"].(string); ok && strings.Contains(code, "PAT") {
hookCalled = true
return fmt.Errorf("%s", sentinelMsg)
}
return nil
},
})
t.Cleanup(func() { edition.Override(&edition.Hooks{}) })
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.URL, false))
cmd := NewRootCommand()
cmd.SetOut(&bytes.Buffer{})
cmd.SetErr(&bytes.Buffer{})
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
err := cmd.Execute()
if err == nil {
t.Fatal("Execute() error = nil, want hook sentinel error")
}
if !hookCalled {
t.Fatal("ClassifyToolResult hook was not called")
}
if !strings.Contains(err.Error(), sentinelMsg) {
t.Fatalf("error = %q, want hook sentinel %q (not generic business error)", err.Error(), sentinelMsg)
}
if strings.Contains(err.Error(), "business_error") || strings.Contains(err.Error(), "mcp_tool_error") {
t.Fatalf("error = %q, should NOT contain generic framework error category", err.Error())
}
}
func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
setupRuntimeCommandTest(t)
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
+273 -18
View File
@@ -19,7 +19,9 @@ import (
"encoding/json"
"fmt"
"io"
"mime"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
@@ -27,10 +29,24 @@ import (
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
"github.com/spf13/cobra"
)
func init() {
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_SKILL_API_HOST",
Category: configmeta.CategoryNetwork,
Description: "覆盖 Skill API 地址",
DefaultValue: "https://mcp.dingtalk.com",
Example: "https://custom-mcp.example.com",
})
}
const (
// legacySkillAPIHost is the legacy skill market host used by the old cli.
legacySkillAPIHost = "https://mcp.dingtalk.com"
// skillDownloadEndpoint is the API endpoint for downloading skills.
skillDownloadEndpoint = "https://aihub.dingtalk.com/cli/download"
// skillDownloadTimeout is the timeout for skill download operations.
@@ -51,6 +67,22 @@ type downloadSkillResult struct {
FileName string `json:"fileName"`
}
// findSkillsResponse represents the legacy skill search API response.
type findSkillsResponse struct {
Success bool `json:"success"`
ErrorCode string `json:"errorCode,omitempty"`
ErrorMsg string `json:"errorMsg,omitempty"`
Result []CliSkillDTO `json:"result,omitempty"`
}
// CliSkillDTO mirrors the old cli response payload for `skill search`.
type CliSkillDTO struct {
SkillID string `json:"skillId"`
Name string `json:"name"`
Desc string `json:"desc"`
Icon string `json:"icon"`
}
// agentSkillPaths maps target names to their relative skill installation paths.
// These paths are relative to the user's home directory.
var agentSkillPaths = map[string]string{
@@ -75,7 +107,7 @@ func buildSkillCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "skill",
Short: "技能管理",
Long: "管理钉钉技能市场的技能。支持下载和安装技能到指定 Agent 目录。",
Long: "管理钉钉技能市场的技能。支持搜索、下载与安装到指定 Agent 目录。",
Args: cobra.NoArgs,
TraverseChildren: true,
DisableAutoGenTag: true,
@@ -84,13 +116,61 @@ func buildSkillCommand() *cobra.Command {
},
}
cmd.AddCommand(newSkillAddCommand())
cmd.AddCommand(
newSkillInstallCommand(),
newSkillGetCommand(),
newSkillSearchCommand(),
newSkillFindHintCommand(),
newSkillAddHintCommand(),
)
return cmd
}
func newSkillAddCommand() *cobra.Command {
func newSkillGetCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "add <skillId> <target>",
Use: "get",
Short: "获取技能压缩文件",
Long: "从服务端下载技能包到本地临时目录。命令执行成功后会输出临时目录路径,供调用方使用。",
Example: " dws skill get --skill-id <skillId>",
DisableAutoGenTag: true,
RunE: runSkillGet,
}
cmd.Flags().String("skill-id", "", "技能 ID(必填)")
_ = cmd.MarkFlagRequired("skill-id")
return cmd
}
func newSkillSearchCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "search",
Short: "从钉钉技能市场搜索技能",
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
Example: " dws skill search --query 关键词",
DisableAutoGenTag: true,
RunE: runSkillFind,
}
cmd.Flags().String("query", "", "搜索关键词(必填)")
_ = cmd.MarkFlagRequired("query")
cmd.Flags().String("scopes", "", "查询范围,空格分隔。备选值:DingtalkMarket(钉钉市场)、OrgInternal(企业内部)。为空默认查市场技能")
return cmd
}
func newSkillFindHintCommand() *cobra.Command {
return &cobra.Command{
Use: "find",
Short: "兼容旧用法,提示使用 skill search",
Hidden: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill search --query <关键词>")
return nil
},
}
}
func newSkillInstallCommand() *cobra.Command {
cmd := &cobra.Command{
Use: "install <skillId> <target>",
Short: "下载并安装技能到指定目录",
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
@@ -107,9 +187,9 @@ func newSkillAddCommand() *cobra.Command {
. -> 当前目录
示例:
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
dws skill install skill-123 qoder # 安装到 ~/.qoder/skills/
dws skill install skill-123 claude # 安装到 ~/.claude/skills/
dws skill install skill-123 . # 安装到当前目录`, supportedTargets()),
Args: cobra.ExactArgs(2),
DisableAutoGenTag: true,
RunE: runSkillAdd,
@@ -118,6 +198,96 @@ func newSkillAddCommand() *cobra.Command {
return cmd
}
func newSkillAddHintCommand() *cobra.Command {
return &cobra.Command{
Use: "add",
Short: "兼容旧用法,提示使用 skill install",
Hidden: true,
DisableAutoGenTag: true,
RunE: func(cmd *cobra.Command, args []string) error {
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill install <skillId> <target>")
return nil
},
}
}
func runSkillGet(cmd *cobra.Command, args []string) error {
skillID, _ := cmd.Flags().GetString("skill-id")
accessToken, err := loadSkillAccessToken()
if err != nil {
return err
}
apiURL := fmt.Sprintf("%s/cli/install?skillId=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(skillID)))
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "⬇️ 下载技能包...")
tmpDir, err := downloadSkillToTmpDir(cmd.Context(), apiURL, accessToken)
if err != nil {
return err
}
_, _ = fmt.Fprintln(cmd.OutOrStdout(), tmpDir)
return nil
}
func runSkillFind(cmd *cobra.Command, args []string) error {
keyword, _ := cmd.Flags().GetString("query")
scopes, _ := cmd.Flags().GetString("scopes")
accessToken, err := loadSkillAccessToken()
if err != nil {
return err
}
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
if scopes != "" {
apiURL += "&scopes=" + url.QueryEscape(scopes)
}
req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, apiURL, nil)
if err != nil {
return apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
}
req.Header.Set("x-user-access-token", accessToken)
client := &http.Client{Timeout: 30 * time.Second}
resp, err := client.Do(req)
if err != nil {
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %v", err), apperrors.WithRetryable(true))
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return parseLegacySkillAPIError(resp)
}
var result findSkillsResponse
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
return apperrors.NewAPI(fmt.Sprintf("failed to parse search response: %v", err))
}
if !result.Success {
errMsg := strings.TrimSpace(result.ErrorMsg)
if errMsg == "" {
errMsg = strings.TrimSpace(result.ErrorCode)
}
if errMsg == "" {
errMsg = "unknown error"
}
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %s", errMsg))
}
if len(result.Result) == 0 {
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "未找到匹配的技能")
return nil
}
for _, skill := range result.Result {
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "SkillID: %s\n", skill.SkillID)
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Name: %s\n", skill.Name)
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Desc: %s\n", skill.Desc)
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "---")
}
return nil
}
func runSkillAdd(cmd *cobra.Command, args []string) error {
skillID := strings.TrimSpace(args[0])
target := strings.TrimSpace(args[1])
@@ -132,13 +302,9 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
}
// Load auth token
configDir := defaultConfigDir()
tokenData, err := authpkg.LoadTokenData(configDir)
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
apperrors.WithHint("请先执行 'dws auth login' 登录"),
apperrors.WithActions("dws auth login"))
accessToken, err := loadSkillAccessToken()
if err != nil {
return err
}
ctx, cancel := context.WithTimeout(cmd.Context(), skillDownloadTimeout)
@@ -148,7 +314,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
// Step 1: Get download URL from API
fmt.Fprintf(w, "正在获取技能信息...\n")
downloadResp, err := fetchSkillDownloadInfo(ctx, tokenData.AccessToken, skillID)
downloadResp, err := fetchSkillDownloadInfo(ctx, accessToken, skillID)
if err != nil {
return err
}
@@ -189,6 +355,33 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
return nil
}
func loadSkillAccessToken() (string, error) {
configDir := defaultConfigDir()
tokenData, err := authpkg.LoadTokenData(configDir)
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
return "", skillAuthError()
}
return tokenData.AccessToken, nil
}
func skillAuthError() error {
if edition.Get().IsEmbedded {
return apperrors.NewAuth("认证信息已失效",
apperrors.WithReason("not_authenticated"),
apperrors.WithHint("请先完成钉钉账号登录后重试"))
}
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
apperrors.WithHint("请先执行 'dws auth login' 登录"),
apperrors.WithActions("dws auth login"))
}
func skillAPIHost() string {
if override := strings.TrimSpace(os.Getenv("DWS_SKILL_API_HOST")); override != "" {
return strings.TrimRight(override, "/")
}
return legacySkillAPIHost
}
// resolveSkillTargetPath resolves the target argument to an absolute path.
func resolveSkillTargetPath(target string) (string, error) {
target = strings.TrimSpace(target)
@@ -236,9 +429,7 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
defer resp.Body.Close()
if resp.StatusCode == http.StatusUnauthorized {
return nil, apperrors.NewAuth("authentication failed. Please run 'dws auth login' to refresh your token",
apperrors.WithHint("请执行 'dws auth login' 重新登录"),
apperrors.WithActions("dws auth login"))
return nil, skillAuthError()
}
if resp.StatusCode != http.StatusOK {
@@ -259,6 +450,70 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
return &result, nil
}
func downloadSkillToTmpDir(ctx context.Context, apiURL, accessToken string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
if err != nil {
return "", apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
}
req.Header.Set("x-user-access-token", accessToken)
client := &http.Client{Timeout: skillDownloadTimeout}
resp, err := client.Do(req)
if err != nil {
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill package: %v", err), apperrors.WithRetryable(true))
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", parseLegacySkillAPIError(resp)
}
tmpDir, err := os.MkdirTemp("", "dws-skill-*")
if err != nil {
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp dir: %v", err))
}
filename := filenameFromDisposition(resp.Header.Get("Content-Disposition"))
destPath := filepath.Join(tmpDir, filename)
file, err := os.Create(destPath)
if err != nil {
os.RemoveAll(tmpDir)
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
}
defer file.Close()
if _, err := io.Copy(file, resp.Body); err != nil {
os.RemoveAll(tmpDir)
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
}
return tmpDir, nil
}
func filenameFromDisposition(cd string) string {
if cd != "" {
if _, params, err := mime.ParseMediaType(cd); err == nil {
if name := strings.TrimSpace(params["filename"]); name != "" {
return name
}
}
}
return "skill.zip"
}
func parseLegacySkillAPIError(resp *http.Response) error {
switch resp.StatusCode {
case http.StatusUnauthorized:
return skillAuthError()
case http.StatusBadRequest:
return apperrors.NewValidation("request parameters are invalid")
case http.StatusNotFound:
return apperrors.NewValidation("skill does not exist or corresponding file was not found")
default:
return apperrors.NewAPI(fmt.Sprintf("skill API returned HTTP %d", resp.StatusCode),
apperrors.WithRetryable(resp.StatusCode >= 500))
}
}
// downloadSkillFile downloads the skill zip file to a temporary location.
func downloadSkillFile(ctx context.Context, downloadURL, fileName string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
+65 -13
View File
@@ -319,7 +319,7 @@ func TestExtractSkillZipPreventZipSlip(t *testing.T) {
}
}
func TestSkillAddCommandValidation(t *testing.T) {
func TestSkillInstallCommandValidation(t *testing.T) {
tests := []struct {
name string
args []string
@@ -328,19 +328,19 @@ func TestSkillAddCommandValidation(t *testing.T) {
}{
{
name: "missing arguments",
args: []string{"skill", "add"},
args: []string{"skill", "install"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
{
name: "missing target",
args: []string{"skill", "add", "skill-123"},
args: []string{"skill", "install", "skill-123"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
{
name: "too many arguments",
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
args: []string{"skill", "install", "skill-123", "qoder", "extra"},
wantErr: true,
errMsg: "accepts 2 arg(s)",
},
@@ -366,7 +366,7 @@ func TestSkillAddCommandValidation(t *testing.T) {
}
}
func TestSkillAddInvalidTarget(t *testing.T) {
func TestSkillInstallInvalidTarget(t *testing.T) {
// Setup: Create config directory with valid token
tempDir := t.TempDir()
configDir := filepath.Join(tempDir, "config")
@@ -380,11 +380,11 @@ func TestSkillAddInvalidTarget(t *testing.T) {
RefreshExpAt: time.Now().Add(24 * time.Hour),
})
if err != nil {
t.Fatalf("failed to save token data: %v", err)
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
}
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
cmd.SetArgs([]string{"skill", "install", "skill-123", "invalid-target"})
var out bytes.Buffer
cmd.SetOut(&out)
@@ -399,7 +399,7 @@ func TestSkillAddInvalidTarget(t *testing.T) {
}
}
func TestSkillAddRequiresAuth(t *testing.T) {
func TestSkillInstallRequiresAuth(t *testing.T) {
// Setup: Create config directory without token
tempDir := t.TempDir()
configDir := filepath.Join(tempDir, "config")
@@ -411,7 +411,7 @@ func TestSkillAddRequiresAuth(t *testing.T) {
}
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
cmd.SetArgs([]string{"skill", "install", "skill-123", "qoder"})
var out bytes.Buffer
cmd.SetOut(&out)
@@ -563,14 +563,16 @@ func TestSkillCommandHelp(t *testing.T) {
if !strings.Contains(output, "技能") {
t.Errorf("help should mention '技能', got: %s", output)
}
if !strings.Contains(output, "add") {
t.Errorf("help should mention 'add' subcommand, got: %s", output)
for _, subcmd := range []string{"install", "search", "get"} {
if !strings.Contains(output, subcmd) {
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
}
}
}
func TestSkillAddCommandHelp(t *testing.T) {
func TestSkillInstallCommandHelp(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "add", "--help"})
cmd.SetArgs([]string{"skill", "install", "--help"})
var out bytes.Buffer
cmd.SetOut(&out)
@@ -590,6 +592,56 @@ func TestSkillAddCommandHelp(t *testing.T) {
}
}
func TestSkillGetCommandValidation(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "get"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
err := cmd.Execute()
if err == nil {
t.Fatal("Execute() error = nil, want missing required flag error")
}
if !strings.Contains(err.Error(), "required flag") {
t.Fatalf("error = %v, want required flag message", err)
}
}
func TestSkillSearchCommandValidation(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "search"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
err := cmd.Execute()
if err == nil {
t.Fatal("Execute() error = nil, want missing required flag error")
}
if !strings.Contains(err.Error(), "required flag") {
t.Fatalf("error = %v, want required flag message", err)
}
}
func TestSkillFindHintCommand(t *testing.T) {
cmd := NewRootCommand()
cmd.SetArgs([]string{"skill", "find"})
var out bytes.Buffer
cmd.SetOut(&out)
cmd.SetErr(&out)
if err := cmd.Execute(); err != nil {
t.Fatalf("Execute() error = %v", err)
}
if !strings.Contains(out.String(), "dws skill search --query") {
t.Fatalf("output = %q, want legacy hint", out.String())
}
}
func TestDownloadSkillFileSuccess(t *testing.T) {
// Create a mock server that returns a zip file
expectedContent := []byte("fake zip content")
+119
View File
@@ -0,0 +1,119 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"log/slog"
"strings"
"sync"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
)
const stdioEndpointScheme = "stdio://"
var (
stdioMu sync.RWMutex
stdioClients = make(map[string]*transport.StdioClient)
)
// RegisterStdioClient stores a StdioClient keyed by its canonical product ID
// (the CLI.ID used in the server descriptor). The runner looks up this client
// when a stdio:// endpoint is resolved at execution time.
func RegisterStdioClient(productID string, client *transport.StdioClient) {
stdioMu.Lock()
defer stdioMu.Unlock()
stdioClients[productID] = client
}
// LookupStdioClient returns the StdioClient registered for the given product ID.
// The productID can be either the full key (pluginName/serverKey) or just the serverKey.
// This supports backward compatibility with existing CanonicalProduct values.
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
stdioMu.RLock()
defer stdioMu.RUnlock()
// Try exact match first
if c, ok := stdioClients[productID]; ok {
return c, true
}
// If not found, try matching by serverKey suffix (for backward compatibility)
for id, c := range stdioClients {
if idx := strings.LastIndex(id, "/"); idx >= 0 {
if id[idx+1:] == productID {
return c, true
}
}
}
return nil, false
}
// StdioEndpoint returns a virtual endpoint URL for a stdio-based MCP server.
// Format: stdio://{pluginName}/{serverKey}
func StdioEndpoint(pluginName, serverKey string) string {
return stdioEndpointScheme + pluginName + "/" + serverKey
}
// IsStdioEndpoint returns true if the endpoint uses the stdio:// scheme.
func IsStdioEndpoint(endpoint string) bool {
return strings.HasPrefix(endpoint, stdioEndpointScheme)
}
// StopAllStdioClients stops all registered stdio clients.
// This should be called on program exit to terminate child processes.
func StopAllStdioClients() {
stdioMu.Lock()
defer stdioMu.Unlock()
for id, client := range stdioClients {
if err := client.Stop(); err != nil {
slog.Warn("failed to stop stdio client", "id", id, "error", err)
}
}
stdioClients = make(map[string]*transport.StdioClient)
}
// StopStdioClient stops a specific stdio client by product ID.
// Returns true if the client was found and stopped, false otherwise.
func StopStdioClient(productID string) bool {
stdioMu.Lock()
defer stdioMu.Unlock()
client, ok := stdioClients[productID]
if !ok {
return false
}
if err := client.Stop(); err != nil {
slog.Warn("failed to stop stdio client", "id", productID, "error", err)
}
delete(stdioClients, productID)
return true
}
// StopStdioClientsByPlugin stops all stdio clients belonging to a plugin.
// The productID format is "pluginName/serverKey". This function stops all
// clients whose productID has the given pluginName prefix.
func StopStdioClientsByPlugin(pluginName string) int {
stdioMu.Lock()
defer stdioMu.Unlock()
prefix := pluginName + "/"
count := 0
for id, client := range stdioClients {
if len(id) > len(prefix) && id[:len(prefix)] == prefix {
if err := client.Stop(); err != nil {
slog.Warn("failed to stop stdio client", "id", id, "error", err)
}
delete(stdioClients, id)
count++
}
}
return count
}
+72
View File
@@ -0,0 +1,72 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
)
func TestStdioEndpoint(t *testing.T) {
endpoint := StdioEndpoint("hello-plugin", "hello")
want := "stdio://hello-plugin/hello"
if endpoint != want {
t.Errorf("StdioEndpoint() = %q, want %q", endpoint, want)
}
}
func TestIsStdioEndpoint(t *testing.T) {
tests := []struct {
endpoint string
want bool
}{
{"stdio://hello-plugin/hello", true},
{"stdio://conference/local", true},
{"https://mcp.dingtalk.com", false},
{"", false},
}
for _, tt := range tests {
if got := IsStdioEndpoint(tt.endpoint); got != tt.want {
t.Errorf("IsStdioEndpoint(%q) = %v, want %v", tt.endpoint, got, tt.want)
}
}
}
func TestStdioClientRegistry(t *testing.T) {
// Clean up after test
defer func() {
stdioMu.Lock()
delete(stdioClients, "test-product")
stdioMu.Unlock()
}()
// Initially not found
if _, ok := LookupStdioClient("test-product"); ok {
t.Error("expected LookupStdioClient to return false for unregistered product")
}
// Register a client
client := transport.NewStdioClient("echo", nil, nil)
RegisterStdioClient("test-product", client)
// Now should be found
got, ok := LookupStdioClient("test-product")
if !ok {
t.Fatal("expected LookupStdioClient to return true after registration")
}
if got != client {
t.Error("LookupStdioClient returned different client instance")
}
}
+215 -11
View File
@@ -15,16 +15,45 @@ package app
import (
"context"
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
)
// Environment variable to enable performance timing output.
const PerfTimingEnv = "DWS_PERF_TIMING"
func init() {
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_PERF_DEBUG",
Category: configmeta.CategoryDebug,
Description: "启用性能计时输出到 stderr",
Example: "1",
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_PERF_REPORT",
Category: configmeta.CategoryDebug,
Description: "JSON 性能报告输出路径 (auto=~/.dws/perf/latest.json)",
Example: "auto",
})
}
const (
// PerfDebugEnv is the environment variable to enable performance timing output.
PerfDebugEnv = "DWS_PERF_DEBUG"
// PerfReportEnv is the environment variable to enable JSON perf report output.
// Set to "auto" to write to ~/.dws/perf/latest.json, or a custom file path.
PerfReportEnv = "DWS_PERF_REPORT"
perfReportDir = "perf"
perfReportFile = "latest.json"
)
// timingContextKey is the context key for TimingCollector.
type timingContextKey struct{}
@@ -107,32 +136,46 @@ func (tc *TimingCollector) Entries() []TimingEntry {
return result
}
// formatDuration returns a human-friendly duration string.
// Sub-µs → "0µs", sub-ms → microsecond precision (e.g. "142µs"), else → ms.
func formatDuration(d time.Duration) string {
switch {
case d < time.Microsecond:
return "0µs"
case d < time.Millisecond:
return d.Truncate(time.Microsecond).String()
default:
return d.Truncate(time.Millisecond).String()
}
}
// Print writes a summary of all timing entries to the given writer.
func (tc *TimingCollector) Print(w io.Writer) {
if tc == nil || w == nil {
return
}
entries := tc.Entries()
total := tc.Total()
if len(entries) == 0 {
fmt.Fprintf(w, "\n[Timing] Total: %v (no detailed entries)\n", tc.Total().Truncate(time.Millisecond))
fmt.Fprintf(w, "\n[Perf] Total: %v (no detailed entries)\n", formatDuration(total))
return
}
fmt.Fprintln(w)
fmt.Fprintln(w, "[Timing] Execution breakdown:")
fmt.Fprintln(w, "[Perf] Execution breakdown:")
for _, e := range entries {
fmt.Fprintf(w, " %-30s %v\n", e.Name, e.Duration.Truncate(time.Millisecond))
fmt.Fprintf(w, " %-30s %v\n", e.Name, formatDuration(e.Duration))
}
fmt.Fprintf(w, " %-30s %v\n", "──────────────────────────────", "──────────")
fmt.Fprintf(w, " %-30s %v\n", "Total", tc.Total().Truncate(time.Millisecond))
fmt.Fprintf(w, " %-30s %v\n", "Total", formatDuration(total))
}
// PrintIfEnabled prints timing info to stderr if DWS_PERF_TIMING is set.
// PrintIfEnabled prints timing info to stderr if DWS_PERF_DEBUG is set.
func (tc *TimingCollector) PrintIfEnabled() {
if tc == nil {
return
}
if os.Getenv(PerfTimingEnv) == "" {
if os.Getenv(PerfDebugEnv) == "" {
return
}
tc.Print(os.Stderr)
@@ -171,7 +214,168 @@ func StartTiming(ctx context.Context, name string) func() {
return tc.StartTimer(name)
}
// IsPerfTimingEnabled returns true if performance timing output is enabled.
func IsPerfTimingEnabled() bool {
return os.Getenv(PerfTimingEnv) != ""
// IsPerfDebugEnabled returns true if performance debug output is enabled.
func IsPerfDebugEnabled() bool {
return os.Getenv(PerfDebugEnv) != ""
}
// ── Structured Performance Report ──────────────────────────────────────
// PerfPhase is a single phase in the performance report.
type PerfPhase struct {
Name string `json:"name"`
DurationMs int64 `json:"duration_ms"`
Seq int `json:"seq"`
}
// PerfReport is the JSON-serialisable performance report.
type PerfReport struct {
Kind string `json:"kind"`
Version string `json:"version"`
CLIVersion string `json:"cli_version"`
Command string `json:"command"`
Timestamp time.Time `json:"timestamp"`
TotalMs int64 `json:"total_ms"`
Phases []PerfPhase `json:"phases"`
Slowest string `json:"slowest"`
OverheadMs int64 `json:"overhead_ms"`
}
// BuildReport constructs a PerfReport from the collected timing entries.
func (tc *TimingCollector) BuildReport(cliVersion, command string) PerfReport {
entries := tc.Entries()
total := tc.Total()
totalMs := total.Milliseconds()
phases := make([]PerfPhase, len(entries))
var sumMs int64
var slowestName string
var slowestMs int64
for i, e := range entries {
ms := e.Duration.Milliseconds()
phases[i] = PerfPhase{
Name: e.Name,
DurationMs: ms,
Seq: e.Seq,
}
sumMs += ms
if ms > slowestMs {
slowestMs = ms
slowestName = e.Name
}
}
overhead := totalMs - sumMs
if overhead < 0 {
overhead = 0
}
return PerfReport{
Kind: "perf_report",
Version: "1",
CLIVersion: cliVersion,
Command: command,
Timestamp: time.Now(),
TotalMs: totalMs,
Phases: phases,
Slowest: slowestName,
OverheadMs: overhead,
}
}
// WriteReportIfEnabled checks DWS_PERF_REPORT and writes a JSON report if set.
func (tc *TimingCollector) WriteReportIfEnabled(cliVersion, command string) {
if tc == nil {
return
}
dest := os.Getenv(PerfReportEnv)
if dest == "" {
return
}
report := tc.BuildReport(cliVersion, command)
data, err := json.MarshalIndent(report, "", " ")
if err != nil {
return
}
path := resolvePerfReportPath(dest)
if path == "" {
return
}
dir := filepath.Dir(path)
if err := os.MkdirAll(dir, 0o700); err != nil {
return
}
tmp := path + ".tmp"
if err := os.WriteFile(tmp, data, 0o600); err != nil {
_ = os.Remove(tmp)
return
}
_ = os.Rename(tmp, path)
}
// LoadLatestReport reads the default perf report file (~/.dws/perf/latest.json).
func LoadLatestReport() (*PerfReport, error) {
path := defaultPerfReportPath()
data, err := os.ReadFile(path)
if err != nil {
return nil, err
}
var report PerfReport
if err := json.Unmarshal(data, &report); err != nil {
return nil, err
}
return &report, nil
}
// resolvePerfReportPath resolves the DWS_PERF_REPORT value to an absolute path.
func resolvePerfReportPath(dest string) string {
if dest == "auto" {
return defaultPerfReportPath()
}
return dest
}
func defaultPerfReportPath() string {
home, err := os.UserHomeDir()
if err != nil {
return ""
}
return filepath.Join(home, ".dws", perfReportDir, perfReportFile)
}
// sensitiveFlags are flag names whose values should be masked in commands.
var sensitiveFlags = map[string]bool{
"--token": true,
"--client-secret": true,
"--client-id": true,
}
// SanitizeCommand redacts sensitive flag values from a command arg slice.
func SanitizeCommand(args []string) string {
sanitized := make([]string, 0, len(args))
skipNext := false
for _, arg := range args {
if skipNext {
sanitized = append(sanitized, "***")
skipNext = false
continue
}
if idx := strings.IndexByte(arg, '='); idx > 0 {
key := arg[:idx]
if sensitiveFlags[key] {
sanitized = append(sanitized, key+"=***")
continue
}
}
if sensitiveFlags[arg] {
skipNext = true
}
sanitized = append(sanitized, arg)
}
return strings.Join(sanitized, " ")
}
+292 -12
View File
@@ -16,7 +16,9 @@ package app
import (
"bytes"
"context"
"encoding/json"
"os"
"path/filepath"
"strings"
"testing"
"time"
@@ -87,8 +89,8 @@ func TestTimingCollector_Print(t *testing.T) {
tc.Print(&buf)
output := buf.String()
if !strings.Contains(output, "[Timing]") {
t.Error("output should contain [Timing] header")
if !strings.Contains(output, "[Perf]") {
t.Error("output should contain [Perf] header")
}
if !strings.Contains(output, "auth_token") {
t.Error("output should contain 'auth_token'")
@@ -103,8 +105,8 @@ func TestTimingCollector_Print(t *testing.T) {
func TestTimingCollector_PrintIfEnabled(t *testing.T) {
// Set environment variable
os.Setenv(PerfTimingEnv, "1")
defer os.Unsetenv(PerfTimingEnv)
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
tc := NewTimingCollector()
tc.Record("test_op", 10*time.Millisecond)
@@ -136,6 +138,7 @@ func TestTimingCollector_ContextIntegration(t *testing.T) {
}
func TestTimingCollectorFromContext_NilContext(t *testing.T) {
//lint:ignore SA1012 Testing explicit nil-context guard in TimingCollectorFromContext.
tc := TimingCollectorFromContext(nil)
if tc != nil {
t.Error("TimingCollectorFromContext(nil) should return nil")
@@ -156,18 +159,295 @@ func TestStartTiming_NoCollector(t *testing.T) {
stop()
}
func TestIsPerfTimingEnabled(t *testing.T) {
func TestIsPerfDebugEnabled(t *testing.T) {
// Clear the env var first
os.Unsetenv(PerfTimingEnv)
os.Unsetenv(PerfDebugEnv)
if IsPerfTimingEnabled() {
t.Error("IsPerfTimingEnabled should return false when env var is not set")
if IsPerfDebugEnabled() {
t.Error("IsPerfDebugEnabled should return false when env var is not set")
}
os.Setenv(PerfTimingEnv, "1")
defer os.Unsetenv(PerfTimingEnv)
os.Setenv(PerfDebugEnv, "1")
defer os.Unsetenv(PerfDebugEnv)
if !IsPerfTimingEnabled() {
t.Error("IsPerfTimingEnabled should return true when env var is set")
if !IsPerfDebugEnabled() {
t.Error("IsPerfDebugEnabled should return true when env var is set")
}
}
// ── PerfReport tests ────────────────────────────────────────────────────
func TestBuildReport(t *testing.T) {
tc := NewTimingCollector()
tc.Record("cmd_init", 45*time.Millisecond)
tc.Record("auth_keychain", 72*time.Millisecond)
tc.Record("mcp_call", 620*time.Millisecond)
report := tc.BuildReport("v1.0.8", "dws aitable list-records")
if report.Kind != "perf_report" {
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
}
if report.Version != "1" {
t.Errorf("expected version '1', got %q", report.Version)
}
if report.CLIVersion != "v1.0.8" {
t.Errorf("expected cli_version 'v1.0.8', got %q", report.CLIVersion)
}
if report.Command != "dws aitable list-records" {
t.Errorf("expected command 'dws aitable list-records', got %q", report.Command)
}
if len(report.Phases) != 3 {
t.Fatalf("expected 3 phases, got %d", len(report.Phases))
}
if report.Phases[0].Name != "cmd_init" || report.Phases[0].DurationMs != 45 {
t.Errorf("unexpected first phase: %+v", report.Phases[0])
}
if report.Slowest != "mcp_call" {
t.Errorf("expected slowest 'mcp_call', got %q", report.Slowest)
}
if report.TotalMs < 0 {
t.Errorf("total_ms should be >= 0, got %d", report.TotalMs)
}
if report.OverheadMs < 0 {
t.Errorf("overhead_ms should be >= 0, got %d", report.OverheadMs)
}
}
func TestBuildReportEmpty(t *testing.T) {
tc := NewTimingCollector()
report := tc.BuildReport("dev", "dws version")
if len(report.Phases) != 0 {
t.Errorf("expected 0 phases, got %d", len(report.Phases))
}
if report.Slowest != "" {
t.Errorf("expected empty slowest, got %q", report.Slowest)
}
}
func TestBuildReportJSON(t *testing.T) {
tc := NewTimingCollector()
tc.Record("cmd_init", 10*time.Millisecond)
report := tc.BuildReport("v1.0.0", "dws version")
data, err := json.Marshal(report)
if err != nil {
t.Fatalf("json.Marshal failed: %v", err)
}
var parsed map[string]any
if err := json.Unmarshal(data, &parsed); err != nil {
t.Fatalf("json.Unmarshal failed: %v", err)
}
requiredKeys := []string{"kind", "version", "cli_version", "command", "timestamp", "total_ms", "phases", "slowest", "overhead_ms"}
for _, key := range requiredKeys {
if _, ok := parsed[key]; !ok {
t.Errorf("missing key %q in JSON output", key)
}
}
}
func TestWriteReportIfEnabled(t *testing.T) {
dir := t.TempDir()
reportPath := filepath.Join(dir, "report.json")
t.Setenv(PerfReportEnv, reportPath)
tc := NewTimingCollector()
tc.Record("cmd_init", 50*time.Millisecond)
tc.Record("mcp_call", 200*time.Millisecond)
tc.WriteReportIfEnabled("v1.0.0", "dws version")
data, err := os.ReadFile(reportPath)
if err != nil {
t.Fatalf("report file not written: %v", err)
}
var report PerfReport
if err := json.Unmarshal(data, &report); err != nil {
t.Fatalf("invalid JSON in report: %v", err)
}
if report.Kind != "perf_report" {
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
}
if len(report.Phases) != 2 {
t.Errorf("expected 2 phases, got %d", len(report.Phases))
}
}
func TestWriteReportIfEnabled_Auto(t *testing.T) {
tmpHome := t.TempDir()
expected := filepath.Join(tmpHome, ".dws", "perf", "latest.json")
// Temporarily override HOME for defaultPerfReportPath
t.Setenv("HOME", tmpHome)
t.Setenv(PerfReportEnv, "auto")
tc := NewTimingCollector()
tc.Record("cmd_init", 10*time.Millisecond)
tc.WriteReportIfEnabled("v1.0.0", "dws version")
if _, err := os.Stat(expected); err != nil {
t.Fatalf("expected report at %s: %v", expected, err)
}
}
func TestWriteReportIfEnabled_Disabled(t *testing.T) {
t.Setenv(PerfReportEnv, "")
tc := NewTimingCollector()
tc.Record("op", 10*time.Millisecond)
tc.WriteReportIfEnabled("v1.0.0", "dws version")
// No file should be written; no error expected
}
func TestWriteReportIfEnabled_NilCollector(t *testing.T) {
t.Setenv(PerfReportEnv, "/tmp/should-not-exist.json")
var tc *TimingCollector
tc.WriteReportIfEnabled("v1.0.0", "dws version")
}
func TestLoadLatestReport(t *testing.T) {
tmpHome := t.TempDir()
t.Setenv("HOME", tmpHome)
perfDir := filepath.Join(tmpHome, ".dws", "perf")
if err := os.MkdirAll(perfDir, 0o700); err != nil {
t.Fatal(err)
}
report := PerfReport{
Kind: "perf_report",
Version: "1",
CLIVersion: "v1.0.0",
Command: "dws version",
TotalMs: 100,
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}},
Slowest: "cmd_init",
OverheadMs: 50,
}
data, _ := json.MarshalIndent(report, "", " ")
if err := os.WriteFile(filepath.Join(perfDir, "latest.json"), data, 0o600); err != nil {
t.Fatal(err)
}
loaded, err := LoadLatestReport()
if err != nil {
t.Fatalf("LoadLatestReport failed: %v", err)
}
if loaded.CLIVersion != "v1.0.0" {
t.Errorf("expected cli_version 'v1.0.0', got %q", loaded.CLIVersion)
}
if len(loaded.Phases) != 1 {
t.Errorf("expected 1 phase, got %d", len(loaded.Phases))
}
}
func TestLoadLatestReport_NotFound(t *testing.T) {
tmpHome := t.TempDir()
t.Setenv("HOME", tmpHome)
_, err := LoadLatestReport()
if err == nil {
t.Error("expected error when report file does not exist")
}
}
func TestSanitizeCommand(t *testing.T) {
tests := []struct {
name string
args []string
want string
}{
{
name: "no sensitive flags",
args: []string{"dws", "aitable", "list-records"},
want: "dws aitable list-records",
},
{
name: "token with space-separated value",
args: []string{"dws", "--token", "secret123", "version"},
want: "dws --token *** version",
},
{
name: "token with equals sign",
args: []string{"dws", "--token=secret123", "version"},
want: "dws --token=*** version",
},
{
name: "client-secret space-separated",
args: []string{"dws", "--client-secret", "mysecret", "--client-id", "myid", "auth"},
want: "dws --client-secret *** --client-id *** auth",
},
{
name: "client-id with equals",
args: []string{"dws", "--client-id=abc123"},
want: "dws --client-id=***",
},
{
name: "empty args",
args: []string{},
want: "",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := SanitizeCommand(tt.args)
if got != tt.want {
t.Errorf("SanitizeCommand(%v) = %q, want %q", tt.args, got, tt.want)
}
})
}
}
func TestResolvePerfReportPath_Auto(t *testing.T) {
p := resolvePerfReportPath("auto")
if p == "" {
t.Skip("HOME not available")
}
if !strings.HasSuffix(p, filepath.Join("perf", "latest.json")) {
t.Errorf("expected path ending in perf/latest.json, got %q", p)
}
}
func TestResolvePerfReportPath_Custom(t *testing.T) {
p := resolvePerfReportPath("/tmp/my-report.json")
if p != "/tmp/my-report.json" {
t.Errorf("expected '/tmp/my-report.json', got %q", p)
}
}
func TestPrintPerfReportSummary(t *testing.T) {
report := &PerfReport{
Command: "dws version",
Timestamp: time.Now(),
TotalMs: 300,
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}, {Name: "mcp_call", DurationMs: 200, Seq: 1}},
Slowest: "mcp_call",
OverheadMs: 50,
}
var buf bytes.Buffer
printPerfReportSummary(&buf, report)
out := buf.String()
if !strings.Contains(out, "cmd_init") {
t.Error("output should contain 'cmd_init'")
}
if !strings.Contains(out, "mcp_call") {
t.Error("output should contain 'mcp_call'")
}
if !strings.Contains(out, "← 最慢") {
t.Error("output should contain '← 最慢' marker")
}
if !strings.Contains(out, "总耗时") {
t.Error("output should contain '总耗时'")
}
if !strings.Contains(out, "框架开销") {
t.Error("output should contain '框架开销'")
}
}
+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")
}
}
+24
View File
@@ -15,6 +15,21 @@ package app
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).
func Version() string {
@@ -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 }
+847
View File
@@ -0,0 +1,847 @@
// 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: &CLIAuthResult{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 == nil || !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: &CLIAuthResult{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 == nil || !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: &CLIAuthResult{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 == nil || 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: &CLIAuthResult{CLIAuthEnabled: true},
}
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && 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: &CLIAuthResult{CLIAuthEnabled: false},
}
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && 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,
FlowID: "test-flow-id",
}, "", "")
case strings.HasSuffix(r.URL.Path, DevicePollPath):
// New terminal API: return APPROVED status
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{
"status": "APPROVED",
"authCode": "test-auth-code",
},
})
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"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,
terminalBaseURL: 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,
FlowID: "test-flow-id",
}, "", "")
case strings.HasSuffix(r.URL.Path, DevicePollPath):
// New terminal API: return APPROVED status
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{
"status": "APPROVED",
"authCode": "test-auth-code",
},
})
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"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: &CLIAuthResult{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,
terminalBaseURL: 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,
FlowID: "test-flow-id",
}, "", "")
case strings.HasSuffix(r.URL.Path, DevicePollPath):
// New terminal API: return APPROVED status
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{
"status": "APPROVED",
"authCode": "test-auth-code",
},
})
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
json.NewEncoder(w).Encode(map[string]any{
"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: &CLIAuthResult{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,
terminalBaseURL: 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)
}
+182 -59
View File
@@ -26,39 +26,44 @@ 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"
)
const (
// defaultPollInterval is the default seconds between device token polls.
defaultPollInterval = 5
// The server-side Redis TTL is 10 minutes; a 2-second interval keeps the
// user-perceived latency low while staying well within rate limits.
defaultPollInterval = 2
// maxPollInterval caps the polling interval to prevent DoS via slow_down.
maxPollInterval = 30
// maxPollTotalWait caps the total wait time for device authorization.
maxPollTotalWait = 15 * time.Minute
// Aligned with the server-side Redis TTL (10 minutes).
maxPollTotalWait = 10 * time.Minute
)
type DeviceFlowProvider struct {
configDir string
clientID string
scope string
baseURL string
logger *slog.Logger
Output io.Writer
httpClient *http.Client
configDir string
clientID string
scope string
baseURL string
terminalBaseURL string
logger *slog.Logger
Output io.Writer
httpClient *http.Client
}
func NewDeviceFlowProvider(configDir string, logger *slog.Logger) *DeviceFlowProvider {
return &DeviceFlowProvider{
configDir: configDir,
clientID: ClientID(),
scope: DefaultScopes,
baseURL: DefaultDeviceBaseURL,
logger: logger,
Output: os.Stderr,
httpClient: &http.Client{Timeout: 30 * time.Second},
configDir: configDir,
clientID: ClientID(),
scope: DefaultScopes,
baseURL: DefaultDeviceBaseURL,
terminalBaseURL: GetMCPBaseURL(),
logger: logger,
Output: os.Stderr,
httpClient: &http.Client{Timeout: 30 * time.Second},
}
}
@@ -66,6 +71,18 @@ func (p *DeviceFlowProvider) SetBaseURL(baseURL string) {
p.baseURL = strings.TrimRight(baseURL, "/")
}
// SetTerminalBaseURL sets the terminal API base URL for device flow polling.
func (p *DeviceFlowProvider) SetTerminalBaseURL(baseURL string) {
p.terminalBaseURL = strings.TrimRight(baseURL, "/")
}
// SetScope overrides the OAuth scope for the device flow.
func (p *DeviceFlowProvider) SetScope(scope string) {
if p != nil {
p.scope = scope
}
}
func (p *DeviceFlowProvider) output() io.Writer {
if p != nil && p.Output != nil {
return p.Output
@@ -80,6 +97,7 @@ type DeviceAuthResponse struct {
VerificationURIComplete string `json:"verificationUriComplete"`
ExpiresIn int `json:"expiresIn"`
Interval int `json:"interval"`
FlowID string `json:"flowId"`
}
type DeviceTokenResponse struct {
@@ -88,6 +106,20 @@ type DeviceTokenResponse struct {
Error string `json:"error"`
}
// DevicePollResponse represents the response from the terminal API poll endpoint.
type DevicePollResponse struct {
Success bool `json:"success"`
Code string `json:"code,omitempty"`
Message string `json:"message,omitempty"`
Data DevicePollData `json:"data"`
}
type DevicePollData struct {
Status string `json:"status"`
AuthCode string `json:"authCode,omitempty"`
FlowID string `json:"flowId,omitempty"`
}
type serviceResult struct {
Success bool `json:"success"`
Result json.RawMessage `json:"result"`
@@ -153,6 +185,11 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
if err != nil {
return nil, err
}
if tokenResult == nil {
// FlowID was empty — no polling happened; authorization URL was already
// printed, so the user can handle it manually.
return nil, nil
}
_, _ = fmt.Fprintln(p.output(), "")
dfPrintStep(p.output(), 3, i18n.T("使用授权码换取 Access Token..."), 0)
@@ -167,41 +204,70 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
}
// Check if CLI auth is enabled for this organization
// 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 {
if p.logger != nil {
p.logger.Warn("failed to check CLI auth status", "error", authErr)
}
// Continue anyway - fail open for better UX
} 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(), dfRed(i18n.T("⚠️ 无法检查 CLI 数据访问权限状态")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请检查网络连接后重试。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, fmt.Errorf("%s: %w", i18n.T("检查 CLI 授权状态失败"), authErr)
}
denialReason := classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
if denialReason != "" {
_, _ = fmt.Fprintln(p.output(), "")
switch denialReason {
case "user_forbidden":
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织已禁止所有成员使用 CLI")))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("该组织已禁止所有成员使用 CLI"))
case "user_not_allowed":
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 您不在该组织的 CLI 授权人员范围内")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员将您加入 CLI 授权人员名单。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员"))
case "channel_not_allowed":
ch := os.Getenv("DWS_CHANNEL")
_, _ = fmt.Fprintf(p.output(), dfRed(i18n.T("⚠️ 当前渠道 %s 未获得该组织授权"))+"\n", ch)
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员开通该渠道的访问权限。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, fmt.Errorf(i18n.T("当前渠道 %s 未获得该组织授权,请联系组织管理员"), ch)
case "channel_required":
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 当前组织已开启渠道管控")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请升级到最新版本的 CLI 后重试。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("当前组织已开启渠道管控,请升级到最新版本的 CLI 后重试"))
case "no_auth":
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 认证已失效")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请执行 dws auth 重新登录。"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("认证已失效,请执行 dws auth 重新登录"))
default:
// cli_not_enabled or unknown — show existing admin-apply flow
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
_, _ = fmt.Fprintln(p.output(), "")
// Try to get super admin list
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
// Show up to 3 admins
maxAdmins := 3
if len(admins.Result) < maxAdmins {
maxAdmins = len(admins.Result)
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
maxAdmins := 3
if len(admins.Result) < maxAdmins {
maxAdmins = len(admins.Result)
}
var adminNames []string
for i := 0; i < maxAdmins; i++ {
adminNames = append(adminNames, admins.Result[i].Name)
}
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
}
var adminNames []string
for i := 0; i < maxAdmins; i++ {
adminNames = append(adminNames, admins.Result[i].Name)
}
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
}
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), i18n.T(" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("管理员操作入口:"), config.GetDeveloperSettingsURL())
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
}
}
// Save token data with associated client ID for refresh
@@ -210,6 +276,13 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
}
// Always persist clientId to app.json so future process startups
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
if p.clientID != "" {
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
}
// Persist app credentials if using custom client credentials
oauthProvider.persistAppConfigIfNeeded()
@@ -278,7 +351,41 @@ func (p *DeviceFlowProvider) pollDeviceToken(ctx context.Context, deviceCode str
return &resp, nil
}
// pollDeviceStatus polls the terminal API for device authorization status.
//
// Note: The server returns success=false for REJECTED and EXPIRED terminal
// states (with a valid data.Status value). These are normal business outcomes,
// not transport errors, so we return the response to the caller and let the
// status-switch handle them.
func (p *DeviceFlowProvider) pollDeviceStatus(ctx context.Context, flowID string) (*DevicePollResponse, error) {
endpoint := fmt.Sprintf("%s%s?flowId=%s", p.terminalBaseURL, DevicePollPath, url.QueryEscape(flowID))
body, err := p.doGet(ctx, endpoint)
if err != nil {
return nil, err
}
var resp DevicePollResponse
if err := json.Unmarshal(body, &resp); err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("解析响应失败"), err)
}
// REJECTED/EXPIRED carry success=false but have a valid data.Status;
// only treat as a real server error when data.Status is empty.
if !resp.Success && resp.Data.Status == "" {
return nil, fmt.Errorf("%s: [%s] %s", i18n.T("服务端返回错误"), resp.Code, resp.Message)
}
return &resp, nil
}
func (p *DeviceFlowProvider) waitForAuthorization(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
// No FlowID from server — cannot poll status; return immediately so the
// user can still see the authorization URL printed earlier and handle it
// manually (same pattern as pat_auth_retry.go L451).
if auth.FlowID == "" {
dfPrintDim(p.output(), i18n.T(" 服务端未返回 flowId,跳过轮询,请在浏览器中手动完成授权后重试"))
_, _ = fmt.Fprintln(p.output(), "")
return nil, nil
}
startTime := time.Now()
interval := time.Duration(auth.Interval) * time.Second
deadline := time.Duration(auth.ExpiresIn) * time.Second
@@ -301,7 +408,7 @@ func (p *DeviceFlowProvider) waitForAuthorization(ctx context.Context, auth *Dev
elapsedSec := int(time.Since(startTime).Seconds())
dfPrintPollStatus(p.output(), pollCount, elapsedSec)
resp, err := p.pollDeviceToken(ctx, auth.DeviceCode)
pollResp, err := p.pollDeviceStatus(ctx, auth.FlowID)
if err != nil {
dfPrintPollResult(p.output(), "network_error", i18n.T("网络错误,继续重试..."))
if p.logger != nil {
@@ -310,27 +417,20 @@ func (p *DeviceFlowProvider) waitForAuthorization(ctx context.Context, auth *Dev
continue
}
if resp.Error == "" {
switch pollResp.Data.Status {
case StatusApproved:
dfPrintPollResult(p.output(), "authorized", i18n.T("授权成功!"))
return resp, nil
}
switch resp.Error {
case "authorization_pending":
return &DeviceTokenResponse{AuthCode: pollResp.Data.AuthCode}, nil
case StatusPending:
dfPrintPollResult(p.output(), "pending", i18n.T("等待用户授权..."))
case "slow_down":
interval += 5 * time.Second
if interval > maxPollInterval*time.Second {
interval = maxPollInterval * time.Second
}
dfPrintPollResult(p.output(), "slow_down", fmt.Sprintf(i18n.T("轮询过快,间隔增加至 %ds"), int(interval.Seconds())))
case "access_denied":
case StatusRejected:
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("用户拒绝了授权请求"))
case "expired_token":
case StatusExpired:
_, _ = fmt.Fprintln(p.output(), "")
return nil, errors.New(i18n.T("设备授权码已过期"))
default:
dfPrintPollResult(p.output(), "unknown", fmt.Sprintf(i18n.T("未知错误: %s"), resp.Error))
dfPrintPollResult(p.output(), "unknown", fmt.Sprintf(i18n.T("未知状态: %s"), pollResp.Data.Status))
}
}
}
@@ -358,6 +458,29 @@ func (p *DeviceFlowProvider) postForm(ctx context.Context, endpoint string, para
return body, nil
}
// doGet performs an HTTP GET request and returns the response body.
func (p *DeviceFlowProvider) doGet(ctx context.Context, endpoint string) ([]byte, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("创建请求失败"), err)
}
resp, err := p.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("发送请求失败"), err)
}
defer resp.Body.Close()
body, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
if err != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("读取响应失败"), err)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncateBody(body, 200))
}
return body, nil
}
// truncateBody returns a string of at most maxLen bytes from body, appending
// "...(truncated)" when the content exceeds the limit. This prevents leaking
// potentially sensitive response payloads in error messages.
+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 auth
// Device flow authorization status constants.
// Shared across device_flow.go and pat_auth_retry.go to avoid maintaining
// string literals in multiple places.
const (
StatusPending = "PENDING"
StatusApproved = "APPROVED"
StatusRejected = "REJECTED"
StatusExpired = "EXPIRED"
StatusCancelled = "CANCELLED"
)
// ParseDeviceFlowStatus normalizes a raw status string from the device flow
// poll response into a canonical status constant. When the server returns an
// empty status with success=false, it falls back to StatusExpired (server
// error / flow not found).
func ParseDeviceFlowStatus(rawStatus string, success bool) string {
switch rawStatus {
case StatusApproved, StatusRejected, StatusExpired, StatusPending, StatusCancelled:
return rawStatus
default:
if rawStatus == "" && !success {
return StatusExpired
}
return rawStatus
}
}
+37 -11
View File
@@ -20,6 +20,7 @@ import (
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
@@ -89,22 +90,42 @@ func TestWaitForAuthorizationSucceedsAfterPending(t *testing.T) {
var calls atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// New terminal API uses GET method
if r.Method != http.MethodGet {
t.Fatalf("method = %s, want GET", r.Method)
}
if !strings.Contains(r.URL.RawQuery, "flowId=") {
t.Fatal("flowId query parameter should be present")
}
if calls.Add(1) <= 2 {
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
// Return PENDING status
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{"status": "PENDING"},
})
return
}
writeServiceResult(w, true, DeviceTokenResponse{AuthCode: "final-auth-code"}, "", "")
// Return APPROVED status with authCode
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{
"status": "APPROVED",
"authCode": "final-auth-code",
},
})
}))
defer server.Close()
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
provider.Output = io.Discard
provider.SetBaseURL(server.URL)
provider.SetTerminalBaseURL(server.URL)
resp, err := provider.waitForAuthorization(context.Background(), &DeviceAuthResponse{
DeviceCode: "dc-1",
ExpiresIn: 10,
Interval: 1,
FlowID: "test-flow-id",
ExpiresIn: 10,
Interval: 1,
})
if err != nil {
t.Fatalf("waitForAuthorization() error = %v", err)
@@ -121,21 +142,26 @@ func TestWaitForAuthorizationHonorsContextCancellation(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
// New terminal API uses GET method
w.Header().Set("Content-Type", "application/json")
json.NewEncoder(w).Encode(map[string]interface{}{
"success": true,
"data": map[string]string{"status": "PENDING"},
})
}))
defer server.Close()
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
provider.Output = io.Discard
provider.SetBaseURL(server.URL)
provider.SetTerminalBaseURL(server.URL)
ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
defer cancel()
if _, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
DeviceCode: "dc-2",
ExpiresIn: 60,
Interval: 1,
FlowID: "test-flow-id-2",
ExpiresIn: 60,
Interval: 1,
}); err == nil {
t.Fatal("waitForAuthorization() error = nil, want context cancellation")
}
+51 -1
View File
@@ -18,8 +18,34 @@ import (
"path/filepath"
"strings"
"sync"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
func init() {
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_CLIENT_ID",
Category: configmeta.CategoryAuth,
Description: "OAuth AppKey (DingTalk 应用凭证)",
DefaultValue: "(内置)",
Sensitive: true,
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_CLIENT_SECRET",
Category: configmeta.CategoryAuth,
Description: "OAuth AppSecret (DingTalk 应用凭证)",
DefaultValue: "(内置)",
Sensitive: true,
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_CHANNEL",
Category: configmeta.CategoryExternal,
Description: "第三方渠道编码 (channelCode),如 Qoderwork",
})
}
const (
// AuthorizeURL is the DingTalk OAuth authorization page.
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
@@ -58,6 +84,14 @@ const (
// DeviceGrantType is the grant_type value defined by RFC 8628.
DeviceGrantType = "urn:ietf:params:oauth:grant-type:device_code"
// Terminal API base URL for developer settings page.
DefaultTerminalBaseURL = "https://open-dev.dingtalk.com"
// DevicePollPath is the device flow polling path (used with MCP base URL).
DevicePollPath = "/cli/oauth/device/poll"
// DeveloperSettingsPath is the path to the organization developer settings page.
DeveloperSettingsPath = "/fe/old#/developerSettings"
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
LogoutContinueURL = "https://login.dingtalk.com"
@@ -74,6 +108,19 @@ const (
MCPRevokeTokenPath = "/oauth2/revokeToken"
)
// GetTerminalBaseURL returns the terminal base URL with priority:
// 1. ~/.dws/terminal_url file content (for pre-release environment)
// 2. Default value (https://open-dev.dingtalk.com)
func GetTerminalBaseURL() string {
return config.GetTerminalBaseURL()
}
// GetDeveloperSettingsURL returns the full URL to the organization developer
// settings page, derived from the terminal base URL.
func GetDeveloperSettingsURL() string {
return config.GetDeveloperSettingsURL()
}
// GetMCPBaseURL returns the MCP base URL with priority:
// 1. ~/.dws/mcp_url file content (for pre-release environment)
// 2. Default value (https://mcp.dingtalk.com)
@@ -110,7 +157,7 @@ func SetClientIDFromMCP(id string) {
func IsClientIDFromMCP() bool {
clientMu.RLock()
defer clientMu.RUnlock()
return clientIDFromMCP
return clientIDFromMCP || edition.Get().AuthClientFromMCP
}
// GetUserAccessTokenURL returns the appropriate token exchange URL.
@@ -189,6 +236,9 @@ func ClientID() string {
if override != "" {
return override
}
if id := edition.Get().AuthClientID; id != "" {
return id
}
// Try loading from persisted app config
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
return id
+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
}
+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
+39 -10
View File
@@ -124,6 +124,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
token *TokenData
err error
cliAuthDisabled bool
denialReason string
}
resultCh := make(chan callbackResult, 1)
errCh := make(chan error, 1)
@@ -234,21 +235,32 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
}
callbackTokenMu.Unlock()
// Check CLI auth enabled status
// Check CLI auth enabled status (fail-closed: treat errors as disabled)
authStatus, statusErr := p.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
cliAuthDisabled := statusErr == nil && authStatus.Success && !authStatus.Result.CLIAuthEnabled
var denialReason string
if statusErr != nil {
denialReason = "unknown"
} else {
denialReason = classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
}
cliAuthEnabled := denialReason == ""
// Update CLI auth disabled state
callbackTokenMu.Lock()
callbackAuthDisabled = cliAuthDisabled
callbackAuthDisabled = !cliAuthEnabled
callbackTokenMu.Unlock()
// Display appropriate HTML based on CLI auth status
// Display appropriate HTML based on auth status and denial reason
w.Header().Set("Content-Type", "text/html; charset=utf-8")
if cliAuthDisabled {
_, _ = fmt.Fprint(w, notEnabledHTML)
} else {
switch {
case cliAuthEnabled:
_, _ = fmt.Fprint(w, successHTML)
case denialReason == "user_forbidden" || denialReason == "user_not_allowed":
_, _ = fmt.Fprint(w, accessDeniedHTML)
case denialReason == "channel_not_allowed" || denialReason == "channel_required":
_, _ = fmt.Fprint(w, channelDeniedHTML)
default:
_, _ = fmt.Fprint(w, notEnabledHTML)
}
// Ensure response is flushed to client
if f, ok := w.(http.Flusher); ok {
@@ -256,7 +268,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
}
// Notify main goroutine with full result
select {
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: cliAuthDisabled}:
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled, denialReason: denialReason}:
default:
}
})
@@ -395,8 +407,18 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), result.err)
}
// Handle CLI auth disabled - keep server running for user to apply
// Handle CLI auth disabled - for terminal denial reasons, exit immediately
// (page shows accessDeniedHTML/channelDeniedHTML with no apply button,
// so polling for apply submission would hang forever).
// Error messages are kept consistent with the text shown on the HTML pages.
if result.cliAuthDisabled {
switch result.denialReason {
case "user_forbidden", "user_not_allowed":
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员将您加入授权名单"))
case "channel_not_allowed", "channel_required":
return nil, errors.New(i18n.T("当前渠道未获得该组织授权,或组织已开启渠道管控,请联系组织管理员开通渠道访问权限,或升级到最新版本的 CLI"))
}
_, _ = fmt.Fprintln(p.output(), "")
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
@@ -435,7 +457,7 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
// Check if CLI auth is now enabled (admin approved)
if currentToken != nil {
authStatus, err := p.CheckCLIAuthEnabled(ctx, currentToken.AccessToken)
if err == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled {
if err == nil && classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL")) == "" {
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
time.Sleep(2 * time.Second)
result.token = currentToken
@@ -463,6 +485,13 @@ continueLogin:
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
}
// Always persist clientId to app.json so future process startups
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
if p.clientID != "" {
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
}
// Persist app credentials if using custom client credentials
p.persistAppConfigIfNeeded()
+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"
+65 -17
View File
@@ -20,7 +20,11 @@ import (
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
// TokenData holds the OAuth token set persisted to disk.
@@ -61,44 +65,88 @@ func (t *TokenData) HasPersistentCode() bool {
return t != nil && t.PersistentCode != ""
}
// SaveTokenData saves TokenData to the platform keychain.
// Uses the new keychain-based storage with random master key for better security.
const tokenJSONFile = "token.json"
// TokenMarker is a lightweight file the host application reads to detect
// whether the CLI has a valid token without accessing the keychain.
type TokenMarker struct {
UpdatedAt string `json:"updated_at"`
}
// WriteTokenMarker writes a token.json marker containing only an updated_at
// timestamp. The host application uses this file's presence and mtime to
// decide whether it needs to trigger a new auth exchange.
func WriteTokenMarker(configDir string) error {
marker := TokenMarker{UpdatedAt: time.Now().Format(time.RFC3339)}
data, _ := json.MarshalIndent(marker, "", " ")
if err := os.MkdirAll(configDir, 0o700); err != nil {
return err
}
tmp := filepath.Join(configDir, tokenJSONFile+".tmp")
if err := os.WriteFile(tmp, data, 0o600); err != nil {
return err
}
return os.Rename(tmp, filepath.Join(configDir, tokenJSONFile))
}
// DeleteTokenMarker removes the token.json marker file.
func DeleteTokenMarker(configDir string) error {
return os.Remove(filepath.Join(configDir, tokenJSONFile))
}
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to the default keychain-based storage.
func SaveTokenData(configDir string, data *TokenData) error {
if h := edition.Get(); h.SaveToken != nil {
jsonData, err := json.MarshalIndent(data, "", " ")
if err != nil {
return fmt.Errorf("marshaling token data for hook: %w", err)
}
return h.SaveToken(configDir, jsonData)
}
return SaveTokenDataKeychain(data)
}
// LoadTokenData reads TokenData from the platform keychain.
// On first call, it attempts to migrate legacy .data file if present.
// LoadTokenData reads TokenData. When an edition hook (LoadToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to keychain with legacy .data migration.
func LoadTokenData(configDir string) (*TokenData, error) {
// Try loading from new keychain first
if h := edition.Get(); h.LoadToken != nil {
jsonData, err := h.LoadToken(configDir)
if err != nil {
return nil, err
}
var td TokenData
if err := json.Unmarshal(jsonData, &td); err != nil {
return nil, fmt.Errorf("parsing token data from hook: %w", err)
}
return &td, nil
}
// Default: keychain with legacy .data migration
if TokenDataExistsKeychain() {
return LoadTokenDataKeychain()
}
// Fallback: try legacy .data file and migrate
data, err := LoadSecureTokenData(configDir)
if err != nil {
return nil, err
}
// Migrate to keychain for future use
if err := SaveTokenDataKeychain(data); err == nil {
// Successfully migrated, delete legacy file
_ = DeleteSecureData(configDir)
}
return data, nil
}
// DeleteTokenData removes token data from both keychain and legacy storage.
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to keychain + legacy cleanup.
func DeleteTokenData(configDir string) error {
// Delete from keychain
if h := edition.Get(); h.DeleteToken != nil {
return h.DeleteToken(configDir)
}
keychainErr := DeleteTokenDataKeychain()
// Also clean up any legacy .data file
legacyErr := DeleteSecureData(configDir)
// Return keychain error if any, otherwise legacy error
if keychainErr != nil {
return keychainErr
}
+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
}
+102 -9
View File
@@ -17,25 +17,107 @@ 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,
})
configmeta.Register(configmeta.ConfigItem{
Name: "DWS_PLUGIN_COLD_TIMEOUT",
Category: configmeta.CategoryCore,
Description: "插件 MCP 冷启动发现的超时时长(Go duration 格式,如 2s / 1500ms)。设置后同时覆盖 HTTP 与 stdio 插件的冷启动预算;未设置时使用内置默认值(HTTP 无鉴权 1s / 有鉴权 1.5s / stdio 2s)。",
DefaultValue: "",
Example: "3s",
})
}
// CatalogDegradedReason identifies why catalog discovery returned empty.
type CatalogDegradedReason string
const (
DegradedUnauthenticated CatalogDegradedReason = "unauthenticated"
DegradedMarketUnreachable CatalogDegradedReason = "market_unreachable"
DegradedRuntimeAllFailed CatalogDegradedReason = "runtime_all_failed"
)
// CatalogDegraded is returned by EnvironmentLoader.Load when discovery
// fails for a diagnosable reason. Callers that need graceful degradation
// (e.g. the runtime runner) can check errors.As and fall back to an
// empty catalog; callers like the schema command can surface the hint.
type CatalogDegraded struct {
Reason CatalogDegradedReason
Hint string
ServerCount int // number of servers discovered (only set for runtime_all_failed)
}
func (e *CatalogDegraded) Error() string { return string(e.Reason) + ": " + e.Hint }
func degradedHint(reason CatalogDegradedReason, serverCount int) string {
embedded := edition.Get().IsEmbedded
switch reason {
case DegradedUnauthenticated:
if embedded {
return "未登录,请重新认证"
}
return "未登录,无法发现 MCP 服务。请先执行: dws auth login"
case DegradedMarketUnreachable:
if embedded {
return "无法连接 MCP 市场,请检查网络"
}
return "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络"
case DegradedRuntimeAllFailed:
if embedded {
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试", serverCount)
}
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试或执行: dws cache refresh", serverCount)
default:
return "MCP 服务发现失败"
}
}
func newCatalogDegraded(reason CatalogDegradedReason, serverCount int) *CatalogDegraded {
return &CatalogDegraded{
Reason: reason,
Hint: degradedHint(reason, serverCount),
ServerCount: serverCount,
}
}
const (
CatalogFixtureEnv = "DWS_CATALOG_FIXTURE"
CacheDirEnv = "DWS_CACHE_DIR"
PluginColdTimeoutEnv = "DWS_PLUGIN_COLD_TIMEOUT"
DefaultMarketBaseURL = "https://mcp.dingtalk.com"
// defaultDiscoveryTimeout bounds the time spent on live registry discovery.
defaultDiscoveryTimeout = 10 * time.Second
// Tightened to 4s so a slow/unreachable discovery endpoint cannot block
// every CLI command invocation. See issue #119.
defaultDiscoveryTimeout = 4 * time.Second
)
type CatalogLoader interface {
@@ -92,6 +174,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 +209,23 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
// Startup command construction should not block on synchronous discovery
// just because the cache has aged past the short revalidation window.
cached := l.loadFromCache(store)
if cached.Available {
if cached.Available && len(cached.Catalog.Products) > 0 {
return cached.Catalog, nil
}
transportClient := transport.NewClient(nil)
hasAuth := false
if l.AuthTokenFunc != nil {
if token := l.AuthTokenFunc(ctx); token != "" {
transportClient = transportClient.WithAuth(token, nil)
hasAuth = true
}
}
if !hasAuth {
return ir.Catalog{}, newCatalogDegraded(DegradedUnauthenticated, 0)
}
// Use a bounded context so discovery doesn't hang in test or CI environments.
timeout := defaultDiscoveryTimeout
if l.DiscoveryTimeout > 0 {
@@ -147,14 +239,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 +277,10 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
refreshed, failures := service.DiscoverAllRuntime(discoverCtx, toRefresh)
if len(unchangedRuntime) == 0 && len(refreshed) == 0 && len(failures) > 0 {
if cached.Available {
if cached.Available && len(cached.Catalog.Products) > 0 {
return cached.Catalog, nil
}
return ir.Catalog{}, nil
return ir.Catalog{}, newCatalogDegraded(DegradedRuntimeAllFailed, len(servers))
}
refreshedByKey := make(map[string]discovery.RuntimeServer, len(refreshed))
+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
+71 -1
View File
@@ -26,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
@@ -139,6 +140,11 @@ func NewDirectCommand(route Route, runner executor.Runner) *cobra.Command {
for key, value := range bindingParams {
params[key] = value
}
// Collect schema-derived flags (from buildFlagsFromDetailSchema)
// that are not covered by explicit bindings.
collectSchemaFlags(cmd, route.Bindings, params)
if route.Normalizer != nil {
if err := route.Normalizer(cmd, params); err != nil {
return err
@@ -246,6 +252,70 @@ func ApplyBindings(cmd *cobra.Command, bindings []FlagBinding) {
_ = cmd.Flags().MarkHidden("params")
}
// collectSchemaFlags picks up flags created by buildFlagsFromDetailSchema that
// have no explicit FlagBinding. This bridges the gap for plugin-defined tools
// whose parameters come from the MCP inputSchema rather than CLIToolOverride.Flags.
func collectSchemaFlags(cmd *cobra.Command, bindings []FlagBinding, params map[string]any) {
// Build a set of flag names already covered by bindings.
bound := make(map[string]bool, len(bindings)*2)
for _, b := range bindings {
if n := strings.TrimSpace(b.FlagName); n != "" {
bound[n] = true
}
if a := strings.TrimSpace(b.Alias); a != "" {
bound[a] = true
}
}
// Reserved/internal flags that should never be forwarded as tool params.
skip := map[string]bool{
"json": true, "params": true, "help": true,
"format": true, "fields": true, "jq": true,
"debug": true, "verbose": true, "dry-run": true,
"yes": true, "mock": true, "timeout": true,
"client-id": true, "client-secret": true,
}
cmd.Flags().Visit(func(f *pflag.Flag) {
if bound[f.Name] || skip[f.Name] {
return
}
// Convert flag name back to the original parameter name (kebab → snake/camel)
// For simplicity, use the flag name as-is since MCP tools typically
// use snake_case which maps to kebab-case flags.
paramName := toOriginalParamName(f.Name)
if _, exists := params[paramName]; exists {
return // already set by --json/--params
}
switch f.Value.Type() {
case "int":
if v, err := cmd.Flags().GetInt(f.Name); err == nil {
params[paramName] = v
}
case "bool":
if v, err := cmd.Flags().GetBool(f.Name); err == nil {
params[paramName] = v
}
case "stringSlice":
if v, err := cmd.Flags().GetStringSlice(f.Name); err == nil {
params[paramName] = v
}
default:
if v, err := cmd.Flags().GetString(f.Name); err == nil {
params[paramName] = v
}
}
})
}
// toOriginalParamName converts a kebab-case flag name back to the original
// MCP parameter name. Since toKebabCase converts both camelCase and snake_case
// to kebab-case, we default to snake_case (the MCP convention).
func toOriginalParamName(flagName string) string {
return strings.ReplaceAll(flagName, "-", "_")
}
func CollectBindings(cmd *cobra.Command, bindings []FlagBinding, existing map[string]any) (map[string]any, error) {
if existing == nil {
existing = map[string]any{}
+99
View File
@@ -146,3 +146,102 @@ func TestCollectBindingsParsesJSONFlagValue(t *testing.T) {
t.Fatalf("config.options = %#v, want array of 1", config["options"])
}
}
func TestCollectSchemaFlagsPicksUpUnboundFlags(t *testing.T) {
t.Parallel()
// Simulate a plugin command with schema-generated flags but no bindings.
cmd := &cobra.Command{Use: "greet"}
cmd.Flags().String("name", "", "Name of person")
cmd.Flags().String("language", "en", "Language")
cmd.Flags().Int("count", 0, "Repeat count")
cmd.Flags().Bool("loud", false, "Loud mode")
cmd.Flags().String("json", "", "")
cmd.Flags().String("params", "", "")
// User sets --name and --count but not --language
_ = cmd.Flags().Set("name", "Alice")
_ = cmd.Flags().Set("count", "3")
_ = cmd.Flags().Set("loud", "true")
params := make(map[string]any)
collectSchemaFlags(cmd, nil, params)
if params["name"] != "Alice" {
t.Errorf("name = %v, want Alice", params["name"])
}
if params["count"] != 3 {
t.Errorf("count = %v, want 3", params["count"])
}
if params["loud"] != true {
t.Errorf("loud = %v, want true", params["loud"])
}
// language was not set by user, should not appear
if _, exists := params["language"]; exists {
t.Errorf("language should not be in params (not set by user)")
}
// json/params are reserved, should not appear
if _, exists := params["json"]; exists {
t.Error("json should be skipped")
}
}
func TestCollectSchemaFlagsSkipsBoundFlags(t *testing.T) {
t.Parallel()
bindings := []FlagBinding{
{FlagName: "dept-id", Property: "deptId", Kind: ValueString},
}
cmd := &cobra.Command{Use: "test"}
ApplyBindings(cmd, bindings)
// Also add a schema-generated flag
cmd.Flags().String("title", "", "Title")
_ = cmd.Flags().Set("dept-id", "D001")
_ = cmd.Flags().Set("title", "Hello")
params := make(map[string]any)
collectSchemaFlags(cmd, bindings, params)
// dept-id is bound, should NOT be collected by collectSchemaFlags
if _, exists := params["dept_id"]; exists {
t.Error("dept-id should be skipped (already has binding)")
}
// title is unbound, should be collected
if params["title"] != "Hello" {
t.Errorf("title = %v, want Hello", params["title"])
}
}
func TestCollectSchemaFlagsSkipsGlobalFlags(t *testing.T) {
t.Parallel()
cmd := &cobra.Command{Use: "test"}
cmd.Flags().String("name", "", "Name")
cmd.Flags().Bool("debug", false, "Debug")
cmd.Flags().Bool("verbose", false, "Verbose")
cmd.Flags().Bool("dry-run", false, "Dry run")
cmd.Flags().String("format", "json", "Format")
cmd.Flags().String("json", "", "")
cmd.Flags().String("params", "", "")
_ = cmd.Flags().Set("name", "Bob")
_ = cmd.Flags().Set("debug", "true")
_ = cmd.Flags().Set("verbose", "true")
_ = cmd.Flags().Set("dry-run", "true")
_ = cmd.Flags().Set("format", "table")
params := make(map[string]any)
collectSchemaFlags(cmd, nil, params)
if params["name"] != "Bob" {
t.Errorf("name = %v, want Bob", params["name"])
}
// Global flags should be skipped
for _, skip := range []string{"debug", "verbose", "dry_run", "format"} {
if _, exists := params[skip]; exists {
t.Errorf("%s should be skipped (global flag)", skip)
}
}
}
+130 -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"
@@ -41,6 +59,11 @@ type Service struct {
Tenant string
AuthIdentity string
Logger *slog.Logger
// PerServerTimeout overrides the default per-server discovery timeout
// when greater than zero. Useful for tests and for callers that need a
// tighter or looser bound. When zero, defaultPerServerDiscoveryTimeout
// applies.
PerServerTimeout time.Duration
}
type RuntimeServer struct {
@@ -152,29 +175,121 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
}, nil
}
// defaultPerServerDiscoveryTimeout bounds the time spent discovering tools on
// a single registry-listed server. Tightened to 2s so a slow/unreachable
// server cannot stall every CLI command — a healthy MCP endpoint negotiates
// well under a second. See issue #119.
const defaultPerServerDiscoveryTimeout = 2 * time.Second
func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.ServerDescriptor) ([]RuntimeServer, []RuntimeFailure) {
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
}
perServerTimeout := defaultPerServerDiscoveryTimeout
if s.PerServerTimeout > 0 {
perServerTimeout = s.PerServerTimeout
}
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, perServerTimeout)
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
+23 -2
View File
@@ -20,6 +20,8 @@ import (
"fmt"
"io"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// Category represents a stable error class with a documented exit code.
@@ -198,12 +200,31 @@ func NewInternal(message string, opts ...Option) error {
return newError(CategoryInternal, message, opts...)
}
// ExitCoder is implemented by errors that provide their own exit code.
// Edition-specific error types (e.g. PATError, CLIError) implement this
// so the framework can resolve exit codes without importing edition packages.
type ExitCoder interface {
ExitCode() int
}
// RawStderrError is implemented by errors that must output raw content
// directly to stderr, bypassing all CLI formatting (e.g. "Error:" prefix).
// PAT authorization errors use this to pass JSON through to the desktop runtime.
type RawStderrError interface {
error
RawStderr() string
}
// ExitCode maps any error to a stable exit code.
func ExitCode(err error) int {
var typed *Error
if stderrors.As(err, &typed) {
return typed.ExitCode()
}
var ec ExitCoder
if stderrors.As(err, &ec) {
return ec.ExitCode()
}
return 5
}
@@ -257,7 +278,7 @@ func PrintJSON(w io.Writer, err error) error {
switch typed.ServerDiag.ServerErrorCode {
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
errorPayload["action_url"] = config.GetDeveloperSettingsURL()
}
}
if typed.ServerDiag.TechnicalDetail != "" {
@@ -323,7 +344,7 @@ func PrintHumanAt(w io.Writer, err error, v Verbosity) error {
switch typed.ServerDiag.ServerErrorCode {
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
lines = append(lines, "Action: 开启地址: "+config.GetDeveloperSettingsURL())
}
if len(typed.Actions) > 0 {
+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())
}
}
+289
View File
@@ -0,0 +1,289 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package errors
import (
"encoding/json"
stderrors "errors"
"fmt"
"strings"
)
// ExitCodePermission is the process exit code for PAT authorisation failures.
const ExitCodePermission = 4
// PATError represents a PAT (Personal Action Token) authorization failure
// that should be passed through to stderr as raw JSON without any CLI-layer
// wrapping. The host application (e.g. RewindDesktop) parses the JSON to
// display its own authorisation UI.
type PATError struct {
RawJSON string
}
func (e *PATError) Error() string { return e.RawJSON }
// ExitCode returns the documented exit code for PAT permission errors (4).
func (e *PATError) ExitCode() int { return ExitCodePermission }
// RawStderr returns the raw JSON to be written directly to stderr.
func (e *PATError) RawStderr() string { return e.RawJSON }
// patNoPermissionCodes are PAT error codes that should be passed through
// as transparent PATError without CLI-level wrapping.
var patNoPermissionCodes = map[string]bool{
"PAT_NO_PERMISSION": true,
"PAT_LOW_RISK_NO_PERMISSION": true,
"PAT_MEDIUM_RISK_NO_PERMISSION": true,
"PAT_HIGH_RISK_NO_PERMISSION": true,
}
// patAuthRequiredCodes are error codes that trigger the PAT authorization
// flow (e.g. the server auto-created a CLI app and returned auth details).
var patAuthRequiredCodes = map[string]bool{
"AGENT_CODE_NOT_EXISTS": true,
}
// IsPATError reports whether err is a *PATError.
func IsPATError(err error) bool {
_, ok := err.(*PATError)
return ok
}
// IsPATNoPermissionCode reports whether code is a known PAT permission error code.
func IsPATNoPermissionCode(code string) bool {
return patNoPermissionCodes[code]
}
// ---- DWS gateway auth errors (shared between PAT & general auth) ----------
// dwsGatewayErrors is the set of DWS gateway-level auth error codes.
var dwsGatewayErrors = map[string]bool{
"DWS_SERVICE_UNAUTHORIZED": true,
"DWS_AUTH_SERVICE_FAILED": true,
}
// getDWSGatewayErrorCode extracts a DWS gateway error code from errBody
// (supports both errorCode and error_code field names).
func getDWSGatewayErrorCode(errBody map[string]any) (string, bool) {
for _, key := range []string{"errorCode", "error_code"} {
if code, ok := errBody[key].(string); ok && dwsGatewayErrors[code] {
return code, true
}
}
return "", false
}
// isNotLoggedInError checks if the error body indicates missing authentication.
func isNotLoggedInError(body map[string]any) bool {
if errMsg, ok := body["error"].(string); ok {
if strings.Contains(errMsg, "Missing service_id or access_key") {
return true
}
}
return false
}
// isBusinessError checks if a parsed JSON body represents a business-level error.
func isBusinessError(body map[string]any) bool {
if _, ok := body["error"].(string); ok {
return true
}
if v, ok := body["success"].(bool); ok && !v {
return true
}
if v, ok := body["success"].(string); ok && strings.EqualFold(v, "false") {
return true
}
return false
}
// ---- Classification functions -----------------------------------------------
// ClassifyToolResultContent checks a raw MCP tool result content map for
// DWS gateway auth errors and PAT permission error codes. This is intended
// for use as the edition.Hooks.ClassifyToolResult callback so the framework's
// runner returns a typed error before its generic business-error classification.
//
// Check order: DWS gateway auth > PAT permission.
func ClassifyToolResultContent(content map[string]any) error {
if _, ok := getDWSGatewayErrorCode(content); ok {
raw, _ := json.Marshal(content)
return NewAuth(string(raw),
WithReason("gateway_auth_expired"),
WithHint(authExpiredHint()),
)
}
for _, key := range []string{"code", "errorCode"} {
if code, ok := content[key].(string); ok && patNoPermissionCodes[code] {
return &PATError{RawJSON: cleanPATJSON(content, code)}
}
}
return nil
}
// ClassifyMCPResponseText classifies a text response returned by an MCP tool call.
// Returns a typed error for known gateway auth failures, PAT interceptions,
// and business-level errors embedded in HTTP-200 JSON bodies.
//
// Check order: DWS gateway > PAT permission > generic business error.
func ClassifyMCPResponseText(text string) error {
var body map[string]any
if json.Unmarshal([]byte(text), &body) != nil {
return nil
}
if _, ok := getDWSGatewayErrorCode(body); ok {
return NewAuth(text,
WithReason("gateway_auth_expired"),
WithHint(authExpiredHint()),
)
}
if isNotLoggedInError(body) {
return NewAuth("当前未登录",
WithReason("not_configured"),
WithHint(notLoggedInHint()),
WithActions("dws auth login"),
)
}
for _, key := range []string{"code", "errorCode"} {
if code, ok := body[key].(string); ok && patNoPermissionCodes[code] {
return &PATError{RawJSON: cleanPATJSON(body, code)}
}
}
if isBusinessError(body) {
return NewAPI(text,
WithReason("business_error"),
WithHint(suggestForBusinessErrorText(body)),
)
}
return nil
}
// ---- Hints -----------------------------------------------------------------
func authExpiredHint() string {
return "Re-authenticate: dws auth login"
}
func notLoggedInHint() string {
return "请先登录:dws auth login"
}
func suggestForBusinessErrorText(body map[string]any) string {
msg := ""
if v, ok := body["errorMsg"].(string); ok {
msg = v
} else if v, ok := body["message"].(string); ok {
msg = v
} else if v, ok := body["error"].(string); ok {
msg = v
}
switch {
case strings.Contains(msg, "搜索内容不能为空"):
return "请提供非空搜索关键词: dws doc search --query \"关键词\""
case strings.Contains(msg, "User has no permission to access this email"):
return "请确认邮箱地址正确,查看可用邮箱: dws mail mailbox list"
case strings.Contains(msg, "频率超限") || strings.Contains(msg, "rate limit"):
return "API rate limit exceeded, wait a moment and retry"
case strings.Contains(msg, "参数错误") || strings.Contains(msg, "param error"):
return "Check input parameters. Use --help for available flags"
default:
return "MCP tool returned a business error; check parameters and refer to skill documentation."
}
}
// ---- PAT JSON helpers ------------------------------------------------------
var patTopLevelStrip = map[string]bool{
"success": true, "code": true, "errorCode": true, "error_code": true,
"message": true, "error": true, "trace_id": true, "class": true,
}
func cleanPATJSON(body map[string]any, code string) string {
out := map[string]any{
"success": false,
"code": code,
}
if data, ok := body["data"]; ok {
out["data"] = stripClassFields(data)
} else {
fallback := map[string]any{}
for k, v := range body {
if !patTopLevelStrip[k] {
fallback[k] = v
}
}
if len(fallback) > 0 {
out["data"] = stripClassFields(fallback)
}
}
b, err := json.MarshalIndent(out, "", " ")
if err != nil {
return fmt.Sprintf(`{"success":false,"code":"%s"}`, code)
}
return string(b)
}
// ---- Runner adapter functions ------------------------------------------------
// These match the function signatures referenced by runner.go's PAT check
// framework (ClassifyPatAuthCheck / AsPatAuthCheckError).
// ClassifyPatAuthCheck is the open-source fallback that checks a tool-call
// Content map for PAT permission codes and auth-required codes. Returns a
// non-nil *PATError when the content carries a recognised PAT/auth error.
func ClassifyPatAuthCheck(content map[string]any) *PATError {
for _, key := range []string{"code", "errorCode"} {
if code, ok := content[key].(string); ok {
if patNoPermissionCodes[code] || patAuthRequiredCodes[code] {
return &PATError{RawJSON: cleanPATJSON(content, code)}
}
}
}
return nil
}
// AsPatAuthCheckError extracts a *PATError from an error chain.
func AsPatAuthCheckError(err error) *PATError {
var patErr *PATError
if stderrors.As(err, &patErr) {
return patErr
}
return nil
}
func stripClassFields(v any) any {
switch val := v.(type) {
case map[string]any:
clean := make(map[string]any, len(val))
for k, item := range val {
if k == "class" {
continue
}
clean[k] = stripClassFields(item)
}
return clean
case []any:
clean := make([]any, len(val))
for i, item := range val {
clean[i] = stripClassFields(item)
}
return clean
default:
return v
}
}
+521
View File
@@ -0,0 +1,521 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package errors
import (
stderrors "errors"
"strings"
"testing"
)
// ---------------------------------------------------------------------------
// PATError basic behaviour
// ---------------------------------------------------------------------------
func TestPATError_Implements(t *testing.T) {
t.Parallel()
raw := `{"success":false,"code":"PAT_NO_PERMISSION"}`
pe := &PATError{RawJSON: raw}
if pe.Error() != raw {
t.Errorf("Error() = %q, want %q", pe.Error(), raw)
}
if pe.ExitCode() != ExitCodePermission {
t.Errorf("ExitCode() = %d, want %d", pe.ExitCode(), ExitCodePermission)
}
if pe.RawStderr() != raw {
t.Errorf("RawStderr() = %q, want %q", pe.RawStderr(), raw)
}
}
func TestIsPATError_True(t *testing.T) {
t.Parallel()
err := &PATError{RawJSON: "{}"}
if !IsPATError(err) {
t.Fatal("expected IsPATError to return true for *PATError")
}
}
func TestIsPATError_False(t *testing.T) {
t.Parallel()
err := stderrors.New("some other error")
if IsPATError(err) {
t.Fatal("expected IsPATError to return false for non-PATError")
}
}
// ---------------------------------------------------------------------------
// IsPATNoPermissionCode
// ---------------------------------------------------------------------------
func TestIsPATNoPermissionCode(t *testing.T) {
t.Parallel()
cases := []struct {
code string
want bool
}{
{"PAT_NO_PERMISSION", true},
{"PAT_LOW_RISK_NO_PERMISSION", true},
{"PAT_MEDIUM_RISK_NO_PERMISSION", true},
{"PAT_HIGH_RISK_NO_PERMISSION", true},
{"AGENT_CODE_NOT_EXISTS", false},
{"UNKNOWN_CODE", false},
{"", false},
}
for _, tc := range cases {
if got := IsPATNoPermissionCode(tc.code); got != tc.want {
t.Errorf("IsPATNoPermissionCode(%q) = %v, want %v", tc.code, got, tc.want)
}
}
}
// ---------------------------------------------------------------------------
// getDWSGatewayErrorCode
// ---------------------------------------------------------------------------
func TestGetDWSGatewayErrorCode_ErrorCode(t *testing.T) {
t.Parallel()
body := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED"}
code, ok := getDWSGatewayErrorCode(body)
if !ok || code != "DWS_SERVICE_UNAUTHORIZED" {
t.Errorf("got (%q, %v), want (DWS_SERVICE_UNAUTHORIZED, true)", code, ok)
}
}
func TestGetDWSGatewayErrorCode_ErrorCodeUnderscore(t *testing.T) {
t.Parallel()
body := map[string]any{"error_code": "DWS_AUTH_SERVICE_FAILED"}
code, ok := getDWSGatewayErrorCode(body)
if !ok || code != "DWS_AUTH_SERVICE_FAILED" {
t.Errorf("got (%q, %v), want (DWS_AUTH_SERVICE_FAILED, true)", code, ok)
}
}
func TestGetDWSGatewayErrorCode_Unknown(t *testing.T) {
t.Parallel()
body := map[string]any{"errorCode": "SOME_OTHER_ERROR"}
_, ok := getDWSGatewayErrorCode(body)
if ok {
t.Fatal("expected ok=false for unknown error code")
}
}
func TestGetDWSGatewayErrorCode_Empty(t *testing.T) {
t.Parallel()
body := map[string]any{}
_, ok := getDWSGatewayErrorCode(body)
if ok {
t.Fatal("expected ok=false for empty body")
}
}
// ---------------------------------------------------------------------------
// isNotLoggedInError
// ---------------------------------------------------------------------------
func TestIsNotLoggedInError_True(t *testing.T) {
t.Parallel()
body := map[string]any{"error": "Missing service_id or access_key in request headers"}
if !isNotLoggedInError(body) {
t.Fatal("expected true for Missing service_id message")
}
}
func TestIsNotLoggedInError_False(t *testing.T) {
t.Parallel()
body := map[string]any{"error": "something else happened"}
if isNotLoggedInError(body) {
t.Fatal("expected false for unrelated error message")
}
}
func TestIsNotLoggedInError_NoErrorField(t *testing.T) {
t.Parallel()
body := map[string]any{"message": "Missing service_id or access_key"}
if isNotLoggedInError(body) {
t.Fatal("expected false when error field is absent")
}
}
// ---------------------------------------------------------------------------
// isBusinessError
// ---------------------------------------------------------------------------
func TestIsBusinessError_ErrorField(t *testing.T) {
t.Parallel()
body := map[string]any{"error": "some error message"}
if !isBusinessError(body) {
t.Fatal("expected true when 'error' field is present")
}
}
func TestIsBusinessError_SuccessBoolFalse(t *testing.T) {
t.Parallel()
body := map[string]any{"success": false}
if !isBusinessError(body) {
t.Fatal("expected true when success=false (bool)")
}
}
func TestIsBusinessError_SuccessStringFalse(t *testing.T) {
t.Parallel()
body := map[string]any{"success": "False"}
if !isBusinessError(body) {
t.Fatal("expected true when success=\"False\" (string)")
}
}
func TestIsBusinessError_SuccessTrue(t *testing.T) {
t.Parallel()
body := map[string]any{"success": true, "data": "ok"}
if isBusinessError(body) {
t.Fatal("expected false when success=true")
}
}
func TestIsBusinessError_EmptyBody(t *testing.T) {
t.Parallel()
body := map[string]any{"data": "hello"}
if isBusinessError(body) {
t.Fatal("expected false for body without error indicators")
}
}
// ---------------------------------------------------------------------------
// ClassifyToolResultContent
// ---------------------------------------------------------------------------
func TestClassifyToolResultContent_GatewayAuth(t *testing.T) {
t.Parallel()
content := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED", "message": "expired"}
err := ClassifyToolResultContent(content)
if err == nil {
t.Fatal("expected non-nil error for gateway auth")
}
var typed *Error
if !stderrors.As(err, &typed) {
t.Fatalf("expected *Error, got %T", err)
}
if typed.Category != CategoryAuth {
t.Errorf("Category = %v, want %v", typed.Category, CategoryAuth)
}
if typed.Reason != "gateway_auth_expired" {
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
}
}
func TestClassifyToolResultContent_PATPermission(t *testing.T) {
t.Parallel()
content := map[string]any{
"code": "PAT_NO_PERMISSION",
"data": map[string]any{"desc": "需要授权"},
}
err := ClassifyToolResultContent(content)
if err == nil {
t.Fatal("expected non-nil error for PAT permission")
}
var patErr *PATError
if !stderrors.As(err, &patErr) {
t.Fatalf("expected *PATError, got %T", err)
}
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
t.Errorf("RawJSON should contain PAT_NO_PERMISSION, got: %s", patErr.RawJSON)
}
}
func TestClassifyToolResultContent_NoError(t *testing.T) {
t.Parallel()
content := map[string]any{"success": true, "data": "ok"}
if err := ClassifyToolResultContent(content); err != nil {
t.Fatalf("expected nil error, got %v", err)
}
}
// ---------------------------------------------------------------------------
// ClassifyMCPResponseText
// ---------------------------------------------------------------------------
func TestClassifyMCPResponseText_GatewayAuth(t *testing.T) {
t.Parallel()
text := `{"errorCode":"DWS_SERVICE_UNAUTHORIZED","message":"token expired"}`
err := ClassifyMCPResponseText(text)
if err == nil {
t.Fatal("expected non-nil error")
}
var typed *Error
if !stderrors.As(err, &typed) {
t.Fatalf("expected *Error, got %T", err)
}
if typed.Reason != "gateway_auth_expired" {
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
}
}
func TestClassifyMCPResponseText_NotLoggedIn(t *testing.T) {
t.Parallel()
text := `{"error":"Missing service_id or access_key in headers"}`
err := ClassifyMCPResponseText(text)
if err == nil {
t.Fatal("expected non-nil error")
}
var typed *Error
if !stderrors.As(err, &typed) {
t.Fatalf("expected *Error, got %T", err)
}
if typed.Reason != "not_configured" {
t.Errorf("Reason = %q, want not_configured", typed.Reason)
}
}
func TestClassifyMCPResponseText_PATPermission(t *testing.T) {
t.Parallel()
text := `{"code":"PAT_HIGH_RISK_NO_PERMISSION","data":{"desc":"high risk"}}`
err := ClassifyMCPResponseText(text)
if err == nil {
t.Fatal("expected non-nil error")
}
var patErr *PATError
if !stderrors.As(err, &patErr) {
t.Fatalf("expected *PATError, got %T", err)
}
if !strings.Contains(patErr.RawJSON, "PAT_HIGH_RISK_NO_PERMISSION") {
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
}
}
func TestClassifyMCPResponseText_BusinessError(t *testing.T) {
t.Parallel()
text := `{"success":false,"errorMsg":"搜索内容不能为空"}`
err := ClassifyMCPResponseText(text)
if err == nil {
t.Fatal("expected non-nil error")
}
var typed *Error
if !stderrors.As(err, &typed) {
t.Fatalf("expected *Error, got %T", err)
}
if typed.Reason != "business_error" {
t.Errorf("Reason = %q, want business_error", typed.Reason)
}
if !strings.Contains(typed.Hint, "搜索关键词") {
t.Errorf("Hint should contain search suggestion, got: %s", typed.Hint)
}
}
func TestClassifyMCPResponseText_InvalidJSON(t *testing.T) {
t.Parallel()
text := "not json at all"
if err := ClassifyMCPResponseText(text); err != nil {
t.Fatalf("expected nil for invalid JSON, got %v", err)
}
}
func TestClassifyMCPResponseText_NoError(t *testing.T) {
t.Parallel()
text := `{"success":true,"data":"hello"}`
if err := ClassifyMCPResponseText(text); err != nil {
t.Fatalf("expected nil for success response, got %v", err)
}
}
// ---------------------------------------------------------------------------
// ClassifyPatAuthCheck
// ---------------------------------------------------------------------------
func TestClassifyPatAuthCheck_PATNoPermission(t *testing.T) {
t.Parallel()
content := map[string]any{"code": "PAT_NO_PERMISSION", "data": map[string]any{"flowId": "f1"}}
patErr := ClassifyPatAuthCheck(content)
if patErr == nil {
t.Fatal("expected non-nil *PATError")
}
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
}
}
func TestClassifyPatAuthCheck_AgentCodeNotExists(t *testing.T) {
t.Parallel()
content := map[string]any{"errorCode": "AGENT_CODE_NOT_EXISTS", "data": map[string]any{"clientId": "c1"}}
patErr := ClassifyPatAuthCheck(content)
if patErr == nil {
t.Fatal("expected non-nil *PATError for AGENT_CODE_NOT_EXISTS")
}
if !strings.Contains(patErr.RawJSON, "AGENT_CODE_NOT_EXISTS") {
t.Errorf("RawJSON should contain AGENT_CODE_NOT_EXISTS, got: %s", patErr.RawJSON)
}
}
func TestClassifyPatAuthCheck_NoMatch(t *testing.T) {
t.Parallel()
content := map[string]any{"code": "SOME_BUSINESS_ERROR", "message": "oops"}
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
t.Fatalf("expected nil, got %v", patErr)
}
}
func TestClassifyPatAuthCheck_EmptyContent(t *testing.T) {
t.Parallel()
content := map[string]any{}
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
t.Fatalf("expected nil for empty content, got %v", patErr)
}
}
// ---------------------------------------------------------------------------
// AsPatAuthCheckError
// ---------------------------------------------------------------------------
func TestAsPatAuthCheckError_Wrapped(t *testing.T) {
t.Parallel()
inner := &PATError{RawJSON: `{"code":"PAT_NO_PERMISSION"}`}
wrapped := stderrors.Join(stderrors.New("context"), inner)
got := AsPatAuthCheckError(wrapped)
if got == nil {
t.Fatal("expected non-nil *PATError from wrapped error")
}
if got.RawJSON != inner.RawJSON {
t.Errorf("RawJSON = %q, want %q", got.RawJSON, inner.RawJSON)
}
}
func TestAsPatAuthCheckError_NotPAT(t *testing.T) {
t.Parallel()
err := stderrors.New("just a plain error")
if got := AsPatAuthCheckError(err); got != nil {
t.Fatalf("expected nil for non-PAT error, got %v", got)
}
}
// ---------------------------------------------------------------------------
// cleanPATJSON
// ---------------------------------------------------------------------------
func TestCleanPATJSON_WithData(t *testing.T) {
t.Parallel()
body := map[string]any{
"success": false,
"code": "PAT_NO_PERMISSION",
"data": map[string]any{
"desc": "需要授权",
"flowId": "f123",
"class": "com.foo.Bar",
},
}
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
if !strings.Contains(result, "PAT_NO_PERMISSION") {
t.Errorf("expected code in output, got: %s", result)
}
if !strings.Contains(result, "flowId") {
t.Errorf("expected flowId in data, got: %s", result)
}
if strings.Contains(result, "class") {
t.Errorf("expected class field to be stripped, got: %s", result)
}
}
func TestCleanPATJSON_WithoutData(t *testing.T) {
t.Parallel()
body := map[string]any{
"success": false,
"code": "PAT_NO_PERMISSION",
"message": "no permission",
"extra": "value",
}
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
if !strings.Contains(result, "extra") {
t.Errorf("expected extra field in fallback data, got: %s", result)
}
// Top-level stripped fields should not appear
if strings.Contains(result, `"message"`) {
t.Errorf("expected message to be stripped from top level, got: %s", result)
}
}
// ---------------------------------------------------------------------------
// stripClassFields
// ---------------------------------------------------------------------------
func TestStripClassFields_Map(t *testing.T) {
t.Parallel()
input := map[string]any{
"name": "test",
"class": "com.foo.Bar",
"nested": map[string]any{
"value": 42,
"class": "com.baz.Qux",
},
}
result := stripClassFields(input).(map[string]any)
if _, ok := result["class"]; ok {
t.Error("top-level class should be removed")
}
nested := result["nested"].(map[string]any)
if _, ok := nested["class"]; ok {
t.Error("nested class should be removed")
}
if nested["value"] != 42 {
t.Errorf("nested value should be preserved, got %v", nested["value"])
}
}
func TestStripClassFields_Array(t *testing.T) {
t.Parallel()
input := []any{
map[string]any{"id": 1, "class": "Foo"},
map[string]any{"id": 2},
}
result := stripClassFields(input).([]any)
first := result[0].(map[string]any)
if _, ok := first["class"]; ok {
t.Error("class in array element should be removed")
}
if first["id"] != 1 {
t.Error("other fields in array element should be preserved")
}
}
func TestStripClassFields_Scalar(t *testing.T) {
t.Parallel()
if stripClassFields("hello") != "hello" {
t.Error("scalar string should pass through unchanged")
}
if stripClassFields(42) != 42 {
t.Error("scalar int should pass through unchanged")
}
}
// ---------------------------------------------------------------------------
// suggestForBusinessErrorText
// ---------------------------------------------------------------------------
func TestSuggestForBusinessErrorText(t *testing.T) {
t.Parallel()
cases := []struct {
body map[string]any
contains string
}{
{map[string]any{"errorMsg": "搜索内容不能为空"}, "搜索关键词"},
{map[string]any{"message": "User has no permission to access this email"}, "邮箱"},
{map[string]any{"error": "频率超限"}, "rate limit"},
{map[string]any{"errorMsg": "参数错误"}, "parameters"},
{map[string]any{"error": "unknown"}, "business error"},
}
for _, tc := range cases {
hint := suggestForBusinessErrorText(tc.body)
if !strings.Contains(strings.ToLower(hint), strings.ToLower(tc.contains)) {
t.Errorf("suggestForBusinessErrorText(%v) = %q, want to contain %q", tc.body, hint, tc.contains)
}
}
}
+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
}
+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
+36 -1
View File
@@ -155,5 +155,40 @@
"返回数据缺少 uploadUrl 或 fileToken": "Response data missing uploadUrl or fileToken",
"附件工作流": "Attachment workflow",
"页码 (必填)": "page number (required)",
"🔐 登录钉钉": "🔐 Login to DingTalk"
"⚠️ 无法检查 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.",
"管理员操作入口:": "Admin settings: ",
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "CLI data access is not enabled for this organization, please contact admin to enable it",
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ CLI data access is not enabled for this organization, please submit an authorization request in the browser...",
"✅ 权限已开启,继续登录...": "✅ Permission enabled, continuing login...",
"等待管理员审批中": "Waiting for admin approval",
"等待提交申请中": "Waiting to submit request",
"操作超时,请重新登录": "Operation timed out, please re-login",
"检查组织 CLI 授权状态...": "Checking organization CLI auth status...",
"🔐 登录钉钉": "🔐 Login to DingTalk",
"插件管理": "Manage plugins",
"列出已安装的插件": "List installed plugins",
"安装插件": "Install a plugin",
"查看插件详情": "Show plugin details",
"启用插件": "Enable a plugin",
"禁用插件": "Disable a plugin",
"卸载已安装的插件": "Remove an installed plugin",
"校验 plugin.json": "Validate a plugin.json",
"脚手架生成新插件目录": "Scaffold a new plugin directory",
"将本地目录注册为开发态插件": "Register a local directory as a dev plugin",
"管理插件配置": "Manage plugin configuration",
"设置插件配置项": "Set a plugin config value",
"读取插件配置项": "Get a plugin config value",
"列出插件所有配置项": "List all config values for a plugin",
"删除插件配置项": "Remove a plugin config value",
"将插件 stdio server 编译为原生二进制": "Build plugin's stdio server into a native binary",
"覆盖 OAuth 客户端 ID (钉钉 AppKey)": "Override OAuth client ID (DingTalk AppKey)",
"覆盖 OAuth 客户端密钥 (钉钉 AppSecret)": "Override OAuth client secret (DingTalk AppSecret)",
"查看任意命令的帮助信息": "Help about any command",
"显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。": "Help provides help for any command in the application.\nSimply type dws help [path to command] for full details."
}
+36 -1
View File
@@ -155,5 +155,40 @@
"返回数据缺少 uploadUrl 或 fileToken": "返回数据缺少 uploadUrl 或 fileToken",
"附件工作流": "附件工作流",
"页码 (必填)": "页码 (必填)",
"🔐 登录钉钉": "🔐 登录钉钉"
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ 无法检查 CLI 数据访问权限状态",
" 请检查网络连接后重试。": " 请检查网络连接后重试。",
"检查 CLI 授权状态失败": "检查 CLI 授权状态失败",
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ 该组织尚未开启 CLI 数据访问权限",
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。",
" 组织主管理员:": " 组织主管理员:",
" 请联系组织主管理员开启后重新登录。": " 请联系组织主管理员开启后重新登录。",
"管理员操作入口:": "管理员操作入口:",
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "该组织尚未开启 CLI 数据访问权限,请联系管理员开启",
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...",
"✅ 权限已开启,继续登录...": "✅ 权限已开启,继续登录...",
"等待管理员审批中": "等待管理员审批中",
"等待提交申请中": "等待提交申请中",
"操作超时,请重新登录": "操作超时,请重新登录",
"检查组织 CLI 授权状态...": "检查组织 CLI 授权状态...",
"🔐 登录钉钉": "🔐 登录钉钉",
"插件管理": "插件管理",
"列出已安装的插件": "列出已安装的插件",
"安装插件": "安装插件",
"查看插件详情": "查看插件详情",
"启用插件": "启用插件",
"禁用插件": "禁用插件",
"卸载已安装的插件": "卸载已安装的插件",
"校验 plugin.json": "校验 plugin.json",
"脚手架生成新插件目录": "脚手架生成新插件目录",
"将本地目录注册为开发态插件": "将本地目录注册为开发态插件",
"管理插件配置": "管理插件配置",
"设置插件配置项": "设置插件配置项",
"读取插件配置项": "读取插件配置项",
"列出插件所有配置项": "列出插件所有配置项",
"删除插件配置项": "删除插件配置项",
"将插件 stdio server 编译为原生二进制": "将插件 stdio server 编译为原生二进制",
"覆盖 OAuth 客户端 ID (钉钉 AppKey)": "覆盖 OAuth 客户端 ID (钉钉 AppKey)",
"覆盖 OAuth 客户端密钥 (钉钉 AppSecret)": "覆盖 OAuth 客户端密钥 (钉钉 AppSecret)",
"查看任意命令的帮助信息": "查看任意命令的帮助信息",
"显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。": "显示任意命令的帮助文案。\n用法:dws help [命令路径] 查看完整说明。"
}
+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 (
+14 -6
View File
@@ -16,6 +16,7 @@ package logging
import (
"context"
"log/slog"
"runtime"
"time"
)
@@ -52,14 +53,15 @@ func LogRequestBody(logger *slog.Logger, method, executionId string, toolName st
)
}
// LogResponse logs a JSON-RPC response at Debug level.
func LogResponse(logger *slog.Logger, method, endpoint string, statusCode int, respSize int, duration time.Duration, err error) {
// LogResponse logs a JSON-RPC response at Debug level (Warn on error).
func LogResponse(logger *slog.Logger, method, endpoint, executionId string, statusCode int, respSize int, duration time.Duration, err error) {
if logger == nil {
return
}
attrs := []slog.Attr{
slog.String("method", method),
slog.String("endpoint", redactEndpoint(endpoint)),
slog.String("execution_id", executionId),
slog.Int("status", statusCode),
slog.Int("resp_size", respSize),
slog.String("duration", duration.Truncate(time.Millisecond).String()),
@@ -137,18 +139,24 @@ func LogErrorClassified(logger *slog.Logger, method, executionId, category, reas
}
// LogCommandStart logs the beginning of a command execution.
func LogCommandStart(logger *slog.Logger, executionId, command, product, tool, version string, authPresent bool) {
func LogCommandStart(logger *slog.Logger, executionId, product, tool, endpoint, version string, authPresent bool, timeoutSec int) {
if logger == nil {
return
}
logger.Info("command_start",
attrs := []slog.Attr{
slog.String("execution_id", executionId),
slog.String("command", command),
slog.String("product", product),
slog.String("tool", tool),
slog.String("endpoint", redactEndpoint(endpoint)),
slog.String("cli_version", version),
slog.String("os", runtime.GOOS),
slog.String("arch", runtime.GOARCH),
slog.Bool("auth_token_present", authPresent),
)
}
if timeoutSec > 0 {
attrs = append(attrs, slog.Int("timeout_sec", timeoutSec))
}
logger.LogAttrs(context.TODO(), slog.LevelInfo, "command_start", attrs...)
}
// LogCommandEnd logs the end of a command execution.
+4 -4
View File
@@ -52,7 +52,7 @@ func TestLogResponseSuccess(t *testing.T) {
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", 200, 1024, 150*time.Millisecond, nil)
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", "exec-1", 200, 1024, 150*time.Millisecond, nil)
out := buf.String()
if !strings.Contains(out, "jsonrpc_response") {
@@ -72,7 +72,7 @@ func TestLogResponseError(t *testing.T) {
var buf bytes.Buffer
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", 500, 0, 2*time.Second, errors.New("connection refused"))
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", "exec-2", 500, 0, 2*time.Second, errors.New("connection refused"))
out := buf.String()
if !strings.Contains(out, "WARN") {
@@ -87,12 +87,12 @@ func TestLogRequestNilLogger(t *testing.T) {
t.Parallel()
// Should not panic
LogRequest(nil, "test", "http://localhost", "", 0)
LogResponse(nil, "test", "http://localhost", 200, 0, 0, nil)
LogResponse(nil, "test", "http://localhost", "", 200, 0, 0, nil)
LogRequestBody(nil, "tools/call", "exec-1", "tool", nil)
LogResponseBody(nil, "tools/call", "exec-1", 200, nil, "")
LogRetryAttempt(nil, "tools/call", "exec-1", 0, 2, 429, 0, nil)
LogErrorClassified(nil, "tools/call", "exec-1", "api", "timeout", 0, 0, true, "")
LogCommandStart(nil, "exec-1", "dws test", "doc", "list", "1.0.0", false)
LogCommandStart(nil, "exec-1", "doc", "list", "https://mcp.example.com", "1.0.0", false, 0)
LogCommandEnd(nil, "exec-1", "doc", "list", true, 0, "", "")
}
+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())
}
+137
View File
@@ -0,0 +1,137 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package pat
import (
"context"
"encoding/json"
"fmt"
"os"
"github.com/fatih/color"
"github.com/spf13/cobra"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
var validGrantTypes = map[string]bool{
"once": true,
"session": true,
"permanent": true,
}
func newChmodCommand(caller edition.ToolCaller) *cobra.Command {
chmodCmd := &cobra.Command{
Use: "chmod <scope>...",
Short: "授予指定权限",
Long: `授予指定 scope 的操作权限。
scope 格式: <product>.<entity>:<permission>
例: aitable.record:read chat.group:write calendar.event:read
grantType 规则:
once 一次性,执行一次后自动失效
session 当前会话有效(默认),需要 --session-id
permanent 永久有效`,
Args: cobra.MinimumNArgs(1),
Example: ` dws pat chmod aitable.record:read --agentCode agt-xxxx --grant-type session --session-id session-xxx
dws pat chmod chat.message:list --grant-type once --agentCode agt-xxxx
dws pat chmod aitable.record:read aitable.record:write --agentCode agt-xxxx --grant-type permanent`,
RunE: func(cmd *cobra.Command, args []string) error {
agentCode, _ := cmd.Flags().GetString("agentCode")
if agentCode == "" {
return fmt.Errorf("flag --agentCode is required\n hint: dws pat chmod <scope>... --agentCode <id>")
}
scopes := args
grantType, _ := cmd.Flags().GetString("grant-type")
sessionID, _ := cmd.Flags().GetString("session-id")
if !validGrantTypes[grantType] {
return fmt.Errorf("invalid --grant-type %q, must be one of: once, session, permanent", grantType)
}
if grantType == "session" && sessionID == "" && os.Getenv("DWS_SESSION_ID") == "" {
return fmt.Errorf("--session-id is required when --grant-type is session\n hint: dws pat chmod <scope> --agentCode <id> --grant-type session --session-id <id>")
}
if caller != nil && caller.DryRun() {
bold := color.New(color.FgYellow, color.Bold)
bold.Println("[DRY-RUN] Preview only, not executed:")
fmt.Printf("%-16s%s\n", "Tool:", "个人授权")
fmt.Printf("%-16s%s\n", "AgentCode:", agentCode)
fmt.Printf("%-16s%v\n", "Scope:", scopes)
fmt.Printf("%-16s%s\n", "GrantType:", grantType)
if sessionID != "" {
fmt.Printf("%-16s%s\n", "SessionID:", sessionID)
}
return nil
}
if caller == nil {
return fmt.Errorf("internal error: tool runtime not initialized")
}
toolArgs := map[string]any{
"agentCode": agentCode,
"scope": scopes,
"grantType": grantType,
}
if sessionID == "" {
sessionID = os.Getenv("DWS_SESSION_ID")
}
if sessionID != "" {
toolArgs["sessionId"] = sessionID
}
ctx := context.Background()
result, err := caller.CallTool(ctx, "pat", "个人授权", toolArgs)
if err != nil {
return fmt.Errorf("pat chmod failed: %w", err)
}
return handleToolResult(result)
},
}
chmodCmd.Flags().String("agentCode", "", "Agent 唯一标识(必填)")
_ = chmodCmd.MarkFlagRequired("agentCode")
chmodCmd.Flags().String("grant-type", "session", "授权策略: once|session|permanent")
chmodCmd.Flags().String("session-id", "", "会话标识(session 模式下必填)")
return chmodCmd
}
// handleToolResult processes a ToolResult and writes output to stdout.
func handleToolResult(result *edition.ToolResult) error {
if result == nil {
return fmt.Errorf("empty tool result")
}
for _, c := range result.Content {
if c.Type != "text" || c.Text == "" {
continue
}
if respErr := apperrors.ClassifyMCPResponseText(c.Text); respErr != nil {
return respErr
}
fmt.Println(c.Text)
return nil
}
data, err := json.MarshalIndent(result, "", " ")
if err != nil {
return fmt.Errorf("failed to marshal result: %w", err)
}
fmt.Println(string(data))
return nil
}
+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 pat implements the "dws pat" command group for PAT (Personal Action
// Token) authorization management.
package pat
import (
"github.com/spf13/cobra"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
// RegisterCommands adds the pat command tree to rootCmd.
func RegisterCommands(root *cobra.Command, c edition.ToolCaller) {
patCmd := &cobra.Command{
Use: "pat",
Short: "行为授权管理",
Long: `管理行为授权(PAT)。
命令结构:
dws pat chmod <scope>... 授予指定权限`,
RunE: cmdutil.GroupRunE,
}
patCmd.AddCommand(newChmodCommand(c))
root.AddCommand(patCmd)
}
@@ -244,6 +244,194 @@ func TestFullPipelineEndToEnd(t *testing.T) {
}
}
// TestFullFivePhasePipeline exercises all five phases in order:
// Register → PreParse → PostParse → PreRequest → PostResponse,
// simulating a complete command lifecycle from registration through
// response output.
func TestFullFivePhasePipeline(t *testing.T) {
engine := pipeline.NewEngine()
engine.RegisterAll(
RegisterHandler{},
AliasHandler{},
StickyHandler{},
ParamNameHandler{},
ParamValueHandler{},
PreRequestHandler{},
PostResponseHandler{},
)
// Verify all five phases have handlers.
for _, phase := range []pipeline.Phase{
pipeline.Register,
pipeline.PreParse,
pipeline.PostParse,
pipeline.PreRequest,
pipeline.PostResponse,
} {
if !engine.HasHandlers(phase) {
t.Fatalf("engine missing handlers for phase %v", phase)
}
}
// Phase 1: Register — command tree being built.
ctx := &pipeline.Context{
Command: "aitable",
}
if err := engine.RunPhase(pipeline.Register, ctx); err != nil {
t.Fatalf("Register error: %v", err)
}
// Phase 2: PreParse — fix raw argv.
ctx.Args = []string{
"--userId", "u001",
"--pageSize50",
"--verbosetrue",
}
ctx.FlagSpecs = flagSpecs("user-id", "page-size", "verbose")
if err := engine.RunPhase(pipeline.PreParse, ctx); err != nil {
t.Fatalf("PreParse error: %v", err)
}
want := "--user-id u001 --page-size 50 --verbose true"
got := strings.Join(ctx.Args, " ")
if got != want {
t.Errorf("after PreParse: Args = %q, want %q", got, want)
}
preParseCorrections := len(ctx.Corrections)
// Phase 3: PostParse — simulate Cobra having parsed the corrected
// args into structured params, then normalise values.
ctx.Command = "aitable.query_records"
ctx.Params = map[string]any{
"user_id": "u001",
"page_size": "1,000",
"verbose": "yes",
}
ctx.Schema = map[string]any{
"properties": map[string]any{
"user_id": map[string]any{"type": "string"},
"page_size": map[string]any{"type": "integer"},
"verbose": map[string]any{"type": "boolean"},
},
}
if err := engine.RunPhase(pipeline.PostParse, ctx); err != nil {
t.Fatalf("PostParse error: %v", err)
}
if got := ctx.Params["verbose"]; got != true {
t.Errorf("verbose = %v (%T), want true (bool)", got, got)
}
if got := ctx.Params["page_size"]; got != int64(1000) {
t.Errorf("page_size = %v, want 1000", got)
}
postParseCorrections := len(ctx.Corrections) - preParseCorrections
if postParseCorrections != 2 {
t.Errorf("PostParse corrections = %d, want 2", postParseCorrections)
}
// Phase 4: PreRequest — inspect final payload before dispatch.
ctx.Payload = ctx.Params
if err := engine.RunPhase(pipeline.PreRequest, ctx); err != nil {
t.Fatalf("PreRequest error: %v", err)
}
// Verify payload was not corrupted.
if ctx.Payload["user_id"] != "u001" {
t.Error("PreRequest corrupted Payload")
}
// Phase 5: PostResponse — process response before output.
ctx.Response = map[string]any{
"records": []any{
map[string]any{"id": "rec001", "fields": map[string]any{"name": "test"}},
},
"total": 1,
}
if err := engine.RunPhase(pipeline.PostResponse, ctx); err != nil {
t.Fatalf("PostResponse error: %v", err)
}
// Verify response was not corrupted.
if ctx.Response["total"] != 1 {
t.Error("PostResponse corrupted Response")
}
}
// TestFullFivePhasePipelineWithEngineRun exercises all five phases
// using Engine.Run (single shot) to verify the ordering is correct
// end-to-end.
func TestFullFivePhasePipelineWithEngineRun(t *testing.T) {
var seq []string
engine := pipeline.NewEngine()
engine.RegisterAll(
&phaseTracker{name: "reg", phase: pipeline.Register, seq: &seq},
&phaseTracker{name: "pre-parse", phase: pipeline.PreParse, seq: &seq},
&phaseTracker{name: "post-parse", phase: pipeline.PostParse, seq: &seq},
&phaseTracker{name: "pre-req", phase: pipeline.PreRequest, seq: &seq},
&phaseTracker{name: "post-resp", phase: pipeline.PostResponse, seq: &seq},
)
ctx := &pipeline.Context{Command: "test.tool"}
if err := engine.Run(ctx); err != nil {
t.Fatalf("Engine.Run error: %v", err)
}
want := "reg,pre-parse,post-parse,pre-req,post-resp"
got := strings.Join(seq, ",")
if got != want {
t.Errorf("phase execution order = %q, want %q", got, want)
}
}
// TestFivePhasePipelineCorrectHandlerCounts verifies that the
// production-equivalent engine has the expected handler distribution.
func TestFivePhasePipelineCorrectHandlerCounts(t *testing.T) {
engine := pipeline.NewEngine()
engine.RegisterAll(
RegisterHandler{},
AliasHandler{},
StickyHandler{},
ParamNameHandler{},
ParamValueHandler{},
PreRequestHandler{},
PostResponseHandler{},
)
tests := []struct {
phase pipeline.Phase
want int
}{
{pipeline.Register, 1},
{pipeline.PreParse, 3},
{pipeline.PostParse, 1},
{pipeline.PreRequest, 1},
{pipeline.PostResponse, 1},
}
for _, tt := range tests {
if got := len(engine.Handlers(tt.phase)); got != tt.want {
t.Errorf("Handlers(%v) = %d, want %d", tt.phase, got, tt.want)
}
}
if got := engine.HandlerCount(); got != 7 {
t.Errorf("HandlerCount = %d, want 7", got)
}
}
// phaseTracker is a test helper that records its name when Handle is called.
type phaseTracker struct {
name string
phase pipeline.Phase
seq *[]string
}
func (h *phaseTracker) Name() string { return h.name }
func (h *phaseTracker) Phase() pipeline.Phase { return h.phase }
func (h *phaseTracker) Handle(_ *pipeline.Context) error {
*h.seq = append(*h.seq, h.name)
return nil
}
// TestPreParseDoesNotBreakValidArgs verifies that valid, correctly
// formatted args pass through the pipeline without modification.
func TestPreParseDoesNotBreakValidArgs(t *testing.T) {
+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)
}
}
+159
View File
@@ -0,0 +1,159 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package plugin
import (
"encoding/json"
"log/slog"
"os"
"path/filepath"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
)
// UserContext holds the minimal user identity fields injected into
// stdio plugin subprocesses via environment variables.
type UserContext struct {
UserID string
CorpID string
}
// StdioServerClient pairs a transport.StdioClient with its server key.
type StdioServerClient struct {
Key string
Client *transport.StdioClient
}
// StdioClients returns StdioClient instances for all stdio-type MCP
// servers declared by this plugin. uc is the current user's identity;
// if non-nil, DWS_USER_ID and DWS_CORP_ID are injected as environment
// variables so that the subprocess can identify the caller without
// implementing its own auth.
func (p *Plugin) StdioClients(uc *UserContext) []StdioServerClient {
var clients []StdioServerClient
for key, srv := range p.Manifest.MCPServers {
if srv.Type != "stdio" {
continue
}
command := srv.Command
if command == "" {
slog.Warn("plugin: stdio server missing command",
"plugin", p.Manifest.Name, "server", key)
continue
}
// Expand ${DWS_PLUGIN_ROOT} in command and args.
command = expandPluginVars(command, p.Root)
args := make([]string, len(srv.Args))
for i, a := range srv.Args {
args[i] = expandPluginVars(a, p.Root)
}
env := make(map[string]string)
for k, v := range srv.Env {
env[k] = expandPluginVars(v, p.Root)
}
env["DWS_PLUGIN_ROOT"] = p.Root
env["DWS_PLUGIN_DATA"] = filepath.Join(filepath.Dir(filepath.Dir(p.Root)), "data", p.Manifest.Name)
// Inject user identity so the subprocess knows who is calling.
if uc != nil {
if uc.UserID != "" {
env["DWS_USER_ID"] = uc.UserID
}
if uc.CorpID != "" {
env["DWS_CORP_ID"] = uc.CorpID
}
}
sc := transport.NewStdioClient(command, args, env)
clients = append(clients, StdioServerClient{Key: key, Client: sc})
}
return clients
}
// expandPluginVars replaces ${DWS_PLUGIN_ROOT} with the actual plugin
// root path and ${DWS_PLUGIN_DATA} with the data directory.
func expandPluginVars(s, root string) string {
s = strings.ReplaceAll(s, "${DWS_PLUGIN_ROOT}", root)
dataDir := filepath.Join(filepath.Dir(filepath.Dir(root)), "data")
s = strings.ReplaceAll(s, "${DWS_PLUGIN_DATA}", dataDir)
return os.Expand(s, os.Getenv)
}
// ToServerDescriptors converts a loaded plugin's MCP servers into
// market.ServerDescriptor values suitable for SetDynamicServers.
// Only streamable-http servers are converted; stdio servers are
// skipped (they require the stdio transport extension).
func (p *Plugin) ToServerDescriptors() []market.ServerDescriptor {
var descriptors []market.ServerDescriptor
for key, srv := range p.Manifest.MCPServers {
if srv.Type != "streamable-http" {
slog.Debug("plugin: skipping non-http server",
"plugin", p.Manifest.Name,
"server", key,
"type", srv.Type,
)
continue
}
overlay := market.CLIOverlay{}
if len(srv.CLI) > 0 {
if err := json.Unmarshal(srv.CLI, &overlay); err != nil {
slog.Warn("plugin: failed to parse CLIOverlay",
"plugin", p.Manifest.Name,
"server", key,
"error", err,
)
}
}
// Ensure the overlay has an ID — fall back to server key.
if overlay.ID == "" {
overlay.ID = key
}
if overlay.Command == "" {
overlay.Command = key
}
source := "plugin"
if p.IsManaged {
source = "plugin-managed"
}
// Resolve headers: expand environment variable references (e.g. ${DASHSCOPE_API_KEY}).
var resolvedHeaders map[string]string
if len(srv.Headers) > 0 {
resolvedHeaders = make(map[string]string, len(srv.Headers))
for headerKey, headerVal := range srv.Headers {
resolvedHeaders[headerKey] = expandPluginVars(headerVal, p.Root)
}
}
descriptors = append(descriptors, market.ServerDescriptor{
Key: key,
DisplayName: p.Manifest.Name + "/" + key,
Description: p.Manifest.Description,
Endpoint: srv.Endpoint,
Source: source,
CLI: overlay,
HasCLIMeta: len(srv.CLI) > 0,
AuthHeaders: resolvedHeaders,
})
}
return descriptors
}
+124
View File
@@ -0,0 +1,124 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package plugin
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"os/exec"
"path/filepath"
"strings"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
)
const defaultHookTimeout = 30 * time.Second
// HookAdapter wraps a plugin hook entry as a pipeline.Handler.
type HookAdapter struct {
pluginName string
entry HookEntry
phase pipeline.Phase
timeout time.Duration
}
// NewHookAdapter creates a pipeline handler from a plugin hook entry.
func NewHookAdapter(pluginName string, entry HookEntry) *HookAdapter {
phase := parsePhase(entry.Phase)
timeout := defaultHookTimeout
if entry.Timeout > 0 {
timeout = time.Duration(entry.Timeout) * time.Second
}
return &HookAdapter{
pluginName: pluginName,
entry: entry,
phase: phase,
timeout: timeout,
}
}
func (h *HookAdapter) Name() string {
return fmt.Sprintf("plugin-hook:%s/%s", h.pluginName, h.entry.Phase)
}
func (h *HookAdapter) Phase() pipeline.Phase {
return h.phase
}
func (h *HookAdapter) Handle(ctx *pipeline.Context) error {
// Check matcher: if set, only run for matching commands.
if h.entry.Matcher != "" {
matched, err := filepath.Match(h.entry.Matcher, ctx.Command)
if err != nil || !matched {
return nil // skip silently
}
}
// Serialize context to JSON for the hook's stdin.
input, err := json.Marshal(map[string]any{
"command": ctx.Command,
"params": ctx.Params,
"args": ctx.Args,
})
if err != nil {
slog.Warn("plugin hook: failed to serialize context",
"plugin", h.pluginName, "error", err)
return nil
}
timeoutCtx, cancel := context.WithTimeout(context.Background(), h.timeout)
defer cancel()
cmd := exec.CommandContext(timeoutCtx, "sh", "-c", h.entry.Command)
cmd.Stdin = strings.NewReader(string(input))
output, err := cmd.CombinedOutput()
if err != nil {
if exitErr, ok := err.(*exec.ExitError); ok {
code := exitErr.ExitCode()
if code == 2 {
// Exit 2 = abort pipeline.
return fmt.Errorf("plugin hook %s/%s aborted: %s",
h.pluginName, h.entry.Phase, strings.TrimSpace(string(output)))
}
}
slog.Warn("plugin hook failed",
"plugin", h.pluginName,
"phase", h.entry.Phase,
"error", err,
"output", string(output),
)
return nil // non-fatal: log warning and continue
}
return nil
}
func parsePhase(s string) pipeline.Phase {
switch strings.TrimSpace(strings.ToLower(s)) {
case "pre-parse":
return pipeline.PreParse
case "post-parse":
return pipeline.PostParse
case "pre-request":
return pipeline.PreRequest
case "post-response":
return pipeline.PostResponse
default:
return pipeline.PreRequest
}
}
+940
View File
@@ -0,0 +1,940 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package plugin
import (
"bytes"
"encoding/json"
"fmt"
"log/slog"
"net/url"
"os"
"os/exec"
"path/filepath"
"runtime"
"strings"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
)
// Loader scans plugin directories and returns loaded, validated plugins.
type Loader struct {
// PluginsDir is the root directory for all plugins.
// Defaults to ~/.dws/plugins/.
PluginsDir string
// CLIVersion is the current CLI version, used for
// minCLIVersion compatibility checks.
CLIVersion string
}
// NewLoader creates a Loader with default paths.
func NewLoader(cliVersion string) *Loader {
home, _ := os.UserHomeDir()
return &Loader{
PluginsDir: filepath.Join(home, ".dws", "plugins"),
CLIVersion: cliVersion,
}
}
// Settings holds user preferences for plugin management.
type Settings struct {
EnabledPlugins map[string]bool `json:"enabledPlugins,omitempty"`
PluginConfigs map[string]map[string]any `json:"pluginConfigs,omitempty"`
PluginAutoUpdate bool `json:"pluginAutoUpdate,omitempty"`
DevPlugins map[string]string `json:"devPlugins,omitempty"` // name → absolute path
}
// LoadManaged scans ~/.dws/plugins/managed/ and returns all valid
// official plugins. Managed plugins are always enabled.
func (l *Loader) LoadManaged() []*Plugin {
managedDir := filepath.Join(l.PluginsDir, "managed")
return l.scanDir(managedDir, true)
}
// LoadUser scans ~/.dws/plugins/user/ and returns enabled user plugins.
func (l *Loader) LoadUser() []*Plugin {
userDir := filepath.Join(l.PluginsDir, "user")
settings := l.loadSettings()
var plugins []*Plugin
// User plugins may be nested: user/{workspace}/{name}/
entries, err := os.ReadDir(userDir)
if err != nil {
if !os.IsNotExist(err) {
slog.Debug("plugin: cannot read user dir", "path", userDir, "error", err)
}
return nil
}
for _, entry := range entries {
if !entry.IsDir() {
continue
}
entryPath := filepath.Join(userDir, entry.Name())
// Check if this is a direct plugin directory (has plugin.json)
if _, err := os.Stat(filepath.Join(entryPath, "plugin.json")); err == nil {
p := l.loadPlugin(entryPath, false)
if p != nil && isPluginEnabled(settings, p.Manifest.Name) {
plugins = append(plugins, p)
}
continue
}
// Otherwise treat as workspace directory: user/{workspace}/{name}/
subEntries, err := os.ReadDir(entryPath)
if err != nil {
continue
}
for _, sub := range subEntries {
if !sub.IsDir() {
continue
}
subPath := filepath.Join(entryPath, sub.Name())
p := l.loadPlugin(subPath, false)
if p != nil {
qualifiedName := entry.Name() + "/" + p.Manifest.Name
if isPluginEnabled(settings, qualifiedName) {
plugins = append(plugins, p)
}
}
}
}
return plugins
}
// LoadAll loads both managed and user plugins.
func (l *Loader) LoadAll() []*Plugin {
managed := l.LoadManaged()
user := l.LoadUser()
return append(managed, user...)
}
// scanDir reads a directory of plugin subdirectories and loads each one.
func (l *Loader) scanDir(dir string, isManaged bool) []*Plugin {
entries, err := os.ReadDir(dir)
if err != nil {
if !os.IsNotExist(err) {
slog.Debug("plugin: cannot read dir", "path", dir, "error", err)
}
return nil
}
var plugins []*Plugin
for _, entry := range entries {
if !entry.IsDir() {
continue
}
pluginDir := filepath.Join(dir, entry.Name())
p := l.loadPlugin(pluginDir, isManaged)
if p != nil {
plugins = append(plugins, p)
}
}
return plugins
}
// loadPlugin reads and validates a single plugin directory.
func (l *Loader) loadPlugin(dir string, isManaged bool) *Plugin {
manifestPath := filepath.Join(dir, "plugin.json")
manifest, err := ParseManifest(manifestPath)
if err != nil {
slog.Warn("plugin: failed to parse manifest",
"path", manifestPath, "error", err)
return nil
}
if err := manifest.Validate(l.CLIVersion); err != nil {
slog.Warn("plugin: validation failed",
"plugin", manifest.Name, "error", err)
return nil
}
return &Plugin{
Manifest: *manifest,
Root: dir,
IsManaged: isManaged,
}
}
// settingsPath returns the path to settings.json.
// Uses PluginsDir's parent (~/.dws/) for production, PluginsDir itself for tests.
func (l *Loader) settingsPath() string {
// If PluginsDir ends with "plugins", go up one level to ~/.dws/
if filepath.Base(l.PluginsDir) == "plugins" {
return filepath.Join(filepath.Dir(l.PluginsDir), "settings.json")
}
// For test temp dirs, use PluginsDir directly
return filepath.Join(l.PluginsDir, "settings.json")
}
// loadSettings reads settings.json from the parent of PluginsDir.
func (l *Loader) loadSettings() *Settings {
settingsPath := l.settingsPath()
data, err := os.ReadFile(settingsPath)
if err != nil {
return &Settings{}
}
var s Settings
if err := json.Unmarshal(data, &s); err != nil {
slog.Debug("plugin: failed to parse settings.json", "error", err)
return &Settings{}
}
return &s
}
func isPluginEnabled(s *Settings, name string) bool {
if s == nil || s.EnabledPlugins == nil {
return true // default: enabled
}
enabled, exists := s.EnabledPlugins[name]
if !exists {
return true // not in list = enabled
}
return enabled
}
// InstalledPlugins returns the list of all installed plugins with their
// status info. Used by `dws plugin list`.
type PluginInfo struct {
Name string `json:"name"`
Version string `json:"version"`
Type string `json:"type"` // "managed" or "user"
Enabled bool `json:"enabled"`
Path string `json:"path"`
Description string `json:"description,omitempty"`
}
// ListInstalled returns info about all installed plugins.
func (l *Loader) ListInstalled() []PluginInfo {
var result []PluginInfo
settings := l.loadSettings()
// Managed plugins
managedDir := filepath.Join(l.PluginsDir, "managed")
if entries, err := os.ReadDir(managedDir); err == nil {
for _, entry := range entries {
if !entry.IsDir() {
continue
}
dir := filepath.Join(managedDir, entry.Name())
m, err := ParseManifest(filepath.Join(dir, "plugin.json"))
if err != nil {
continue
}
result = append(result, PluginInfo{
Name: m.Name,
Version: m.Version,
Type: "managed",
Enabled: true, // managed plugins always enabled
Path: dir,
Description: m.Description,
})
}
}
// User plugins
userDir := filepath.Join(l.PluginsDir, "user")
if entries, err := os.ReadDir(userDir); err == nil {
for _, entry := range entries {
if !entry.IsDir() {
continue
}
l.collectUserPluginInfos(filepath.Join(userDir, entry.Name()), entry.Name(), settings, &result)
}
}
// Dev plugins
for name, dir := range settings.DevPlugins {
m, err := ParseManifest(filepath.Join(dir, "plugin.json"))
if err != nil {
continue
}
result = append(result, PluginInfo{
Name: name,
Version: m.Version,
Type: "dev",
Enabled: true,
Path: dir,
Description: m.Description,
})
}
return result
}
func (l *Loader) collectUserPluginInfos(dir, prefix string, settings *Settings, result *[]PluginInfo) {
// Direct plugin
if m, err := ParseManifest(filepath.Join(dir, "plugin.json")); err == nil {
qualName := prefix
*result = append(*result, PluginInfo{
Name: qualName,
Version: m.Version,
Type: "user",
Enabled: isPluginEnabled(settings, qualName),
Path: dir,
Description: m.Description,
})
return
}
// Workspace: dir is a workspace, iterate sub-plugins
subEntries, err := os.ReadDir(dir)
if err != nil {
return
}
for _, sub := range subEntries {
if !sub.IsDir() {
continue
}
subDir := filepath.Join(dir, sub.Name())
m, err := ParseManifest(filepath.Join(subDir, "plugin.json"))
if err != nil {
continue
}
qualName := prefix + "/" + m.Name
*result = append(*result, PluginInfo{
Name: qualName,
Version: m.Version,
Type: "user",
Enabled: isPluginEnabled(settings, qualName),
Path: subDir,
Description: m.Description,
})
}
}
// InstallFromDir copies a plugin from a source directory to the user
// plugins directory.
func (l *Loader) InstallFromDir(srcDir string) (*Plugin, error) {
manifestPath := filepath.Join(srcDir, "plugin.json")
manifest, err := ParseManifest(manifestPath)
if err != nil {
return nil, fmt.Errorf("invalid plugin: %w", err)
}
if err := manifest.Validate(l.CLIVersion); err != nil {
return nil, fmt.Errorf("plugin validation failed: %w", err)
}
destDir := filepath.Join(l.PluginsDir, "user", manifest.Name)
if err := copyDir(srcDir, destDir); err != nil {
return nil, fmt.Errorf("install failed: %w", err)
}
// Remove stale files in destDir that no longer exist in srcDir.
removeStaleFiles(srcDir, destDir)
// Run build if configured (compile server to binary).
if manifest.Build != nil {
if err := runBuild(destDir, manifest.Build); err != nil {
// Clean up on build failure.
_ = os.RemoveAll(destDir)
return nil, fmt.Errorf("plugin build failed: %w", err)
}
}
// Enable by default in settings
l.setPluginEnabled(manifest.Name, true)
return &Plugin{
Manifest: *manifest,
Root: destDir,
IsManaged: false,
}, nil
}
// InstallFromGit clones a git repository and installs the plugin.
// The workspace is extracted from the git URL (e.g. github.com/{workspace}/{name}).
func (l *Loader) InstallFromGit(gitURL string) (*Plugin, error) {
workspace, repoName, err := parseGitURL(gitURL)
if err != nil {
return nil, fmt.Errorf("invalid git URL: %w", err)
}
// Clone to temp directory.
tmpDir, err := os.MkdirTemp("", "dws-plugin-git-*")
if err != nil {
return nil, fmt.Errorf("create temp dir: %w", err)
}
defer os.RemoveAll(tmpDir)
cloneDir := filepath.Join(tmpDir, repoName)
cmd := exec.Command("git", "clone", "--depth", "1", gitURL, cloneDir)
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
if err := cmd.Run(); err != nil {
return nil, fmt.Errorf("git clone failed: %w", err)
}
// Parse and validate manifest.
manifest, err := ParseManifest(filepath.Join(cloneDir, "plugin.json"))
if err != nil {
return nil, fmt.Errorf("invalid plugin: %w", err)
}
if err := manifest.Validate(l.CLIVersion); err != nil {
return nil, fmt.Errorf("plugin validation failed: %w", err)
}
// All plugins install to the user directory with workspace nesting:
// ~/.dws/plugins/user/{workspace}/{name}/. There is no privileged
// workspace — every plugin is third-party.
destDir := filepath.Join(l.PluginsDir, config.PluginUserDir, workspace, manifest.Name)
// Remove .git directory before copying.
_ = os.RemoveAll(filepath.Join(cloneDir, ".git"))
if err := copyDir(cloneDir, destDir); err != nil {
return nil, fmt.Errorf("install failed: %w", err)
}
// Run build if configured (compile server to binary).
if manifest.Build != nil {
if err := runBuild(destDir, manifest.Build); err != nil {
// Clean up on build failure.
_ = os.RemoveAll(destDir)
return nil, fmt.Errorf("plugin build failed: %w", err)
}
}
qualifiedName := workspace + "/" + manifest.Name
l.setPluginEnabled(qualifiedName, true)
return &Plugin{
Manifest: *manifest,
Root: destDir,
IsManaged: false,
}, nil
}
// parseGitURL extracts workspace and repo name from a git URL.
// Supports: https://github.com/org/repo.git, git@github.com:org/repo.git
// Rejects file:// and other local protocols to prevent reading local files.
func parseGitURL(gitURL string) (workspace, repoName string, err error) {
gitURL = strings.TrimSpace(gitURL)
// Reject dangerous protocols that could read local files.
lower := strings.ToLower(gitURL)
if strings.HasPrefix(lower, "file://") || strings.HasPrefix(lower, "/") || strings.HasPrefix(lower, ".") {
return "", "", fmt.Errorf("local paths and file:// URLs are not allowed: %q", gitURL)
}
// Handle SSH format: git@github.com:org/repo.git
if strings.HasPrefix(gitURL, "git@") {
parts := strings.SplitN(gitURL, ":", 2)
if len(parts) != 2 {
return "", "", fmt.Errorf("cannot parse SSH URL %q", gitURL)
}
path := strings.TrimSuffix(parts[1], ".git")
segments := strings.Split(path, "/")
if len(segments) < 2 {
return "", "", fmt.Errorf("SSH URL %q must have org/repo format", gitURL)
}
return segments[len(segments)-2], segments[len(segments)-1], nil
}
// Handle HTTPS format.
u, err := url.Parse(gitURL)
if err != nil {
return "", "", fmt.Errorf("cannot parse URL %q: %w", gitURL, err)
}
// Only allow https:// and http:// schemes.
if u.Scheme != "https" && u.Scheme != "http" {
return "", "", fmt.Errorf("unsupported URL scheme %q: only https and ssh are allowed", u.Scheme)
}
path := strings.TrimSuffix(strings.Trim(u.Path, "/"), ".git")
segments := strings.Split(path, "/")
if len(segments) < 2 {
return "", "", fmt.Errorf("URL %q must have org/repo format", gitURL)
}
return segments[len(segments)-2], segments[len(segments)-1], nil
}
// RemovePlugin removes an installed plugin by name. It searches both the
// user and the legacy managed directories; all plugins are equally
// removable.
func (l *Loader) RemovePlugin(name string, keepData bool) error {
pluginDir := l.findUserPluginDir(name)
if pluginDir == "" {
// Fall back to legacy managed/ directory (for plugins installed
// by older CLI builds that wrote under ~/.dws/plugins/managed/).
legacyDir := filepath.Join(l.PluginsDir, config.PluginManagedDir, name)
if _, err := os.Stat(filepath.Join(legacyDir, "plugin.json")); err == nil {
pluginDir = legacyDir
}
}
if pluginDir == "" {
return fmt.Errorf("plugin %q not found", name)
}
if err := os.RemoveAll(pluginDir); err != nil {
return fmt.Errorf("failed to remove plugin: %w", err)
}
if !keepData {
dataDir := filepath.Join(l.PluginsDir, config.PluginDataDir, name)
_ = os.RemoveAll(dataDir)
}
l.purgePluginFromSettings(name)
return nil
}
// purgePluginFromSettings removes all traces of a plugin from settings.json:
// its enabled flag and any persisted pluginConfigs entry. Called after
// RemovePlugin succeeds so settings.json does not retain dangling state for
// plugins that no longer exist on disk.
func (l *Loader) purgePluginFromSettings(name string) {
settings := l.loadSettings()
changed := false
if _, ok := settings.EnabledPlugins[name]; ok {
delete(settings.EnabledPlugins, name)
changed = true
}
if _, ok := settings.PluginConfigs[name]; ok {
delete(settings.PluginConfigs, name)
changed = true
}
if !changed {
return
}
l.saveSettings(settings)
}
// SetEnabled enables or disables a plugin in settings.json.
func (l *Loader) SetEnabled(name string, enabled bool) error {
// Verify plugin exists
if l.findUserPluginDir(name) == "" {
managedDir := filepath.Join(l.PluginsDir, "managed", name)
if _, err := os.Stat(managedDir); err != nil {
return fmt.Errorf("plugin %q not found", name)
}
}
l.setPluginEnabled(name, enabled)
return nil
}
func (l *Loader) findUserPluginDir(name string) string {
// Try direct: user/{name}/
dir := filepath.Join(l.PluginsDir, "user", name)
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err == nil {
return dir
}
// Try workspace: user/{workspace}/{plugin}/
parts := strings.SplitN(name, "/", 2)
if len(parts) == 2 {
dir = filepath.Join(l.PluginsDir, "user", parts[0], parts[1])
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err == nil {
return dir
}
}
return ""
}
func (l *Loader) setPluginEnabled(name string, enabled bool) {
settings := l.loadSettings()
if settings.EnabledPlugins == nil {
settings.EnabledPlugins = make(map[string]bool)
}
settings.EnabledPlugins[name] = enabled
l.saveSettings(settings)
}
func (l *Loader) saveSettings(s *Settings) {
settingsPath := l.settingsPath()
data, err := json.MarshalIndent(s, "", " ")
if err != nil {
slog.Debug("plugin: failed to marshal settings", "error", err)
return
}
_ = os.MkdirAll(filepath.Dir(settingsPath), 0o700)
_ = os.WriteFile(settingsPath, data, 0o600)
}
// GetPluginConfig returns the value of a config key for a plugin.
// It checks pluginConfigs in settings.json first, then falls back to
// the userConfig default in the plugin's manifest.
func (l *Loader) GetPluginConfig(pluginName, key string) (string, bool) {
settings := l.loadSettings()
if settings.PluginConfigs != nil {
if pluginCfg, ok := settings.PluginConfigs[pluginName]; ok {
if val, ok := pluginCfg[key]; ok {
if s, ok := val.(string); ok {
return s, true
}
}
}
}
return "", false
}
// SetPluginConfig persists a config key-value pair for a plugin.
func (l *Loader) SetPluginConfig(pluginName, key, value string) {
settings := l.loadSettings()
if settings.PluginConfigs == nil {
settings.PluginConfigs = make(map[string]map[string]any)
}
if settings.PluginConfigs[pluginName] == nil {
settings.PluginConfigs[pluginName] = make(map[string]any)
}
settings.PluginConfigs[pluginName][key] = value
l.saveSettings(settings)
}
// UnsetPluginConfig removes a config key for a plugin.
func (l *Loader) UnsetPluginConfig(pluginName, key string) bool {
settings := l.loadSettings()
if settings.PluginConfigs == nil {
return false
}
pluginCfg, ok := settings.PluginConfigs[pluginName]
if !ok {
return false
}
if _, exists := pluginCfg[key]; !exists {
return false
}
delete(pluginCfg, key)
if len(pluginCfg) == 0 {
delete(settings.PluginConfigs, pluginName)
}
l.saveSettings(settings)
return true
}
// ListPluginConfig returns all config key-value pairs for a plugin.
func (l *Loader) ListPluginConfig(pluginName string) map[string]string {
settings := l.loadSettings()
result := make(map[string]string)
if settings.PluginConfigs == nil {
return result
}
pluginCfg, ok := settings.PluginConfigs[pluginName]
if !ok {
return result
}
for k, v := range pluginCfg {
if s, ok := v.(string); ok {
result[k] = s
}
}
return result
}
// InjectPluginConfigEnv reads pluginConfigs from settings.json and sets
// environment variables for each configured key. This allows
// expandPluginVars (which calls os.Expand) to resolve ${KEY} references
// in plugin.json headers, endpoints, etc.
//
// Environment variables already set by the user take precedence — only
// keys not already present in the environment are injected.
// dangerousEnvVars contains environment variable names that must never be
// set from plugin config because they can alter process behavior in
// security-critical ways (library injection, executable search path, etc.).
var dangerousEnvVars = map[string]bool{
"PATH": true, "HOME": true, "USER": true, "SHELL": true,
"LD_PRELOAD": true, "LD_LIBRARY_PATH": true,
"DYLD_INSERT_LIBRARIES": true, "DYLD_LIBRARY_PATH": true, "DYLD_FRAMEWORK_PATH": true,
"NODE_OPTIONS": true, "PYTHONPATH": true, "RUBYLIB": true,
"GOPATH": true, "GOROOT": true,
"HTTP_PROXY": true, "HTTPS_PROXY": true, "ALL_PROXY": true, "NO_PROXY": true,
"http_proxy": true, "https_proxy": true, "all_proxy": true, "no_proxy": true,
}
func (l *Loader) InjectPluginConfigEnv() {
settings := l.loadSettings()
if len(settings.PluginConfigs) == 0 {
return
}
for _, pluginCfg := range settings.PluginConfigs {
for key, val := range pluginCfg {
strVal, ok := val.(string)
if !ok || strVal == "" {
continue
}
// Block dangerous environment variable names.
if dangerousEnvVars[key] {
slog.Warn("plugin: blocked dangerous env var from config",
"key", key)
continue
}
// Do not override existing environment variables.
if _, exists := os.LookupEnv(key); exists {
continue
}
_ = os.Setenv(key, strVal)
}
}
}
// LoadDev loads dev plugins registered via `dws plugin dev`.
// Dev plugins are loaded from their source directories without copying.
func (l *Loader) LoadDev() []*Plugin {
settings := l.loadSettings()
if len(settings.DevPlugins) == 0 {
return nil
}
var plugins []*Plugin
for name, dir := range settings.DevPlugins {
if _, err := os.Stat(filepath.Join(dir, "plugin.json")); err != nil {
slog.Debug("plugin: dev plugin directory missing, skipping",
"name", name, "dir", dir)
continue
}
p := l.loadPlugin(dir, false)
if p != nil {
plugins = append(plugins, p)
slog.Debug("plugin: loaded dev plugin", "name", name, "dir", dir)
}
}
return plugins
}
// RegisterDevPlugin registers a source directory as a dev plugin.
func (l *Loader) RegisterDevPlugin(name, absDir string) error {
settings := l.loadSettings()
if settings.DevPlugins == nil {
settings.DevPlugins = make(map[string]string)
}
settings.DevPlugins[name] = absDir
l.saveSettings(settings)
return nil
}
// UnregisterDevPlugin removes a dev plugin registration.
func (l *Loader) UnregisterDevPlugin(name string) error {
settings := l.loadSettings()
if settings.DevPlugins == nil || settings.DevPlugins[name] == "" {
return fmt.Errorf("dev plugin %q is not registered", name)
}
delete(settings.DevPlugins, name)
l.saveSettings(settings)
return nil
}
// SyncSkills copies plugin SKILL.md files into all detected agent
// skill directories (e.g. ~/.claude/skills/dws/, ~/.cursor/skills/dws/).
// This makes plugin skills available to AI agents without CLI releases.
func SyncSkills(plugins []*Plugin) {
if len(plugins) == 0 {
return
}
homeDir, err := os.UserHomeDir()
if err != nil {
slog.Debug("plugin: cannot get home dir for skill sync", "error", err)
return
}
// Known agent skill directories (subset of upgrade/paths.go knownSkillDirs).
agentDirs := []string{
".agents/skills",
".claude/skills",
".cursor/skills",
".qoder/skills",
".codex/skills",
}
for _, p := range plugins {
skillsDir := p.SkillsDir()
if _, err := os.Stat(skillsDir); err != nil {
continue
}
// Walk the plugin's skills directory and copy files to each agent dir.
entries, err := os.ReadDir(skillsDir)
if err != nil {
continue
}
for _, agentDir := range agentDirs {
agentBase := filepath.Join(homeDir, agentDir)
// Only sync to agents that are actually installed (parent dir exists).
parentGate := filepath.Dir(agentBase)
if _, err := os.Stat(parentGate); os.IsNotExist(err) {
continue
}
for _, entry := range entries {
src := filepath.Join(skillsDir, entry.Name())
// Place plugin skills under dws/plugins/{plugin-name}/
dest := filepath.Join(agentBase, "dws", "plugins", p.Manifest.Name, entry.Name())
if entry.IsDir() {
_ = copyDir(src, dest)
} else {
_ = os.MkdirAll(filepath.Dir(dest), 0o755)
data, readErr := os.ReadFile(src)
if readErr == nil {
_ = os.WriteFile(dest, data, 0o644)
}
}
}
}
}
slog.Debug("plugin: skill sync completed", "plugins", len(plugins))
}
// BuildPlugin runs the build command declared in plugin.json.
// It compiles the plugin's stdio server into a native binary so that
// users don't need language runtimes. Returns nil if no build is configured.
func BuildPlugin(pluginDir string) error {
manifest, err := ParseManifest(filepath.Join(pluginDir, "plugin.json"))
if err != nil {
return fmt.Errorf("parse manifest: %w", err)
}
if manifest.Build == nil {
return nil // no build configured
}
return runBuild(pluginDir, manifest.Build)
}
// runBuild executes the build command and verifies the output exists.
func runBuild(pluginDir string, build *BuildConfig) error {
if build.Command == "" {
return fmt.Errorf("build.command is empty")
}
// Validate build.output is a relative path within the plugin directory.
if build.Output != "" {
if filepath.IsAbs(build.Output) {
return fmt.Errorf("build.output must be a relative path, got %q", build.Output)
}
cleanOut := filepath.Clean(build.Output)
if strings.HasPrefix(cleanOut, "..") {
return fmt.Errorf("build.output must not escape plugin directory: %q", build.Output)
}
}
slog.Info("plugin: building", "dir", pluginDir, "command", build.Command)
var cmd *exec.Cmd
if runtime.GOOS == "windows" {
cmd = exec.Command("cmd", "/C", build.Command)
} else {
cmd = exec.Command("sh", "-c", build.Command)
}
cmd.Dir = pluginDir
cmd.Stdout = os.Stdout
cmd.Stderr = os.Stderr
// Pass through environment + plugin root
cmd.Env = append(os.Environ(), "DWS_PLUGIN_ROOT="+pluginDir)
if err := cmd.Run(); err != nil {
return fmt.Errorf("build failed: %w", err)
}
// Verify output binary exists
if build.Output != "" {
outPath := filepath.Join(pluginDir, build.Output)
info, err := os.Stat(outPath)
if err != nil {
return fmt.Errorf("build output not found at %s: %w", build.Output, err)
}
// Ensure the output is executable
if info.Mode()&0o111 == 0 {
_ = os.Chmod(outPath, info.Mode()|0o755)
}
}
slog.Info("plugin: build succeeded", "output", build.Output)
return nil
}
// copyDir recursively copies src to dst, skipping files whose content
// is identical to the destination. This avoids overwriting locked
// executables (e.g. a running stdio plugin on Windows).
// Symlinks are skipped for security (prevents path traversal attacks).
func copyDir(src, dst string) error {
cleanDst := filepath.Clean(dst) + string(os.PathSeparator)
if err := os.MkdirAll(dst, 0o755); err != nil {
return err
}
return filepath.Walk(src, func(path string, info os.FileInfo, err error) error {
if err != nil {
return err
}
// Skip symlinks to prevent path traversal.
if info.Mode()&os.ModeSymlink != 0 {
return nil
}
rel, err := filepath.Rel(src, path)
if err != nil {
return err
}
target := filepath.Join(dst, rel)
// Guard against path traversal via crafted relative paths.
if target != cleanDst[:len(cleanDst)-1] && !strings.HasPrefix(target, cleanDst) {
return fmt.Errorf("path traversal detected: %s", rel)
}
if info.IsDir() {
return os.MkdirAll(target, info.Mode())
}
data, err := os.ReadFile(path)
if err != nil {
return err
}
// Skip if destination already has identical content (cheap size check first).
if targetInfo, statErr := os.Stat(target); statErr == nil && targetInfo.Size() == int64(len(data)) {
if existing, readErr := os.ReadFile(target); readErr == nil && bytes.Equal(existing, data) {
return nil
}
}
return os.WriteFile(target, data, info.Mode())
})
}
// removeStaleFiles deletes files under dst that do not exist in src.
// Best-effort: errors are logged but do not fail the install.
func removeStaleFiles(src, dst string) {
srcSet := make(map[string]struct{})
_ = filepath.Walk(src, func(path string, info os.FileInfo, err error) error {
if err != nil {
return nil
}
rel, relErr := filepath.Rel(src, path)
if relErr != nil {
return nil
}
srcSet[rel] = struct{}{}
return nil
})
_ = filepath.Walk(dst, func(path string, info os.FileInfo, err error) error {
if err != nil {
return nil
}
rel, relErr := filepath.Rel(dst, path)
if relErr != nil {
return nil
}
if rel == "." {
return nil
}
if _, exists := srcSet[rel]; !exists {
if info.IsDir() {
_ = os.RemoveAll(path)
return filepath.SkipDir
}
if removeErr := os.Remove(path); removeErr != nil {
slog.Debug("plugin: failed to remove stale file", "path", path, "error", removeErr)
}
}
return nil
})
}
+193
View File
@@ -0,0 +1,193 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package plugin
import (
"os"
"testing"
)
func TestSetAndGetPluginConfig(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Initially empty.
val, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
if ok {
t.Errorf("expected not found, got %q", val)
}
// Set a value.
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
// Read it back.
val, ok = loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
if !ok {
t.Fatal("expected to find config after set")
}
if val != "sk-test-12345" {
t.Errorf("got %q, want sk-test-12345", val)
}
}
func TestSetPluginConfigMultipleKeys(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
loader.SetPluginConfig("my-plugin", "API_KEY", "key-1")
loader.SetPluginConfig("my-plugin", "API_ENDPOINT", "https://example.com")
loader.SetPluginConfig("other-plugin", "TOKEN", "tok-abc")
val, ok := loader.GetPluginConfig("my-plugin", "API_KEY")
if !ok || val != "key-1" {
t.Errorf("API_KEY = %q (ok=%v), want key-1", val, ok)
}
val, ok = loader.GetPluginConfig("my-plugin", "API_ENDPOINT")
if !ok || val != "https://example.com" {
t.Errorf("API_ENDPOINT = %q (ok=%v), want https://example.com", val, ok)
}
val, ok = loader.GetPluginConfig("other-plugin", "TOKEN")
if !ok || val != "tok-abc" {
t.Errorf("TOKEN = %q (ok=%v), want tok-abc", val, ok)
}
}
func TestUnsetPluginConfig(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Unset on empty returns false.
if loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
t.Error("expected false for unset on empty config")
}
// Set then unset.
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
if !loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
t.Error("expected true for unset of existing key")
}
// Verify it's gone.
_, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
if ok {
t.Error("expected not found after unset")
}
}
func TestUnsetPluginConfigCleansEmptyMap(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
loader.SetPluginConfig("demo-devtool", "KEY1", "val1")
loader.UnsetPluginConfig("demo-devtool", "KEY1")
// After removing the last key, the plugin entry should be cleaned up.
configs := loader.ListPluginConfig("demo-devtool")
if len(configs) != 0 {
t.Errorf("expected empty config map after removing last key, got %v", configs)
}
}
func TestListPluginConfig(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Empty list.
configs := loader.ListPluginConfig("demo-devtool")
if len(configs) != 0 {
t.Errorf("expected empty, got %v", configs)
}
// Set some values.
loader.SetPluginConfig("demo-devtool", "KEY_A", "val-a")
loader.SetPluginConfig("demo-devtool", "KEY_B", "val-b")
configs = loader.ListPluginConfig("demo-devtool")
if len(configs) != 2 {
t.Fatalf("expected 2 configs, got %d", len(configs))
}
if configs["KEY_A"] != "val-a" {
t.Errorf("KEY_A = %q, want val-a", configs["KEY_A"])
}
if configs["KEY_B"] != "val-b" {
t.Errorf("KEY_B = %q, want val-b", configs["KEY_B"])
}
}
func TestInjectPluginConfigEnv(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Use a unique env var name to avoid test pollution.
envKey := "DWS_TEST_INJECT_CONFIG_" + t.Name()
t.Cleanup(func() { os.Unsetenv(envKey) })
loader.SetPluginConfig("demo-devtool", envKey, "injected-value")
// Ensure it's not already set.
os.Unsetenv(envKey)
loader.InjectPluginConfigEnv()
got := os.Getenv(envKey)
if got != "injected-value" {
t.Errorf("env %s = %q, want injected-value", envKey, got)
}
}
func TestInjectPluginConfigEnvDoesNotOverride(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
envKey := "DWS_TEST_INJECT_NOOVERRIDE_" + t.Name()
t.Cleanup(func() { os.Unsetenv(envKey) })
// Pre-set the env var.
os.Setenv(envKey, "user-value")
loader.SetPluginConfig("demo-devtool", envKey, "config-value")
loader.InjectPluginConfigEnv()
got := os.Getenv(envKey)
if got != "user-value" {
t.Errorf("env %s = %q, want user-value (should not be overridden)", envKey, got)
}
}
func TestSetPluginConfigOverwritesExisting(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
loader.SetPluginConfig("demo-devtool", "KEY", "old-value")
loader.SetPluginConfig("demo-devtool", "KEY", "new-value")
val, ok := loader.GetPluginConfig("demo-devtool", "KEY")
if !ok || val != "new-value" {
t.Errorf("got %q (ok=%v), want new-value", val, ok)
}
}
func TestGetPluginConfigWrongPlugin(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
loader.SetPluginConfig("plugin-a", "KEY", "value")
_, ok := loader.GetPluginConfig("plugin-b", "KEY")
if ok {
t.Error("expected not found for different plugin name")
}
}
+258
View File
@@ -0,0 +1,258 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
// Package plugin implements the DWS CLI plugin system. It loads,
// validates, and injects plugin capabilities (MCP servers, skills,
// pipeline hooks) into the existing CLI infrastructure.
package plugin
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
)
// namePattern validates plugin names: lowercase kebab-case, 3–50 chars.
var namePattern = regexp.MustCompile(`^[a-z][a-z0-9-]{2,49}$`)
// Manifest represents the parsed contents of a plugin.json file.
type Manifest struct {
Name string `json:"name"`
Version string `json:"version"`
Description string `json:"description,omitempty"`
Type string `json:"type,omitempty"` // "managed" or "user"
MinCLIVersion string `json:"minCLIVersion,omitempty"`
MCPServers map[string]*MCPServer `json:"mcpServers,omitempty"`
Skills string `json:"skills,omitempty"`
Hooks string `json:"hooks,omitempty"`
Permissions []string `json:"permissions,omitempty"`
UserConfig map[string]ConfigItem `json:"userConfig,omitempty"`
Build *BuildConfig `json:"build,omitempty"`
}
// BuildConfig declares how to compile the plugin's stdio server into
// a native binary. DWS runs this automatically during install so that
// plugin users never need language runtimes or dependency managers.
type BuildConfig struct {
// Command is the shell command to compile the server.
// Executed via "sh -c" in the plugin root directory.
// Examples: "bun build --compile src/server.ts --outfile bin/server"
// "go build -o bin/server ./cmd/server"
// "pip install pyinstaller && pyinstaller --onefile src/server.py -n server --distpath bin/"
Command string `json:"command"`
// Output is the path to the compiled binary, relative to the plugin root.
// Used to verify the build succeeded. Example: "bin/server"
Output string `json:"output"`
}
// MCPServer describes a single MCP server declared by a plugin.
type MCPServer struct {
Type string `json:"type"` // "streamable-http" or "stdio"
Endpoint string `json:"endpoint,omitempty"` // required for streamable-http
Command string `json:"command,omitempty"` // required for stdio
Args []string `json:"args,omitempty"`
Env map[string]string `json:"env,omitempty"`
Headers map[string]string `json:"headers,omitempty"` // custom HTTP headers (e.g. Authorization for third-party APIs)
CLI json.RawMessage `json:"cli,omitempty"` // CLIOverlay, passed through
}
// ConfigItem describes a user-configurable setting for a plugin.
type ConfigItem struct {
Description string `json:"description,omitempty"`
Default string `json:"default,omitempty"`
Sensitive bool `json:"sensitive,omitempty"`
}
// HooksConfig describes pipeline hooks declared in a hooks.json file.
type HooksConfig struct {
Hooks []HookEntry `json:"hooks"`
}
// HookEntry describes a single pipeline hook.
type HookEntry struct {
Phase string `json:"phase"` // "pre-request", "post-response", etc.
Matcher string `json:"matcher,omitempty"` // glob pattern, e.g. "conference.*"
Command string `json:"command"` // shell command to execute
Timeout int `json:"timeout,omitempty"` // seconds, default 30
}
// Plugin is a loaded, validated plugin ready for injection.
type Plugin struct {
Manifest Manifest
Root string // absolute path to plugin directory
IsManaged bool // true for official (DingTalk-Real-AI) plugins
}
// ParseManifest reads and parses a plugin.json file.
func ParseManifest(path string) (*Manifest, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read plugin.json: %w", err)
}
var m Manifest
if err := json.Unmarshal(data, &m); err != nil {
return nil, fmt.Errorf("parse plugin.json: %w", err)
}
return &m, nil
}
// Validate checks that a manifest is well-formed. It returns an error
// describing the first problem found, or nil if the manifest is valid.
// cliVersion is the current CLI version string for compatibility checks.
func (m *Manifest) Validate(cliVersion string) error {
if !namePattern.MatchString(m.Name) {
return fmt.Errorf("invalid plugin name %q: must be lowercase kebab-case, 3–50 chars", m.Name)
}
if !isValidSemver(m.Version) {
return fmt.Errorf("invalid plugin version %q: must be valid semver (e.g. 1.0.0)", m.Version)
}
if m.Type != "" && m.Type != "managed" && m.Type != "user" {
return fmt.Errorf("invalid plugin type %q: must be \"managed\" or \"user\"", m.Type)
}
if m.MinCLIVersion != "" && cliVersion != "" && cliVersion != "dev" {
if compareSemver(cliVersion, m.MinCLIVersion) < 0 {
return fmt.Errorf("plugin requires CLI >= %s, current is %s", m.MinCLIVersion, cliVersion)
}
}
for key, srv := range m.MCPServers {
if err := validateMCPServer(key, srv); err != nil {
return err
}
}
if m.Skills != "" {
if err := validateSafePath(m.Skills); err != nil {
return fmt.Errorf("skills path: %w", err)
}
}
if m.Hooks != "" {
if err := validateSafePath(m.Hooks); err != nil {
return fmt.Errorf("hooks path: %w", err)
}
}
return nil
}
func validateMCPServer(key string, srv *MCPServer) error {
switch srv.Type {
case "streamable-http":
if strings.TrimSpace(srv.Endpoint) == "" {
return fmt.Errorf("mcpServers[%q]: streamable-http requires endpoint", key)
}
case "stdio":
if strings.TrimSpace(srv.Command) == "" {
return fmt.Errorf("mcpServers[%q]: stdio requires command", key)
}
// Reject absolute paths in command to encourage relative paths within plugin root.
if filepath.IsAbs(srv.Command) {
return fmt.Errorf("mcpServers[%q]: command must be a relative path, got %q", key, srv.Command)
}
default:
return fmt.Errorf("mcpServers[%q]: unsupported type %q (must be streamable-http or stdio)", key, srv.Type)
}
return nil
}
// validateSafePath rejects paths containing ".." traversal.
func validateSafePath(p string) error {
cleaned := filepath.Clean(p)
if strings.Contains(cleaned, "..") {
return fmt.Errorf("unsafe path %q: must not contain \"..\"", p)
}
return nil
}
// LoadHooks reads the hooks.json file referenced by the manifest.
func (p *Plugin) LoadHooks() (*HooksConfig, error) {
if p.Manifest.Hooks == "" {
return nil, nil
}
hooksPath := filepath.Join(p.Root, p.Manifest.Hooks)
data, err := os.ReadFile(hooksPath)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("read hooks: %w", err)
}
var cfg HooksConfig
if err := json.Unmarshal(data, &cfg); err != nil {
return nil, fmt.Errorf("parse hooks: %w", err)
}
return &cfg, nil
}
// SkillsDir returns the absolute path to the plugin's skills directory.
func (p *Plugin) SkillsDir() string {
dir := p.Manifest.Skills
if dir == "" {
dir = "./skills/"
}
return filepath.Join(p.Root, dir)
}
// isValidSemver checks if a string is a valid semantic version (major.minor.patch).
func isValidSemver(v string) bool {
parts := strings.SplitN(strings.TrimPrefix(v, "v"), "-", 2)
nums := strings.Split(parts[0], ".")
if len(nums) != 3 {
return false
}
for _, n := range nums {
if _, err := strconv.Atoi(n); err != nil {
return false
}
}
return true
}
// parseSemver extracts major, minor, patch from a version string.
func parseSemver(v string) (int, int, int) {
v = strings.TrimPrefix(v, "v")
parts := strings.SplitN(v, "-", 2) // strip pre-release
nums := strings.Split(parts[0], ".")
if len(nums) != 3 {
return 0, 0, 0
}
major, _ := strconv.Atoi(nums[0])
minor, _ := strconv.Atoi(nums[1])
patch, _ := strconv.Atoi(nums[2])
return major, minor, patch
}
// compareSemver compares two semver strings. Returns -1, 0, or 1.
func compareSemver(a, b string) int {
aMaj, aMin, aPat := parseSemver(a)
bMaj, bMin, bPat := parseSemver(b)
if aMaj != bMaj {
return cmpInt(aMaj, bMaj)
}
if aMin != bMin {
return cmpInt(aMin, bMin)
}
return cmpInt(aPat, bPat)
}
func cmpInt(a, b int) int {
if a < b {
return -1
}
if a > b {
return 1
}
return 0
}
+668
View File
@@ -0,0 +1,668 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package plugin
import (
"encoding/json"
"os"
"path/filepath"
"testing"
)
func TestParseManifest(t *testing.T) {
dir := t.TempDir()
manifestPath := filepath.Join(dir, "plugin.json")
content := `{
"name": "conference",
"version": "1.0.0",
"description": "音视频会议",
"type": "managed",
"minCLIVersion": "0.9.0",
"mcpServers": {
"conference": {
"type": "streamable-http",
"endpoint": "https://mcp.conference.dingtalk.com"
},
"conference-local": {
"type": "stdio",
"command": "${DWS_PLUGIN_ROOT}/bin/conference-local",
"args": ["--mode", "cli"]
}
},
"skills": "./skills/"
}`
if err := os.WriteFile(manifestPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
m, err := ParseManifest(manifestPath)
if err != nil {
t.Fatalf("ParseManifest: %v", err)
}
if m.Name != "conference" {
t.Errorf("name = %q, want conference", m.Name)
}
if m.Version != "1.0.0" {
t.Errorf("version = %q, want 1.0.0", m.Version)
}
if m.Type != "managed" {
t.Errorf("type = %q, want managed", m.Type)
}
if len(m.MCPServers) != 2 {
t.Errorf("mcpServers count = %d, want 2", len(m.MCPServers))
}
if m.MCPServers["conference"].Type != "streamable-http" {
t.Errorf("conference server type = %q, want streamable-http", m.MCPServers["conference"].Type)
}
if m.MCPServers["conference-local"].Type != "stdio" {
t.Errorf("conference-local server type = %q, want stdio", m.MCPServers["conference-local"].Type)
}
}
func TestManifestValidate(t *testing.T) {
tests := []struct {
name string
manifest Manifest
cliVersion string
wantErr bool
}{
{
name: "valid manifest",
manifest: Manifest{
Name: "conference",
Version: "1.0.0",
Type: "managed",
MCPServers: map[string]*MCPServer{
"conference": {Type: "streamable-http", Endpoint: "https://example.com"},
},
},
cliVersion: "1.0.0",
wantErr: false,
},
{
name: "invalid name - too short",
manifest: Manifest{
Name: "ab",
Version: "1.0.0",
},
wantErr: true,
},
{
name: "invalid name - uppercase",
manifest: Manifest{
Name: "MyPlugin",
Version: "1.0.0",
},
wantErr: true,
},
{
name: "invalid version",
manifest: Manifest{
Name: "my-plugin",
Version: "not-semver",
},
wantErr: true,
},
{
name: "invalid type",
manifest: Manifest{
Name: "my-plugin",
Version: "1.0.0",
Type: "invalid",
},
wantErr: true,
},
{
name: "cli version too low",
manifest: Manifest{
Name: "my-plugin",
Version: "1.0.0",
MinCLIVersion: "2.0.0",
},
cliVersion: "1.0.0",
wantErr: true,
},
{
name: "streamable-http without endpoint",
manifest: Manifest{
Name: "my-plugin",
Version: "1.0.0",
MCPServers: map[string]*MCPServer{
"srv": {Type: "streamable-http"},
},
},
wantErr: true,
},
{
name: "stdio without command",
manifest: Manifest{
Name: "my-plugin",
Version: "1.0.0",
MCPServers: map[string]*MCPServer{
"srv": {Type: "stdio"},
},
},
wantErr: true,
},
{
name: "unsafe skills path",
manifest: Manifest{
Name: "my-plugin",
Version: "1.0.0",
Skills: "../../../etc/passwd",
},
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := tt.manifest.Validate(tt.cliVersion)
if (err != nil) != tt.wantErr {
t.Errorf("Validate() error = %v, wantErr %v", err, tt.wantErr)
}
})
}
}
func TestPluginToServerDescriptors(t *testing.T) {
cliOverlay, _ := json.Marshal(map[string]any{
"id": "conference",
"command": "conference",
})
p := &Plugin{
Manifest: Manifest{
Name: "conference",
Description: "音视频会议",
MCPServers: map[string]*MCPServer{
"conference": {
Type: "streamable-http",
Endpoint: "https://mcp.conference.dingtalk.com",
CLI: cliOverlay,
},
"conference-local": {
Type: "stdio",
Command: "/usr/local/bin/conference-local",
},
},
},
Root: "/tmp/plugins/conference",
IsManaged: true,
}
descriptors := p.ToServerDescriptors()
// Only streamable-http should be converted
if len(descriptors) != 1 {
t.Fatalf("got %d descriptors, want 1 (stdio should be skipped)", len(descriptors))
}
d := descriptors[0]
if d.Key != "conference" {
t.Errorf("key = %q, want conference", d.Key)
}
if d.Endpoint != "https://mcp.conference.dingtalk.com" {
t.Errorf("endpoint = %q", d.Endpoint)
}
if d.Source != "plugin-managed" {
t.Errorf("source = %q, want plugin-managed", d.Source)
}
if d.CLI.ID != "conference" {
t.Errorf("cli.id = %q, want conference", d.CLI.ID)
}
}
func TestPluginToServerDescriptorsWithHeaders(t *testing.T) {
cliOverlay, _ := json.Marshal(map[string]any{
"id": "web-search",
"command": "web-search",
})
// Set an environment variable to test expansion
t.Setenv("TEST_API_KEY", "sk-test-12345")
p := &Plugin{
Manifest: Manifest{
Name: "my-plugin",
Description: "Test plugin with headers",
MCPServers: map[string]*MCPServer{
"web-search": {
Type: "streamable-http",
Endpoint: "https://api.example.com/mcp/v1",
CLI: cliOverlay,
Headers: map[string]string{
"Authorization": "Bearer ${TEST_API_KEY}",
"X-Custom": "static-value",
},
},
},
},
Root: "/tmp/plugins/my-plugin",
}
descriptors := p.ToServerDescriptors()
if len(descriptors) != 1 {
t.Fatalf("got %d descriptors, want 1", len(descriptors))
}
d := descriptors[0]
if d.Key != "web-search" {
t.Errorf("key = %q, want web-search", d.Key)
}
if len(d.AuthHeaders) != 2 {
t.Fatalf("AuthHeaders len = %d, want 2", len(d.AuthHeaders))
}
// Environment variable should be expanded
if d.AuthHeaders["Authorization"] != "Bearer sk-test-12345" {
t.Errorf("AuthHeaders[Authorization] = %q, want 'Bearer sk-test-12345'", d.AuthHeaders["Authorization"])
}
if d.AuthHeaders["X-Custom"] != "static-value" {
t.Errorf("AuthHeaders[X-Custom] = %q, want static-value", d.AuthHeaders["X-Custom"])
}
if d.Source != "plugin" {
t.Errorf("source = %q, want plugin (non-managed)", d.Source)
}
}
func TestPluginToServerDescriptorsNoHeaders(t *testing.T) {
cliOverlay, _ := json.Marshal(map[string]any{
"id": "conference",
"command": "conference",
})
p := &Plugin{
Manifest: Manifest{
Name: "conference",
MCPServers: map[string]*MCPServer{
"conference": {
Type: "streamable-http",
Endpoint: "https://mcp.conference.dingtalk.com",
CLI: cliOverlay,
},
},
},
Root: "/tmp/plugins/conference",
}
descriptors := p.ToServerDescriptors()
if len(descriptors) != 1 {
t.Fatalf("got %d descriptors, want 1", len(descriptors))
}
if descriptors[0].AuthHeaders != nil {
t.Errorf("AuthHeaders = %v, want nil for server without headers", descriptors[0].AuthHeaders)
}
}
func TestParseManifestWithHeaders(t *testing.T) {
dir := t.TempDir()
manifestPath := filepath.Join(dir, "plugin.json")
content := `{
"name": "api-plugin",
"version": "1.0.0",
"mcpServers": {
"api-server": {
"type": "streamable-http",
"endpoint": "https://api.example.com/mcp",
"headers": {
"Authorization": "Bearer ${MY_API_KEY}",
"X-Custom-Header": "custom-value"
}
}
}
}`
if err := os.WriteFile(manifestPath, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
m, err := ParseManifest(manifestPath)
if err != nil {
t.Fatalf("ParseManifest: %v", err)
}
srv := m.MCPServers["api-server"]
if srv == nil {
t.Fatal("api-server not found in MCPServers")
}
if len(srv.Headers) != 2 {
t.Fatalf("Headers len = %d, want 2", len(srv.Headers))
}
if srv.Headers["Authorization"] != "Bearer ${MY_API_KEY}" {
t.Errorf("Headers[Authorization] = %q, want raw template", srv.Headers["Authorization"])
}
if srv.Headers["X-Custom-Header"] != "custom-value" {
t.Errorf("Headers[X-Custom-Header] = %q, want custom-value", srv.Headers["X-Custom-Header"])
}
}
func TestLoaderScanEmpty(t *testing.T) {
dir := t.TempDir()
loader := &Loader{
PluginsDir: dir,
CLIVersion: "1.0.0",
}
managed := loader.LoadManaged()
if len(managed) != 0 {
t.Errorf("expected 0 managed plugins, got %d", len(managed))
}
user := loader.LoadUser()
if len(user) != 0 {
t.Errorf("expected 0 user plugins, got %d", len(user))
}
}
func TestLoaderLoadManaged(t *testing.T) {
dir := t.TempDir()
managedDir := filepath.Join(dir, "managed", "conference")
if err := os.MkdirAll(managedDir, 0o755); err != nil {
t.Fatal(err)
}
manifest := `{
"name": "conference",
"version": "1.0.0",
"type": "managed",
"mcpServers": {
"conference": {
"type": "streamable-http",
"endpoint": "https://example.com"
}
}
}`
if err := os.WriteFile(filepath.Join(managedDir, "plugin.json"), []byte(manifest), 0o644); err != nil {
t.Fatal(err)
}
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
plugins := loader.LoadManaged()
if len(plugins) != 1 {
t.Fatalf("expected 1 managed plugin, got %d", len(plugins))
}
if plugins[0].Manifest.Name != "conference" {
t.Errorf("name = %q, want conference", plugins[0].Manifest.Name)
}
if !plugins[0].IsManaged {
t.Error("expected IsManaged = true")
}
}
// TestRemoveLegacyManagedPlugin ensures plugins that were installed under
// the legacy ~/.dws/plugins/managed/ directory are now freely removable
// — the old "cannot be removed" privilege has been dropped.
func TestRemoveLegacyManagedPlugin(t *testing.T) {
dir := t.TempDir()
managedDir := filepath.Join(dir, "managed", "conference")
if err := os.MkdirAll(managedDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(managedDir, "plugin.json"), []byte(`{"name":"conference","version":"1.0.0"}`), 0o644); err != nil {
t.Fatal(err)
}
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
if err := loader.RemovePlugin("conference", false); err != nil {
t.Fatalf("unexpected error removing legacy managed plugin: %v", err)
}
if _, err := os.Stat(managedDir); !os.IsNotExist(err) {
t.Errorf("managed plugin dir should be removed, stat err = %v", err)
}
}
// TestRemovePluginPurgesSettings verifies RemovePlugin fully purges the
// plugin's settings — both its enabled flag and any pluginConfigs entry —
// so settings.json does not retain dangling state for a plugin that no
// longer exists on disk. Covers both the user and legacy managed paths.
func TestRemovePluginPurgesSettings(t *testing.T) {
cases := []struct {
name string
layout string // "user" or "legacy"
pkgName string
}{
{name: "user plugin", layout: "user", pkgName: "my-plugin"},
{name: "legacy managed plugin", layout: "legacy", pkgName: "conference"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
dir := t.TempDir()
var pluginDir string
switch tc.layout {
case "user":
pluginDir = filepath.Join(dir, "user", tc.pkgName)
case "legacy":
pluginDir = filepath.Join(dir, "managed", tc.pkgName)
}
if err := os.MkdirAll(pluginDir, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(pluginDir, "plugin.json"),
[]byte(`{"name":"`+tc.pkgName+`","version":"1.0.0"}`), 0o644); err != nil {
t.Fatal(err)
}
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Seed settings.json with an explicit enabled flag and a
// pluginConfigs entry to verify both get purged.
settings := &Settings{
EnabledPlugins: map[string]bool{tc.pkgName: true, "other-plugin": true},
PluginConfigs: map[string]map[string]any{
tc.pkgName: {"API_KEY": "secret"},
"other-plugin": {"TOKEN": "keep-me"},
},
}
loader.saveSettings(settings)
if err := loader.RemovePlugin(tc.pkgName, false); err != nil {
t.Fatalf("RemovePlugin: %v", err)
}
reloaded := loader.loadSettings()
if _, exists := reloaded.EnabledPlugins[tc.pkgName]; exists {
t.Errorf("EnabledPlugins should not retain removed plugin %q", tc.pkgName)
}
if _, exists := reloaded.PluginConfigs[tc.pkgName]; exists {
t.Errorf("PluginConfigs should not retain removed plugin %q", tc.pkgName)
}
if !reloaded.EnabledPlugins["other-plugin"] {
t.Error("unrelated EnabledPlugins entry should be preserved")
}
if reloaded.PluginConfigs["other-plugin"]["TOKEN"] != "keep-me" {
t.Error("unrelated PluginConfigs entry should be preserved")
}
})
}
}
func TestIsPluginEnabled(t *testing.T) {
s := &Settings{
EnabledPlugins: map[string]bool{
"my-plugin": true,
"disabled": false,
},
}
if !isPluginEnabled(s, "my-plugin") {
t.Error("my-plugin should be enabled")
}
if isPluginEnabled(s, "disabled") {
t.Error("disabled should not be enabled")
}
if !isPluginEnabled(s, "not-in-list") {
t.Error("unlisted plugin should default to enabled")
}
if !isPluginEnabled(nil, "anything") {
t.Error("nil settings should default to enabled")
}
}
func TestParseGitURL(t *testing.T) {
tests := []struct {
name string
url string
wantWS string
wantRepo string
wantErr bool
}{
{
name: "https with .git",
url: "https://github.com/PeterGuy326/hello-plugin.git",
wantWS: "PeterGuy326",
wantRepo: "hello-plugin",
},
{
name: "https without .git",
url: "https://github.com/DingTalk-Real-AI/conference",
wantWS: "DingTalk-Real-AI",
wantRepo: "conference",
},
{
name: "ssh format",
url: "git@github.com:DingTalk-Real-AI/conference.git",
wantWS: "DingTalk-Real-AI",
wantRepo: "conference",
},
{
name: "invalid - no repo",
url: "https://github.com/onlyone",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ws, repo, err := parseGitURL(tt.url)
if (err != nil) != tt.wantErr {
t.Errorf("parseGitURL() error = %v, wantErr %v", err, tt.wantErr)
return
}
if !tt.wantErr {
if ws != tt.wantWS {
t.Errorf("workspace = %q, want %q", ws, tt.wantWS)
}
if repo != tt.wantRepo {
t.Errorf("repo = %q, want %q", repo, tt.wantRepo)
}
}
})
}
}
func TestDevPluginRegistration(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
// Create a dev plugin directory
devDir := filepath.Join(t.TempDir(), "my-dev-plugin")
if err := os.MkdirAll(devDir, 0o755); err != nil {
t.Fatal(err)
}
manifest := `{"name":"my-dev-plugin","version":"0.1.0","type":"user"}`
if err := os.WriteFile(filepath.Join(devDir, "plugin.json"), []byte(manifest), 0o644); err != nil {
t.Fatal(err)
}
// Register dev plugin
if err := loader.RegisterDevPlugin("my-dev-plugin", devDir); err != nil {
t.Fatalf("RegisterDevPlugin: %v", err)
}
// Load dev plugins
plugins := loader.LoadDev()
if len(plugins) != 1 {
t.Fatalf("expected 1 dev plugin, got %d", len(plugins))
}
if plugins[0].Manifest.Name != "my-dev-plugin" {
t.Errorf("name = %q, want my-dev-plugin", plugins[0].Manifest.Name)
}
if plugins[0].Root != devDir {
t.Errorf("root = %q, want %q (should load from source dir, not copy)", plugins[0].Root, devDir)
}
// Unregister
if err := loader.UnregisterDevPlugin("my-dev-plugin"); err != nil {
t.Fatalf("UnregisterDevPlugin: %v", err)
}
// Should be empty now
plugins = loader.LoadDev()
if len(plugins) != 0 {
t.Errorf("expected 0 dev plugins after unregister, got %d", len(plugins))
}
}
func TestUnregisterDevPluginNotFound(t *testing.T) {
dir := t.TempDir()
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
err := loader.UnregisterDevPlugin("nonexistent")
if err == nil {
t.Error("expected error when unregistering nonexistent dev plugin")
}
}
func TestSyncSkills(t *testing.T) {
// Create a plugin with skills
pluginDir := t.TempDir()
skillsDir := filepath.Join(pluginDir, "skills", "test-plugin")
if err := os.MkdirAll(skillsDir, 0o755); err != nil {
t.Fatal(err)
}
skillContent := "# Test Plugin Skill"
if err := os.WriteFile(filepath.Join(skillsDir, "SKILL.md"), []byte(skillContent), 0o644); err != nil {
t.Fatal(err)
}
p := &Plugin{
Manifest: Manifest{
Name: "test-plugin",
Skills: "./skills/test-plugin",
},
Root: pluginDir,
}
// Create a mock agent directory
home, _ := os.UserHomeDir()
agentDir := filepath.Join(home, ".agents", "skills")
// Only run if .agents exists (don't create in CI)
if _, err := os.Stat(filepath.Dir(agentDir)); err == nil {
SyncSkills([]*Plugin{p})
synced := filepath.Join(agentDir, "dws", "plugins", "test-plugin", "SKILL.md")
if _, err := os.Stat(synced); err == nil {
data, _ := os.ReadFile(synced)
if string(data) != skillContent {
t.Errorf("synced content = %q, want %q", string(data), skillContent)
}
// Cleanup
os.RemoveAll(filepath.Join(agentDir, "dws", "plugins", "test-plugin"))
}
}
}
func contains(s, substr string) bool {
return len(s) >= len(substr) && (s == substr || len(s) > 0 && containsSubstring(s, substr))
}
func containsSubstring(s, sub string) bool {
for i := 0; i+len(sub) <= len(s); i++ {
if s[i:i+len(sub)] == sub {
return true
}
}
return false
}

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