Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
383aeefaf6 | ||
|
|
a5bede3a19 | ||
|
|
bbf66e23d6 | ||
|
|
5e168c92cf | ||
|
|
725577103d | ||
|
|
f762117d4e | ||
|
|
750b6c04d6 | ||
|
|
59e51c348a | ||
|
|
a056a9abfb | ||
|
|
33ae780103 | ||
|
|
daf56514f7 | ||
|
|
8bcbceb971 | ||
|
|
df01f36442 | ||
|
|
b0024aa669 | ||
|
|
7b7aeadbbe | ||
|
|
94ad422a9f | ||
|
|
4f1ee37508 | ||
|
|
fec0347cd6 | ||
|
|
93318f4a83 | ||
|
|
a14fd0250c | ||
|
|
c99e228669 | ||
|
|
95d495f290 | ||
|
|
4bf300d862 | ||
|
|
1a1fc531f5 | ||
|
|
9fc570607f | ||
|
|
4851d19141 | ||
|
|
42fb25d150 | ||
|
|
416ad6571d | ||
|
|
9a119fbd64 | ||
|
|
c5decb2f90 | ||
|
|
fae2a4f5f0 | ||
|
|
d25b106e4f | ||
|
|
9f78e51ae7 | ||
|
|
d2752d8b5b | ||
|
|
8ecbff391c | ||
|
|
d259864a2b | ||
|
|
408098bdc1 | ||
|
|
658ec1676c | ||
|
|
e36d3b3474 | ||
|
|
0b9952c58d | ||
|
|
56af1ea091 | ||
|
|
ea5859b92b | ||
|
|
19f2ed5c69 | ||
|
|
efbaf7a49d | ||
|
|
374a9e9b13 | ||
|
|
d7d85c9e67 | ||
|
|
0fa982fe91 | ||
|
|
c4fb1bbd3e | ||
|
|
26d7d8f946 | ||
|
|
05ac342c4b | ||
|
|
5e491aef8f | ||
|
|
202187d5e2 | ||
|
|
13877b1c3a | ||
|
|
0e72e89ba3 | ||
|
|
f1b68271cc | ||
|
|
83efff21cd | ||
|
|
e6a4b35921 | ||
|
|
cc2d97ddba | ||
|
|
b78dd19cf9 | ||
|
|
1f0a75f836 | ||
|
|
16202c83a3 | ||
|
|
f4cc76c77d | ||
|
|
9fef6a9c43 | ||
|
|
810985b03a | ||
|
|
02633c6bd3 | ||
|
|
eb9416aa16 | ||
|
|
65b64af213 | ||
|
|
f1d160a481 | ||
|
|
f8c7f012a1 | ||
|
|
45618a55e6 | ||
|
|
c49583836b | ||
|
|
9dc8dc7065 | ||
|
|
f978e306cc | ||
|
|
aec852f971 | ||
|
|
143f781064 | ||
|
|
953b422295 | ||
|
|
da1a0f1299 | ||
|
|
bc7d19cfd8 | ||
|
|
df3122090f | ||
|
|
713fdf6188 | ||
|
|
70e21b58b4 | ||
|
|
18ebba1bb2 | ||
|
|
937404e6df | ||
|
|
88e155dd23 | ||
|
|
2e2cea0973 | ||
|
|
9b8c13a8b6 | ||
|
|
8238cc9f41 | ||
|
|
e59c4f30b8 | ||
|
|
fd7ef5edc2 | ||
|
|
a8e1acec09 | ||
|
|
31eb10985e | ||
|
|
1436b62a80 | ||
|
|
ec6a27635b | ||
|
|
1727744691 | ||
|
|
afdd47b5a5 | ||
|
|
d968e8e551 | ||
|
|
c649d1a762 | ||
|
|
a1f5d97345 | ||
|
|
58062515a5 | ||
|
|
5614b508f2 | ||
|
|
5e003a41b1 | ||
|
|
4eaeb1dd4a | ||
|
|
84471bd6f0 | ||
|
|
c8e3ac21c2 | ||
|
|
c38892b7cf | ||
|
|
1a0a5324f0 | ||
|
|
c1e9e9e0d6 | ||
|
|
57c93243a0 | ||
|
|
4259336e6d | ||
|
|
a18ce2e54d | ||
|
|
a0dc5d6183 | ||
|
|
5149f6808f | ||
|
|
918db33a8b | ||
|
|
95cbde9187 | ||
|
|
5e194393ff | ||
|
|
3da572a76a | ||
|
|
a116cba8ba | ||
|
|
fe5952fe14 | ||
|
|
2ab45ffd90 | ||
|
|
7c4932154c | ||
|
|
e2dbaa7c78 | ||
|
|
11a0dbc84a | ||
|
|
ee441643dd | ||
|
|
4c5affba99 | ||
|
|
c8148ef2cc | ||
|
|
b89740bad6 | ||
|
|
5c0d2b274c | ||
|
|
0489cd67c8 | ||
|
|
15d495e56e | ||
|
|
c74f1eeb26 | ||
|
|
933615003c | ||
|
|
cc4dd1e87b | ||
|
|
3c75c66d4d | ||
|
|
93d6fdb17e | ||
|
|
110f887181 | ||
|
|
25118d1ec7 |
@@ -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 |
@@ -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
|
||||
|
||||
@@ -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(', ') || '无标签'
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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`
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -27,3 +27,6 @@ test/cli_compat/testdata/
|
||||
credentials*
|
||||
plans
|
||||
_docs
|
||||
dws.zip
|
||||
*.code-workspace
|
||||
/dingtalk-workspace.zip
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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|
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
|
||||
// the same host compatibility hooks as MCP. It mirrors the former body of
|
||||
// getCachedRuntimeToken (excluding process-level cache and timing).
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, disc)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token), nil
|
||||
}
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
return "", tokenErr
|
||||
}
|
||||
manager := authpkg.NewManager(configDir, nil)
|
||||
configureLegacyAuthManagerCompatibility(manager)
|
||||
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
|
||||
return strings.TrimSpace(leg), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
|
||||
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
|
||||
// the active edition config directory, the same process-cached path as MCP is used.
|
||||
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
if t := strings.TrimSpace(explicitToken); t != "" {
|
||||
return t, nil
|
||||
}
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
|
||||
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
|
||||
return tok, nil
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
}
|
||||
tok, err := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if tok != "" {
|
||||
return tok, nil
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
}
|
||||
|
||||
func noCredentialsError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return fmt.Errorf("认证信息已失效,请重新认证")
|
||||
}
|
||||
return fmt.Errorf("no credentials found, run: dws auth login")
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResolveAuxiliaryAccessToken_explicitToken(t *testing.T) {
|
||||
tok, err := ResolveAuxiliaryAccessToken(context.Background(), "/any/dir", " bearer-xyz ")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if tok != "bearer-xyz" {
|
||||
t.Fatalf("got %q, want bearer-xyz", tok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveAuxiliaryAccessToken_emptyConfigDir(t *testing.T) {
|
||||
_, err := ResolveAuxiliaryAccessToken(context.Background(), " ", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty config directory")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
},
|
||||
}
|
||||
|
||||
@@ -44,7 +44,7 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
CorpID: "dingcorp",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
}
|
||||
|
||||
originalTransport := http.DefaultTransport
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import "sync"
|
||||
|
||||
// PluginAuth holds authentication credentials for a plugin-owned
|
||||
// streamable-http MCP server. Each server is keyed by its canonical
|
||||
// product ID (CLI.ID) so that different servers can use independent
|
||||
// tokens without interfering with each other or with the default
|
||||
// DingTalk OAuth token.
|
||||
type PluginAuth struct {
|
||||
// Token is the Bearer token extracted from the plugin's
|
||||
// "Authorization" header (e.g. a third-party API key).
|
||||
Token string
|
||||
|
||||
// ExtraHeaders contains any additional custom HTTP headers
|
||||
// declared by the plugin (excluding Authorization).
|
||||
ExtraHeaders map[string]string
|
||||
|
||||
// TrustedDomains lists the hostnames that the token is allowed
|
||||
// to be sent to. Typically derived from the server endpoint.
|
||||
TrustedDomains []string
|
||||
}
|
||||
|
||||
var (
|
||||
pluginAuthMu sync.RWMutex
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
)
|
||||
|
||||
// RegisterPluginAuth stores authentication credentials for a plugin
|
||||
// server keyed by its canonical product ID. The runner looks up these
|
||||
// credentials at execution time to inject the correct Bearer token
|
||||
// instead of the default DingTalk OAuth token.
|
||||
func RegisterPluginAuth(productID string, auth *PluginAuth) {
|
||||
pluginAuthMu.Lock()
|
||||
defer pluginAuthMu.Unlock()
|
||||
pluginAuthRegistry[productID] = auth
|
||||
}
|
||||
|
||||
// LookupPluginAuth returns the authentication credentials registered
|
||||
// for the given product ID, or nil if none exists.
|
||||
func LookupPluginAuth(productID string) (*PluginAuth, bool) {
|
||||
pluginAuthMu.RLock()
|
||||
defer pluginAuthMu.RUnlock()
|
||||
auth, ok := pluginAuthRegistry[productID]
|
||||
return auth, ok
|
||||
}
|
||||
@@ -0,0 +1,213 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestPluginAuthRegistry(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "test-product")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
// Initially not found
|
||||
if _, ok := LookupPluginAuth("test-product"); ok {
|
||||
t.Error("expected LookupPluginAuth to return false for unregistered product")
|
||||
}
|
||||
|
||||
// Register auth credentials
|
||||
auth := &PluginAuth{
|
||||
Token: "sk-test-token-12345",
|
||||
ExtraHeaders: map[string]string{"X-Custom": "value"},
|
||||
TrustedDomains: []string{"api.example.com", "*.example.com"},
|
||||
}
|
||||
RegisterPluginAuth("test-product", auth)
|
||||
|
||||
// Now should be found
|
||||
got, ok := LookupPluginAuth("test-product")
|
||||
if !ok {
|
||||
t.Fatal("expected LookupPluginAuth to return true after registration")
|
||||
}
|
||||
if got != auth {
|
||||
t.Error("LookupPluginAuth returned different auth instance")
|
||||
}
|
||||
if got.Token != "sk-test-token-12345" {
|
||||
t.Errorf("Token = %q, want sk-test-token-12345", got.Token)
|
||||
}
|
||||
if got.ExtraHeaders["X-Custom"] != "value" {
|
||||
t.Errorf("ExtraHeaders[X-Custom] = %q, want value", got.ExtraHeaders["X-Custom"])
|
||||
}
|
||||
if len(got.TrustedDomains) != 2 {
|
||||
t.Errorf("TrustedDomains len = %d, want 2", len(got.TrustedDomains))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginAuthRegistryIsolation(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "product-a")
|
||||
delete(pluginAuthRegistry, "product-b")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
authA := &PluginAuth{Token: "token-a"}
|
||||
authB := &PluginAuth{Token: "token-b"}
|
||||
|
||||
RegisterPluginAuth("product-a", authA)
|
||||
RegisterPluginAuth("product-b", authB)
|
||||
|
||||
gotA, okA := LookupPluginAuth("product-a")
|
||||
gotB, okB := LookupPluginAuth("product-b")
|
||||
|
||||
if !okA || !okB {
|
||||
t.Fatal("expected both products to be registered")
|
||||
}
|
||||
if gotA.Token != "token-a" {
|
||||
t.Errorf("product-a Token = %q, want token-a", gotA.Token)
|
||||
}
|
||||
if gotB.Token != "token-b" {
|
||||
t.Errorf("product-b Token = %q, want token-b", gotB.Token)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeriveToolCLIName(t *testing.T) {
|
||||
tests := []struct {
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{"web_search", "web-search"},
|
||||
{"maps.search_poi", "search-poi"},
|
||||
{"maps.geo", "geo"},
|
||||
{"simple", "simple"},
|
||||
{"a.b.deep_nested_name", "deep-nested-name"},
|
||||
{"already-kebab", "already-kebab"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.input, func(t *testing.T) {
|
||||
got := deriveToolCLIName(tt.input)
|
||||
if got != tt.want {
|
||||
t.Errorf("deriveToolCLIName(%q) = %q, want %q", tt.input, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPluginAuthFromHeaders(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
pluginAuthMu.Lock()
|
||||
delete(pluginAuthRegistry, "test-srv")
|
||||
pluginAuthMu.Unlock()
|
||||
}()
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Key: "test-srv",
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
CLI: market.CLIOverlay{ID: "test-srv", Command: "test-srv"},
|
||||
AuthHeaders: map[string]string{
|
||||
"Authorization": "Bearer sk-my-secret-key",
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
|
||||
auth, ok := LookupPluginAuth("test-srv")
|
||||
if !ok {
|
||||
t.Fatal("expected auth to be registered after registerPluginAuthFromHeaders")
|
||||
}
|
||||
if auth.Token != "sk-my-secret-key" {
|
||||
t.Errorf("Token = %q, want sk-my-secret-key", auth.Token)
|
||||
}
|
||||
if auth.ExtraHeaders["X-Custom"] != "custom-value" {
|
||||
t.Errorf("ExtraHeaders[X-Custom] = %q, want custom-value", auth.ExtraHeaders["X-Custom"])
|
||||
}
|
||||
if len(auth.TrustedDomains) != 2 {
|
||||
t.Fatalf("TrustedDomains len = %d, want 2", len(auth.TrustedDomains))
|
||||
}
|
||||
if auth.TrustedDomains[0] != "api.example.com" {
|
||||
t.Errorf("TrustedDomains[0] = %q, want api.example.com", auth.TrustedDomains[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPluginAuthFromHeadersNoAuth(t *testing.T) {
|
||||
srv := market.ServerDescriptor{
|
||||
Key: "no-auth-srv",
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
CLI: market.CLIOverlay{ID: "no-auth-srv"},
|
||||
AuthHeaders: map[string]string{
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
|
||||
// Should not register because there's no Authorization header
|
||||
if _, ok := LookupPluginAuth("no-auth-srv"); ok {
|
||||
t.Error("expected no auth registration when Authorization header is missing")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildPluginAuthClient(t *testing.T) {
|
||||
base := transport.NewClient(nil)
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Endpoint: "https://dashscope.aliyuncs.com/compatible-mode/v1/mcp",
|
||||
AuthHeaders: map[string]string{
|
||||
"Authorization": "Bearer sk-test-api-key",
|
||||
"X-Extra": "extra-value",
|
||||
},
|
||||
}
|
||||
|
||||
client := buildPluginAuthClient(base, srv)
|
||||
|
||||
// Should return a different client instance
|
||||
if client == base {
|
||||
t.Error("expected buildPluginAuthClient to return a new client, not the base")
|
||||
}
|
||||
|
||||
// Verify trusted domains
|
||||
if len(client.TrustedDomains) != 2 {
|
||||
t.Fatalf("TrustedDomains len = %d, want 2", len(client.TrustedDomains))
|
||||
}
|
||||
if client.TrustedDomains[0] != "dashscope.aliyuncs.com" {
|
||||
t.Errorf("TrustedDomains[0] = %q, want dashscope.aliyuncs.com", client.TrustedDomains[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildPluginAuthClientNoAuth(t *testing.T) {
|
||||
base := transport.NewClient(nil)
|
||||
|
||||
srv := market.ServerDescriptor{
|
||||
Endpoint: "https://api.example.com/mcp/v1",
|
||||
AuthHeaders: map[string]string{
|
||||
"X-Custom": "custom-value",
|
||||
},
|
||||
}
|
||||
|
||||
client := buildPluginAuthClient(base, srv)
|
||||
|
||||
// Should return the base client when no Authorization header
|
||||
if client != base {
|
||||
t.Error("expected buildPluginAuthClient to return base client when no Authorization header")
|
||||
}
|
||||
}
|
||||
+16
-1
@@ -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()
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newConfigCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "config",
|
||||
Short: "配置管理",
|
||||
Long: "管理 DWS CLI 的配置项。查看所有支持的环境变量及其当前值。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
cmd.AddCommand(newConfigListCommand())
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "列出所有可用配置项",
|
||||
Long: "显示 DWS CLI 支持的全部环境变量配置项,包括名称、分类、描述和默认值。",
|
||||
RunE: runConfigList,
|
||||
}
|
||||
cmd.Flags().String("category", "", "按分类过滤 (core|auth|network|security|runtime|debug|external)")
|
||||
cmd.Flags().Bool("show-values", false, "显示配置项的当前实际值 (敏感信息会脱敏)")
|
||||
cmd.Flags().Bool("show-hidden", false, "包含隐藏的内部调试配置项")
|
||||
cmd.Flags().Bool("json", false, "以 JSON 格式输出")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runConfigList(cmd *cobra.Command, _ []string) error {
|
||||
category, _ := cmd.Flags().GetString("category")
|
||||
showValues, _ := cmd.Flags().GetBool("show-values")
|
||||
showHidden, _ := cmd.Flags().GetBool("show-hidden")
|
||||
jsonOut, _ := cmd.Flags().GetBool("json")
|
||||
|
||||
var items []configmeta.ConfigItem
|
||||
if category != "" {
|
||||
items = configmeta.ByCategory(configmeta.Category(category))
|
||||
} else {
|
||||
items = configmeta.All()
|
||||
}
|
||||
|
||||
if !showHidden {
|
||||
items = filterVisible(items)
|
||||
}
|
||||
|
||||
if jsonOut {
|
||||
return writeConfigJSON(cmd, items, showValues)
|
||||
}
|
||||
return writeConfigTable(cmd, items, showValues)
|
||||
}
|
||||
|
||||
func filterVisible(items []configmeta.ConfigItem) []configmeta.ConfigItem {
|
||||
out := make([]configmeta.ConfigItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
if !item.Hidden {
|
||||
out = append(out, item)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func writeConfigJSON(cmd *cobra.Command, items []configmeta.ConfigItem, showValues bool) error {
|
||||
type jsonItem struct {
|
||||
Name string `json:"name"`
|
||||
Category string `json:"category"`
|
||||
Description string `json:"description"`
|
||||
DefaultValue string `json:"default_value,omitempty"`
|
||||
Example string `json:"example,omitempty"`
|
||||
Sensitive bool `json:"sensitive,omitempty"`
|
||||
CurrentValue string `json:"current_value,omitempty"`
|
||||
IsSet bool `json:"is_set"`
|
||||
}
|
||||
|
||||
result := make([]jsonItem, 0, len(items))
|
||||
for _, item := range items {
|
||||
ji := jsonItem{
|
||||
Name: item.Name,
|
||||
Category: string(item.Category),
|
||||
Description: item.Description,
|
||||
DefaultValue: item.DefaultValue,
|
||||
Example: item.Example,
|
||||
Sensitive: item.Sensitive,
|
||||
}
|
||||
val, ok := configmeta.Resolve(item.Name)
|
||||
ji.IsSet = ok
|
||||
if showValues && ok {
|
||||
ji.CurrentValue = val
|
||||
}
|
||||
result = append(result, ji)
|
||||
}
|
||||
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "config_list",
|
||||
"count": len(result),
|
||||
"configs": result,
|
||||
})
|
||||
}
|
||||
|
||||
func writeConfigTable(cmd *cobra.Command, items []configmeta.ConfigItem, showValues bool) error {
|
||||
w := cmd.OutOrStdout()
|
||||
|
||||
if len(items) == 0 {
|
||||
_, _ = fmt.Fprintln(w, "没有找到匹配的配置项。")
|
||||
return nil
|
||||
}
|
||||
|
||||
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
|
||||
|
||||
if showValues {
|
||||
_, _ = fmt.Fprintln(tw, "分类\t配置项\t描述\t默认值\t当前值")
|
||||
_, _ = fmt.Fprintln(tw, "────\t──────\t────\t──────\t──────")
|
||||
} else {
|
||||
_, _ = fmt.Fprintln(tw, "分类\t配置项\t描述\t默认值")
|
||||
_, _ = fmt.Fprintln(tw, "────\t──────\t────\t──────")
|
||||
}
|
||||
|
||||
for _, item := range items {
|
||||
def := item.DefaultValue
|
||||
if def == "" {
|
||||
def = "(空)"
|
||||
}
|
||||
if showValues {
|
||||
val, ok := configmeta.Resolve(item.Name)
|
||||
display := "(未设置)"
|
||||
if ok {
|
||||
display = val
|
||||
}
|
||||
_, _ = fmt.Fprintf(tw, "%s\t%s\t%s\t%s\t%s\n",
|
||||
item.Category, item.Name, item.Description, def, display)
|
||||
} else {
|
||||
_, _ = fmt.Fprintf(tw, "%s\t%s\t%s\t%s\n",
|
||||
item.Category, item.Name, item.Description, def)
|
||||
}
|
||||
}
|
||||
|
||||
_ = tw.Flush()
|
||||
_, _ = fmt.Fprintf(w, "\n共 %d 个配置项。使用 --show-values 查看当前值,--show-hidden 显示隐藏项。\n", len(items))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,177 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func seedTestConfig(t *testing.T) {
|
||||
t.Helper()
|
||||
configmeta.Reset()
|
||||
t.Cleanup(configmeta.Reset)
|
||||
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CONFIG_DIR", Category: configmeta.CategoryCore,
|
||||
Description: "覆盖默认配置目录", DefaultValue: "~/.dws",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET", Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret", Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CATALOG_FIXTURE", Category: configmeta.CategoryDebug,
|
||||
Description: "目录 Fixture 路径", Hidden: true,
|
||||
})
|
||||
}
|
||||
|
||||
func TestConfigListTable(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CONFIG_DIR") {
|
||||
t.Error("expected DWS_CONFIG_DIR in output")
|
||||
}
|
||||
if !strings.Contains(out, "DWS_CLIENT_SECRET") {
|
||||
t.Error("expected DWS_CLIENT_SECRET in output")
|
||||
}
|
||||
// Hidden items should be excluded by default
|
||||
if strings.Contains(out, "DWS_CATALOG_FIXTURE") {
|
||||
t.Error("expected DWS_CATALOG_FIXTURE to be hidden")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListShowHidden(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--show-hidden"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CATALOG_FIXTURE") {
|
||||
t.Error("expected DWS_CATALOG_FIXTURE with --show-hidden")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListCategory(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--category", "auth"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "DWS_CLIENT_SECRET") {
|
||||
t.Error("expected DWS_CLIENT_SECRET for auth category")
|
||||
}
|
||||
if strings.Contains(out, "DWS_CONFIG_DIR") {
|
||||
t.Error("DWS_CONFIG_DIR should not appear for auth category")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListJSON(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--json", "--show-hidden"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var result map[string]any
|
||||
if err := json.Unmarshal(buf.Bytes(), &result); err != nil {
|
||||
t.Fatalf("invalid JSON output: %v", err)
|
||||
}
|
||||
if result["kind"] != "config_list" {
|
||||
t.Errorf("expected kind=config_list, got %v", result["kind"])
|
||||
}
|
||||
count, ok := result["count"].(float64)
|
||||
if !ok || count != 3 {
|
||||
t.Errorf("expected count=3, got %v", result["count"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListShowValues(t *testing.T) {
|
||||
seedTestConfig(t)
|
||||
|
||||
t.Setenv("DWS_CONFIG_DIR", "/custom/dir")
|
||||
t.Setenv("DWS_CLIENT_SECRET", "supersecret123")
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{"--show-values"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "/custom/dir") {
|
||||
t.Error("expected actual value for DWS_CONFIG_DIR")
|
||||
}
|
||||
if strings.Contains(out, "supersecret123") {
|
||||
t.Error("sensitive value should be masked")
|
||||
}
|
||||
if !strings.Contains(out, "当前值") {
|
||||
t.Error("expected '当前值' column header")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigListEmpty(t *testing.T) {
|
||||
configmeta.Reset()
|
||||
defer configmeta.Reset()
|
||||
|
||||
cmd := newConfigListCommand()
|
||||
buf := new(bytes.Buffer)
|
||||
cmd.SetOut(buf)
|
||||
cmd.SetArgs([]string{})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "没有找到") {
|
||||
t.Error("expected empty message")
|
||||
}
|
||||
}
|
||||
@@ -157,6 +157,66 @@ func DirectRuntimeProductIDs() map[string]bool {
|
||||
return ids
|
||||
}
|
||||
|
||||
// AppendDynamicServer adds a single server descriptor to the existing
|
||||
// dynamic server registry without replacing the current entries. This
|
||||
// is used by the plugin loader to inject plugin servers alongside
|
||||
// Market-discovered servers.
|
||||
func AppendDynamicServer(server market.ServerDescriptor) {
|
||||
dynamicMu.Lock()
|
||||
defer dynamicMu.Unlock()
|
||||
|
||||
if dynamicEndpoints == nil {
|
||||
dynamicEndpoints = make(map[string]string)
|
||||
}
|
||||
if dynamicProducts == nil {
|
||||
dynamicProducts = make(map[string]bool)
|
||||
}
|
||||
if dynamicAliases == nil {
|
||||
dynamicAliases = make(map[string]string)
|
||||
}
|
||||
if dynamicToolEndpoints == nil {
|
||||
dynamicToolEndpoints = make(map[string]string)
|
||||
}
|
||||
|
||||
if server.CLI.Skip {
|
||||
return
|
||||
}
|
||||
|
||||
id := strings.TrimSpace(server.CLI.ID)
|
||||
endpoint := strings.TrimSpace(server.Endpoint)
|
||||
if id != "" && endpoint != "" {
|
||||
dynamicEndpoints[id] = endpoint
|
||||
dynamicProducts[id] = true
|
||||
}
|
||||
cmd := strings.TrimSpace(server.CLI.Command)
|
||||
if cmd != "" && cmd != id && endpoint != "" {
|
||||
dynamicEndpoints[cmd] = endpoint
|
||||
dynamicProducts[cmd] = true
|
||||
}
|
||||
for _, alias := range server.CLI.Aliases {
|
||||
alias = strings.TrimSpace(alias)
|
||||
if alias != "" && endpoint != "" {
|
||||
dynamicEndpoints[alias] = endpoint
|
||||
dynamicProducts[alias] = true
|
||||
dynamicAliases[alias] = id
|
||||
}
|
||||
}
|
||||
if endpoint != "" {
|
||||
for _, tool := range server.CLI.Tools {
|
||||
toolName := strings.TrimSpace(tool.Name)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
for toolName := range server.CLI.ToolOverrides {
|
||||
toolName = strings.TrimSpace(toolName)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeDirectRuntimeProductID(productID string) string {
|
||||
dynamicMu.RLock()
|
||||
da := dynamicAliases
|
||||
|
||||
@@ -0,0 +1,438 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/upgrade"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// checkStatus represents the outcome of a single doctor check.
|
||||
type checkStatus string
|
||||
|
||||
const (
|
||||
statusPass checkStatus = "pass"
|
||||
statusWarn checkStatus = "warn"
|
||||
statusFail checkStatus = "fail"
|
||||
)
|
||||
|
||||
// checkResult holds the outcome of a single doctor check.
|
||||
type checkResult struct {
|
||||
Name string `json:"name"`
|
||||
Status checkStatus `json:"status"`
|
||||
Message string `json:"message"`
|
||||
Hint string `json:"hint,omitempty"`
|
||||
Detail any `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
func newDoctorCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "doctor",
|
||||
Short: "环境健康检查",
|
||||
Long: "一键检查登录态、网络连通性、缓存状态和版本更新,快速定位常见问题。",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runDoctor,
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "以 JSON 格式输出")
|
||||
cmd.Flags().Int("timeout", 10, "网络检查超时时间 (秒)")
|
||||
cmd.Flags().Bool("perf", false, "额外展示最近一次性能报告")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runDoctor(cmd *cobra.Command, _ []string) error {
|
||||
jsonOut, _ := cmd.Flags().GetBool("json")
|
||||
timeout, _ := cmd.Flags().GetInt("timeout")
|
||||
if timeout <= 0 {
|
||||
timeout = 10
|
||||
}
|
||||
networkTimeout := time.Duration(timeout) * time.Second
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
checks := make([]checkResult, 0, 4)
|
||||
|
||||
authResult := doctorCheckAuth(cmd.Context(), w, jsonOut)
|
||||
checks = append(checks, authResult)
|
||||
|
||||
networkResult := doctorCheckNetwork(cmd.Context(), w, jsonOut, networkTimeout)
|
||||
checks = append(checks, networkResult)
|
||||
|
||||
cacheResult := doctorCheckCache(w, jsonOut)
|
||||
checks = append(checks, cacheResult)
|
||||
|
||||
versionResult := doctorCheckVersion(w, jsonOut, networkTimeout)
|
||||
checks = append(checks, versionResult)
|
||||
|
||||
showPerf, _ := cmd.Flags().GetBool("perf")
|
||||
if showPerf {
|
||||
perfResult := doctorCheckPerf(w, jsonOut)
|
||||
checks = append(checks, perfResult)
|
||||
}
|
||||
|
||||
pass, warn, fail := countResults(checks)
|
||||
|
||||
if jsonOut {
|
||||
result := map[string]any{
|
||||
"kind": "doctor",
|
||||
"checks": checks,
|
||||
"summary": map[string]int{
|
||||
"pass": pass,
|
||||
"warn": warn,
|
||||
"fail": fail,
|
||||
},
|
||||
}
|
||||
if showPerf {
|
||||
if report, err := LoadLatestReport(); err == nil {
|
||||
result["perf_report"] = report
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(w, result)
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "\n诊断完成: %d 项通过, %d 项警告, %d 项失败\n", pass, warn, fail)
|
||||
if fail > 0 {
|
||||
return fmt.Errorf("诊断发现 %d 项失败", fail)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── Auth check ──────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckAuth(ctx context.Context, w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查登录状态... ")
|
||||
}
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
|
||||
data, err := provider.Status()
|
||||
if err != nil || data == nil {
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "未登录"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 进行登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if data.IsAccessTokenValid() || data.IsRefreshTokenValid() {
|
||||
if !data.IsAccessTokenValid() {
|
||||
refreshCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
|
||||
_, refreshErr := provider.GetAccessToken(refreshCtx)
|
||||
cancel()
|
||||
if refreshErr != nil {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusWarn,
|
||||
Message: "Refresh Token 有效, 但自动刷新 Access Token 失败",
|
||||
Hint: "运行 dws auth login 重新登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "登录已过期"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 重新登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Network check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckNetwork(ctx context.Context, w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查网络连通性... ")
|
||||
}
|
||||
|
||||
baseURL := cli.DefaultMarketBaseURL
|
||||
httpClient := &http.Client{Timeout: timeout}
|
||||
client := market.NewClient(baseURL, httpClient)
|
||||
|
||||
start := time.Now()
|
||||
reqCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
_, err := client.FetchServers(reqCtx, 1)
|
||||
latency := time.Since(start)
|
||||
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 不可达: %v", err),
|
||||
Hint: "请检查网络连接或代理设置",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 可达 (延迟 %dms)", latency.Milliseconds()),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Cache check ─────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckCache(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查缓存状态... ")
|
||||
}
|
||||
|
||||
store := cacheStoreFromEnv()
|
||||
files, _, err := cacheDirectoryStats(store.Root)
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("缓存目录不可读: %v", err),
|
||||
Hint: "运行 dws cache clean 清理后重试",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
entries, _ := store.ListToolsCacheEntries(config.DefaultPartition)
|
||||
|
||||
if files == 0 && len(entries) == 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: "缓存为空 (首次使用)",
|
||||
Hint: "运行任意 dws 命令后将自动建立缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
staleCount := 0
|
||||
for _, e := range entries {
|
||||
if e.Freshness == cache.FreshnessStale {
|
||||
staleCount++
|
||||
}
|
||||
}
|
||||
|
||||
if staleCount > 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("%d 个文件, %d 个工具缓存, %d 个已过期", files, len(entries), staleCount),
|
||||
Hint: "运行 dws cache refresh 刷新缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("%d 个文件, %d 个工具缓存", files, len(entries))
|
||||
if len(entries) > 0 {
|
||||
msg += ", 全部新鲜"
|
||||
}
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusPass,
|
||||
Message: msg,
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Version check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckVersion(w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查版本更新... ")
|
||||
}
|
||||
|
||||
currentVer := version
|
||||
|
||||
client := upgrade.NewClient()
|
||||
latest, err := client.FetchLatestRelease()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("无法获取最新版本: %v", err),
|
||||
Hint: "请检查网络连接",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if upgrade.NeedsUpgrade(currentVer, latest.Version) {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("有新版本 (当前 %s, 最新 v%s)", ensureV(currentVer), latest.Version),
|
||||
Hint: "运行 dws upgrade 升级到最新版本",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("已是最新版本 %s", ensureV(currentVer)),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Output helpers ──────────────────────────────────────────────────────
|
||||
|
||||
func printCheckResult(w io.Writer, r checkResult) {
|
||||
icon := statusIcon(r.Status)
|
||||
fmt.Fprintf(w, "%s %s\n", icon, r.Message)
|
||||
if r.Hint != "" {
|
||||
fmt.Fprintf(w, " %s\n", r.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func statusIcon(s checkStatus) string {
|
||||
switch s {
|
||||
case statusPass:
|
||||
return "✅"
|
||||
case statusWarn:
|
||||
return "⚠️"
|
||||
case statusFail:
|
||||
return "❌"
|
||||
default:
|
||||
return "?"
|
||||
}
|
||||
}
|
||||
|
||||
func countResults(checks []checkResult) (pass, warn, fail int) {
|
||||
for _, c := range checks {
|
||||
switch c.Status {
|
||||
case statusPass:
|
||||
pass++
|
||||
case statusWarn:
|
||||
warn++
|
||||
case statusFail:
|
||||
fail++
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// ── Perf report check ──────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckPerf(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查性能报告... ")
|
||||
}
|
||||
|
||||
report, err := LoadLatestReport()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusWarn,
|
||||
Message: "未找到性能报告",
|
||||
Hint: "设置 DWS_PERF_REPORT=auto 后运行任意命令生成报告",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("报告可用 (%s, %s)", report.Command, report.Timestamp.Local().Format("2006-01-02 15:04")),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
printPerfReportSummary(w, report)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func printPerfReportSummary(w io.Writer, report *PerfReport) {
|
||||
fmt.Fprintf(w, "\n最近一次性能报告 (%s, %s):\n",
|
||||
report.Command, report.Timestamp.Local().Format("2006-01-02 15:04"))
|
||||
|
||||
for _, p := range report.Phases {
|
||||
marker := ""
|
||||
if p.Name == report.Slowest {
|
||||
marker = " ← 最慢"
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s %dms%s\n", p.Name, p.DurationMs, marker)
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s ─────────\n", "─────────────────────────")
|
||||
fmt.Fprintf(w, " %-25s %dms (框架开销 %dms)\n", "总耗时", report.TotalMs, report.OverheadMs)
|
||||
}
|
||||
|
||||
func formatLocalTime(t time.Time) string {
|
||||
if t.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return t.Local().Format("2006-01-02 15:04")
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCountResults(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
{Status: statusWarn},
|
||||
{Status: statusFail},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 1 || fail != 1 {
|
||||
t.Errorf("expected (2,1,1), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountResultsAllPass(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 0 || fail != 0 {
|
||||
t.Errorf("expected (2,0,0), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStatusIcon(t *testing.T) {
|
||||
tests := []struct {
|
||||
status checkStatus
|
||||
want string
|
||||
}{
|
||||
{statusPass, "✅"},
|
||||
{statusWarn, "⚠️"},
|
||||
{statusFail, "❌"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
got := statusIcon(tc.status)
|
||||
if got != tc.want {
|
||||
t.Errorf("statusIcon(%q) = %q, want %q", tc.status, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResult(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusFail,
|
||||
Message: "something broke",
|
||||
Hint: "try fixing it",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "❌") {
|
||||
t.Error("expected fail icon")
|
||||
}
|
||||
if !strings.Contains(out, "something broke") {
|
||||
t.Error("expected message")
|
||||
}
|
||||
if !strings.Contains(out, "try fixing it") {
|
||||
t.Error("expected hint")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResultNoHint(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusPass,
|
||||
Message: "all good",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "✅") {
|
||||
t.Error("expected pass icon")
|
||||
}
|
||||
lines := strings.Split(strings.TrimSpace(out), "\n")
|
||||
if len(lines) != 1 {
|
||||
t.Errorf("expected 1 line (no hint), got %d", len(lines))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmpty(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, false)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if !strings.Contains(r.Message, "缓存为空") {
|
||||
t.Errorf("expected empty cache message, got %q", r.Message)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmptyJSON(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, true)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if buf.Len() != 0 {
|
||||
t.Error("expected no output in JSON mode")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCommandStructure(t *testing.T) {
|
||||
cmd := newDoctorCommand()
|
||||
if cmd.Use != "doctor" {
|
||||
t.Errorf("Use = %q, want doctor", cmd.Use)
|
||||
}
|
||||
|
||||
jsonFlag := cmd.Flags().Lookup("json")
|
||||
if jsonFlag == nil {
|
||||
t.Error("expected --json flag")
|
||||
}
|
||||
timeoutFlag := cmd.Flags().Lookup("timeout")
|
||||
if timeoutFlag == nil {
|
||||
t.Error("expected --timeout flag")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckResultJSONMarshal(t *testing.T) {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
data, err := json.Marshal(r)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parsed["name"] != "auth" {
|
||||
t.Errorf("expected name=auth, got %v", parsed["name"])
|
||||
}
|
||||
if parsed["status"] != "pass" {
|
||||
t.Errorf("expected status=pass, got %v", parsed["status"])
|
||||
}
|
||||
if _, hasHint := parsed["hint"]; hasHint {
|
||||
t.Error("empty hint should be omitted")
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,7 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -35,8 +36,8 @@ type GlobalFlags struct {
|
||||
}
|
||||
|
||||
func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", "Override OAuth client ID (DingTalk AppKey)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", "Override OAuth client secret (DingTalk AppSecret)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", i18n.T("覆盖 OAuth 客户端 ID (钉钉 AppKey)"))
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", i18n.T("覆盖 OAuth 客户端密钥 (钉钉 AppSecret)"))
|
||||
cmd.PersistentFlags().BoolVar(&flags.Debug, "debug", false, "显示调试日志")
|
||||
cmd.PersistentFlags().BoolVar(&flags.DryRun, "dry-run", false, "预览操作内容,不实际执行")
|
||||
cmd.PersistentFlags().StringVar(&flags.Fields, "fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -0,0 +1,606 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
|
||||
func TestIsPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected missing_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_PlainString(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user lacks required scope",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user lacks required scope",
|
||||
}
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected plain string with missing_scope to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_NotScopeError(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewValidation("invalid parameter")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("expected validation error NOT to be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if isPatScopeError(nil) {
|
||||
t.Fatal("nil error should not be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_WithReason(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("API error",
|
||||
apperrors.WithReason("missing_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected error with missing_scope reason to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_InsufficientScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("insufficient_scope for resource",
|
||||
apperrors.WithReason("insufficient_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected insufficient_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.ErrorType != "missing_scope" {
|
||||
t.Errorf("expected error type 'missing_scope', got %q", scopeErr.ErrorType)
|
||||
}
|
||||
if !strings.Contains(scopeErr.Hint, "dws auth login") {
|
||||
t.Errorf("expected hint to contain 'dws auth login', got %q", scopeErr.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ExtractsScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user needs calendar:read",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user needs calendar:read",
|
||||
}
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:read" {
|
||||
t.Errorf("expected MissingScope 'calendar:read', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthError_HumanReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:user_mailbox.message:send",
|
||||
Hint: "run `dws auth login --scope \"mail:user_mailbox.message:send\"` to authorize",
|
||||
MissingScope: "mail:user_mailbox.message:send",
|
||||
}
|
||||
PrintPatAuthError(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "missing_scope") {
|
||||
t.Errorf("expected output to contain 'missing_scope', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "dws auth login") {
|
||||
t.Errorf("expected output to contain 'dws auth login', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "需要额外授权") {
|
||||
t.Errorf("expected output to contain Chinese auth prompt, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthJSON_MachineReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:send",
|
||||
Hint: "run dws auth login --scope mail:send",
|
||||
MissingScope: "mail:send",
|
||||
}
|
||||
PrintPatAuthJSON(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, `"ok": false`) {
|
||||
t.Errorf("expected JSON to contain ok: false, got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, `"missing_scope": "mail:send"`) {
|
||||
t.Errorf("expected JSON to contain missing_scope, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_BusinessPermissionDenied(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Generic business "permission denied" should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("User has no permission to access this mailbox, permission denied")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("generic 'permission denied' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_GenericForbidden(t *testing.T) {
|
||||
t.Parallel()
|
||||
// HTTP 403 Forbidden should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("403 Forbidden")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("'403 Forbidden' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ComplexScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "mail:user_mailbox.message:send" {
|
||||
t.Errorf("expected MissingScope 'mail:user_mailbox.message:send', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatScopeError_Error(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "test error message",
|
||||
}
|
||||
if err.Error() != "test error message" {
|
||||
t.Errorf("expected Error() to return OriginalError, got %q", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow integration tests — httptest mock covering four terminal
|
||||
// states: APPROVED, REJECTED, EXPIRED, CANCELLED (ctx cancel).
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// setupPollServer creates an httptest server that responds to
|
||||
// /cli/oauth/device/poll?flowId=<fid> with the given status sequence.
|
||||
// It also writes the server URL into a temp DWS_CONFIG_DIR/mcp_url so that
|
||||
// GetMCPBaseURL() returns the test server address.
|
||||
func setupPollServer(t *testing.T, statuses []authpkg.DevicePollResponse) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var callCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
idx := int(callCount.Add(1)) - 1
|
||||
if idx >= len(statuses) {
|
||||
idx = len(statuses) - 1
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(statuses[idx])
|
||||
}))
|
||||
|
||||
// Write mcp_url so GetMCPBaseURL picks up the test server.
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Approved(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "APPROVED", AuthCode: "code123"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-1", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED, got %q", status)
|
||||
}
|
||||
if authCode != "code123" {
|
||||
t.Errorf("expected authCode 'code123', got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Rejected(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "REJECTED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-2", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "REJECTED" {
|
||||
t.Errorf("expected REJECTED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for REJECTED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Expired(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "EXPIRED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-3", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for EXPIRED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Cancelled(t *testing.T) {
|
||||
// Server always returns PENDING so context cancellation is the only exit.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Cancel immediately after first poll tick.
|
||||
go func() {
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
cancel()
|
||||
}()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-4", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "CANCELLED" {
|
||||
t.Errorf("expected CANCELLED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for CANCELLED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// IsPatRetrying tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsPatRetrying_Default(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
if IsPatRetrying(ctx) {
|
||||
t.Fatal("expected false for plain context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatRetrying_WithValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.WithValue(context.Background(), patRetryingKey, true)
|
||||
if !IsPatRetrying(ctx) {
|
||||
t.Fatal("expected true when pat retry key is set")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPollPatDeviceFlow_ServerErrorFallback(t *testing.T) {
|
||||
// When server returns success=false with empty status, should treat as EXPIRED.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: ""}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-err", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED for server error fallback, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for server error, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_RedirectSkipped(t *testing.T) {
|
||||
// When server returns 302 (SSO redirect), poll should continue until real response.
|
||||
var callCount int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
callCount++
|
||||
if callCount <= 1 {
|
||||
// First call: simulate SSO redirect
|
||||
w.Header().Set("Location", "https://sso.example.com")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
return
|
||||
}
|
||||
// Second call: return APPROVED
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
resp := authpkg.DevicePollResponse{
|
||||
Success: true,
|
||||
Data: authpkg.DevicePollData{Status: "APPROVED"},
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, _, err := pollPatDeviceFlow(ctx, "flow-redirect", tmpDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED after redirect, got %q", status)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// extractPatScopeError edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestExtractPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := extractPatScopeError(nil); got != nil {
|
||||
t.Fatalf("expected nil for nil error, got %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_WithIdentity(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth(`insufficient_scope: identity "app_user" needs calendar:write`)
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.Identity != "app_user" {
|
||||
t.Errorf("expected Identity 'app_user', got %q", scopeErr.Identity)
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:write" {
|
||||
t.Errorf("expected MissingScope 'calendar:write', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// handlePatAuthCheck integration tests — cover the main orchestrator with
|
||||
// mock runner + httptest poll server for APPROVED, REJECTED, EmptyFlowID.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// mockRunner is a simple executor.Runner for testing handlePatAuthCheck.
|
||||
type mockRunner struct {
|
||||
runFunc func(ctx context.Context, inv executor.Invocation) (executor.Result, error)
|
||||
}
|
||||
|
||||
func (m *mockRunner) Run(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
return m.runFunc(ctx, inv)
|
||||
}
|
||||
|
||||
// setupHandlePATServer creates an httptest server for handlePatAuthCheck tests.
|
||||
// It responds to device poll requests with the given status after the first poll.
|
||||
func setupHandlePATServer(t *testing.T, terminalStatus string, authCode string) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var pollCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.Contains(r.URL.Path, "/cli/oauth/device/poll") {
|
||||
idx := int(pollCount.Add(1)) - 1
|
||||
var resp authpkg.DevicePollResponse
|
||||
if idx == 0 {
|
||||
resp = authpkg.DevicePollResponse{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}}
|
||||
} else {
|
||||
resp = authpkg.DevicePollResponse{
|
||||
Success: terminalStatus == "APPROVED",
|
||||
Data: authpkg.DevicePollData{Status: terminalStatus, AuthCode: authCode},
|
||||
}
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func makePATErrorJSON(flowID, clientID string) string {
|
||||
type patData struct {
|
||||
Desc string `json:"desc"`
|
||||
FlowID string `json:"flowId"`
|
||||
URI string `json:"uri"`
|
||||
ClientID string `json:"clientId"`
|
||||
}
|
||||
payload := struct {
|
||||
Code string `json:"code"`
|
||||
Data patData `json:"data"`
|
||||
}{
|
||||
Code: "AGENT_CODE_NOT_EXISTS",
|
||||
Data: patData{
|
||||
Desc: "test auth",
|
||||
FlowID: flowID,
|
||||
URI: "", // empty to avoid opening browser in test
|
||||
ClientID: clientID,
|
||||
},
|
||||
}
|
||||
data, _ := json.Marshal(payload)
|
||||
return string(data)
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Approved(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "APPROVED", "test-auth-code")
|
||||
defer server.Close()
|
||||
|
||||
var retryCalled bool
|
||||
var retryHasKey bool
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
retryCalled = true
|
||||
retryHasKey = IsPatRetrying(ctx)
|
||||
return executor.Result{Response: map[string]any{"ok": true}}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-approved", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !retryCalled {
|
||||
t.Fatal("expected mock runner to be called for retry")
|
||||
}
|
||||
if !retryHasKey {
|
||||
t.Fatal("expected retry context to have patRetryingKey")
|
||||
}
|
||||
// Verify SetClientIDFromMCP was called with the PAT response clientId.
|
||||
if cid := authpkg.ClientID(); cid != "test-client-id" {
|
||||
t.Errorf("expected ClientID 'test-client-id', got %q", cid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Rejected(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "REJECTED", "")
|
||||
defer server.Close()
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called on REJECTED")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-rejected", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for REJECTED")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "用户已拒绝授权") {
|
||||
t.Errorf("expected rejection error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_EmptyFlowID_FallsBackToPATError(t *testing.T) {
|
||||
// No poll server needed — empty flowId means no polling, return PATError directly.
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called when flowId is empty")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, tmpDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected PATError when flowId is empty")
|
||||
}
|
||||
// Should return the original PATError.
|
||||
if _, ok := err.(*apperrors.PATError); !ok {
|
||||
t.Errorf("expected *PATError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,664 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newPluginCommand() *cobra.Command {
|
||||
pluginCmd := newPlaceholderParent("plugin", i18n.T("插件管理"))
|
||||
|
||||
pluginCmd.AddCommand(
|
||||
newPluginListCommand(),
|
||||
newPluginInstallCommand(),
|
||||
newPluginInfoCommand(),
|
||||
newPluginEnableCommand(),
|
||||
newPluginDisableCommand(),
|
||||
newPluginRemoveCommand(),
|
||||
newPluginValidateCommand(),
|
||||
newPluginCreateCommand(),
|
||||
newPluginDevCommand(),
|
||||
newPluginConfigCommand(),
|
||||
newPluginBuildCommand(),
|
||||
)
|
||||
|
||||
return pluginCmd
|
||||
}
|
||||
|
||||
func newPluginListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("列出已安装的插件"),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
if wantJSON {
|
||||
return output.WriteJSON(cmd.OutOrStdout(), plugins)
|
||||
}
|
||||
|
||||
if len(plugins) == 0 {
|
||||
fmt.Fprintln(cmd.OutOrStdout(), "No plugins installed.")
|
||||
return nil
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
"NAME", "VERSION", "TYPE", "STATUS", "DESCRIPTION")
|
||||
fmt.Fprintln(w, strings.Repeat("-", 85))
|
||||
for _, p := range plugins {
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
p.Name, p.Version, p.Type, statusStr(p.Enabled), p.Description)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install",
|
||||
Short: i18n.T("安装插件"),
|
||||
Example: ` dws plugin install --dir ./conference
|
||||
dws plugin install --git https://github.com/DingTalk-Real-AI/conference.git`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dirPath, _ := cmd.Flags().GetString("dir")
|
||||
gitURL, _ := cmd.Flags().GetString("git")
|
||||
|
||||
if dirPath == "" && gitURL == "" {
|
||||
return apperrors.NewValidation("specify install source: --dir <path> or --git <url>")
|
||||
}
|
||||
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if gitURL != "" {
|
||||
p, err := loader.InstallFromGit(gitURL)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
}
|
||||
|
||||
p, err := loader.InstallFromDir(dirPath)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("dir", "", "Install from a local directory")
|
||||
cmd.Flags().String("git", "", "Install from a Git repository")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInfoCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "info <name>",
|
||||
Short: i18n.T("查看插件详情"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
for _, p := range plugins {
|
||||
if p.Name == name {
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Name: %s\n", p.Name)
|
||||
fmt.Fprintf(w, "Version: %s\n", p.Version)
|
||||
fmt.Fprintf(w, "Type: %s\n", p.Type)
|
||||
fmt.Fprintf(w, "Status: %s\n", statusStr(p.Enabled))
|
||||
fmt.Fprintf(w, "Path: %s\n", p.Path)
|
||||
if p.Description != "" {
|
||||
fmt.Fprintf(w, "Description: %s\n", p.Description)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found", name))
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginEnableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "enable <name>",
|
||||
Short: i18n.T("启用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], true); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s enabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginDisableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "disable <name>",
|
||||
Short: i18n.T("禁用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], false); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s disabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginRemoveCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "remove <name>",
|
||||
Short: i18n.T("卸载已安装的插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
// Stop stdio clients before removing to release file locks
|
||||
StopStdioClientsByPlugin(args[0])
|
||||
keepData, _ := cmd.Flags().GetBool("keep-data")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.RemovePlugin(args[0], keepData); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s removed.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("keep-data", false, "Keep plugin data directory")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginValidateCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "validate <dir>",
|
||||
Short: i18n.T("校验 plugin.json"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
m, err := plugin.ParseManifest(dir + "/plugin.json")
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("parse failed: %v", err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Valid: %s (%s)\n", m.Name, m.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginCreateCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create <name>",
|
||||
Short: i18n.T("脚手架生成新插件目录"),
|
||||
Example: ` dws plugin create my-tool
|
||||
dws plugin create my-tool --description "My awesome tool"`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
desc, _ := cmd.Flags().GetString("description")
|
||||
pluginType := "user"
|
||||
|
||||
// Validate name format
|
||||
m := &plugin.Manifest{Name: name, Version: "0.1.0", Type: pluginType}
|
||||
if err := m.Validate(""); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin name: %v", err))
|
||||
}
|
||||
|
||||
dir := filepath.Join(".", name)
|
||||
if _, err := os.Stat(dir); err == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("directory %q already exists", dir))
|
||||
}
|
||||
|
||||
// Create directory structure
|
||||
dirs := []string{
|
||||
dir,
|
||||
filepath.Join(dir, "skills", name),
|
||||
filepath.Join(dir, "hooks"),
|
||||
}
|
||||
for _, d := range dirs {
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
|
||||
}
|
||||
}
|
||||
|
||||
// Write plugin.json
|
||||
pluginJSON := fmt.Sprintf(`{
|
||||
"name": %q,
|
||||
"version": "0.1.0",
|
||||
"description": %q,
|
||||
"type": %q,
|
||||
"minCLIVersion": %q,
|
||||
"mcpServers": {
|
||||
%q: {
|
||||
"type": "stdio",
|
||||
"command": "${DWS_PLUGIN_ROOT}/bin/server",
|
||||
"args": []
|
||||
}
|
||||
},
|
||||
"build": {
|
||||
"command": "echo 'TODO: replace with your build command, e.g.: bun build --compile src/server.ts --outfile bin/server'",
|
||||
"output": "bin/server"
|
||||
},
|
||||
"skills": "./skills/",
|
||||
"hooks": "./hooks/hooks.json"
|
||||
}
|
||||
`, name, desc, pluginType, RawVersion(), name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(pluginJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write plugin.json: %v", err))
|
||||
}
|
||||
|
||||
// Write SKILL.md template
|
||||
skillMD := fmt.Sprintf(`---
|
||||
name: %s
|
||||
description: %s
|
||||
cli_version: ">=%s"
|
||||
---
|
||||
|
||||
# %s
|
||||
|
||||
## Intent Recognition
|
||||
|
||||
Use this skill when the user mentions:
|
||||
- TODO: add your intent keywords here
|
||||
|
||||
## Command Decision Tree
|
||||
|
||||
| User Intent | Command | Required Parameters |
|
||||
|-------------|---------|---------------------|
|
||||
| TODO | `+"`dws %s <sub-command>`"+` | `+"`--param`"+` |
|
||||
|
||||
## Parameter Rules
|
||||
|
||||
### TODO: parameter type
|
||||
- Format description
|
||||
- Conversion rules
|
||||
`, name, desc, RawVersion(), name, name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "skills", name, "SKILL.md"), []byte(skillMD), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write SKILL.md: %v", err))
|
||||
}
|
||||
|
||||
// Write hooks.json template
|
||||
hooksJSON := `{
|
||||
"hooks": []
|
||||
}
|
||||
`
|
||||
if err := os.WriteFile(filepath.Join(dir, "hooks", "hooks.json"), []byte(hooksJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write hooks.json: %v", err))
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Created plugin scaffold at ./%s/\n", name)
|
||||
fmt.Fprintf(w, " %s/\n", name)
|
||||
fmt.Fprintf(w, " ├── plugin.json\n")
|
||||
fmt.Fprintf(w, " ├── skills/%s/SKILL.md\n", name)
|
||||
fmt.Fprintf(w, " └── hooks/hooks.json\n")
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "Next steps:\n")
|
||||
fmt.Fprintf(w, " 1. Edit plugin.json to configure your MCP servers\n")
|
||||
fmt.Fprintf(w, " 2. Edit skills/%s/SKILL.md to describe your commands\n", name)
|
||||
fmt.Fprintf(w, " 3. Run: dws plugin validate ./%s\n", name)
|
||||
fmt.Fprintf(w, " 4. Run: dws plugin dev ./%s\n", name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("description", "", "Plugin description")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginDevCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "dev <dir>",
|
||||
Short: i18n.T("将本地目录注册为开发态插件"),
|
||||
Long: `Registers a plugin from a local source directory for development.
|
||||
The plugin is loaded directly from the source directory on next CLI invocation,
|
||||
without copying files to ~/.dws/plugins/. Use 'dws plugin dev --off <name>'
|
||||
to unregister.`,
|
||||
Example: ` dws plugin dev ./my-tool
|
||||
dws plugin dev --off my-tool`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
off, _ := cmd.Flags().GetBool("off")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if off {
|
||||
// Unregister dev plugin
|
||||
name := args[0]
|
||||
if err := loader.UnregisterDevPlugin(name); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q unregistered.\n", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Register dev plugin
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
// Validate the plugin first
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
|
||||
if err := loader.RegisterDevPlugin(m.Name, absDir); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to register: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q registered from %s\n", m.Name, absDir)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "It will be loaded on next dws invocation.\n")
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "To unregister: dws plugin dev --off %s\n", m.Name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("off", false, "Unregister a dev plugin")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigCommand() *cobra.Command {
|
||||
configCmd := newPlaceholderParent("config", i18n.T("管理插件配置"))
|
||||
configCmd.AddCommand(
|
||||
newPluginConfigSetCommand(),
|
||||
newPluginConfigGetCommand(),
|
||||
newPluginConfigListCommand(),
|
||||
newPluginConfigUnsetCommand(),
|
||||
)
|
||||
return configCmd
|
||||
}
|
||||
|
||||
func newPluginConfigSetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "set <plugin-name> <key> <value>",
|
||||
Short: i18n.T("设置插件配置项"),
|
||||
Long: `Persistently set a configuration value for a plugin.
|
||||
The value is stored in ~/.dws/settings.json and automatically injected
|
||||
as an environment variable when the plugin is loaded.
|
||||
|
||||
Environment variables set by the user (e.g. via export) take precedence
|
||||
over values stored in settings.json.`,
|
||||
Example: ` dws plugin config set demo-devtool DASHSCOPE_API_KEY sk-xxx
|
||||
dws plugin config set my-plugin API_ENDPOINT https://api.example.com`,
|
||||
Args: cobra.ExactArgs(3),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key, value := args[0], args[1], args[2]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// Validate that the plugin exists.
|
||||
plugins := loader.ListInstalled()
|
||||
found := false
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found; use 'dws plugin list' to see installed plugins", pluginName))
|
||||
}
|
||||
|
||||
loader.SetPluginConfig(pluginName, key, value)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config saved: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigGetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "get <plugin-name> <key>",
|
||||
Short: i18n.T("读取插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
val, ok := loader.GetPluginConfig(pluginName, key)
|
||||
if !ok {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintln(cmd.OutOrStdout(), val)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list <plugin-name>",
|
||||
Short: i18n.T("列出插件所有配置项"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
configs := loader.ListPluginConfig(pluginName)
|
||||
|
||||
// Also load the plugin manifest to show declared userConfig keys.
|
||||
declaredKeys := loadDeclaredUserConfig(loader, pluginName)
|
||||
|
||||
if wantJSON {
|
||||
result := make(map[string]any)
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
if sensitive {
|
||||
result[k] = maskSensitiveValue(v)
|
||||
} else {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
// Include declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
entry := map[string]any{
|
||||
"value": nil,
|
||||
"description": ci.Description,
|
||||
"required": ci.Default == "",
|
||||
}
|
||||
result[k] = entry
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "plugin_config",
|
||||
"plugin": pluginName,
|
||||
"config": result,
|
||||
})
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
if len(configs) == 0 && len(declaredKeys) == 0 {
|
||||
fmt.Fprintf(w, "No configuration for plugin %q.\n", pluginName)
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "Configuration for %s:\n\n", pluginName)
|
||||
|
||||
// Show set values.
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
displayVal := v
|
||||
if sensitive {
|
||||
displayVal = maskSensitiveValue(v)
|
||||
}
|
||||
fmt.Fprintf(w, " %s = %s\n", k, displayVal)
|
||||
}
|
||||
|
||||
// Show declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
desc := ""
|
||||
if ci.Description != "" {
|
||||
desc = " # " + ci.Description
|
||||
}
|
||||
fmt.Fprintf(w, " %s = (not set)%s\n", k, desc)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigUnsetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "unset <plugin-name> <key>",
|
||||
Short: i18n.T("删除插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if !loader.UnsetPluginConfig(pluginName, key) {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config removed: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// loadDeclaredUserConfig loads the userConfig section from a plugin's manifest.
|
||||
func loadDeclaredUserConfig(loader *plugin.Loader, pluginName string) map[string]plugin.ConfigItem {
|
||||
plugins := loader.ListInstalled()
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
m, err := plugin.ParseManifest(filepath.Join(p.Path, "plugin.json"))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return m.UserConfig
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// maskSensitiveValue masks a sensitive value, showing only the first 4
|
||||
// and last 2 characters for values longer than 8 characters.
|
||||
func maskSensitiveValue(value string) string {
|
||||
if len(value) <= 8 {
|
||||
return strings.Repeat("*", len(value))
|
||||
}
|
||||
return value[:4] + strings.Repeat("*", len(value)-6) + value[len(value)-2:]
|
||||
}
|
||||
|
||||
func newPluginBuildCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "build <dir>",
|
||||
Short: i18n.T("将插件 stdio server 编译为原生二进制"),
|
||||
Long: `Runs the build command declared in plugin.json to compile the
|
||||
plugin's server into a single executable. This ensures plugin users
|
||||
don't need any language runtime (Node.js, Python, etc.) installed.
|
||||
|
||||
The build configuration is read from the "build" field in plugin.json:
|
||||
|
||||
{
|
||||
"build": {
|
||||
"command": "bun build --compile src/server.ts --outfile bin/server",
|
||||
"output": "bin/server"
|
||||
}
|
||||
}`,
|
||||
Example: ` dws plugin build ./my-plugin
|
||||
dws plugin build .`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
|
||||
if m.Build == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf(
|
||||
"plugin %q has no \"build\" field in plugin.json.\n"+
|
||||
"Add a build config, e.g.:\n\n"+
|
||||
" \"build\": {\n"+
|
||||
" \"command\": \"bun build --compile src/server.js --outfile bin/server\",\n"+
|
||||
" \"output\": \"bin/server\"\n"+
|
||||
" }", m.Name))
|
||||
}
|
||||
|
||||
if err := plugin.BuildPlugin(absDir); err != nil {
|
||||
return apperrors.NewInternal(err.Error())
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Build succeeded: %s\n", m.Build.Output)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func statusStr(enabled bool) string {
|
||||
if enabled {
|
||||
return "enabled"
|
||||
}
|
||||
return "disabled"
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
// TestSharedCacheStoreConcurrentSaveTools verifies that a single *cache.Store
|
||||
// instance is safe for goroutines saving tool snapshots concurrently, as long
|
||||
// as each goroutine targets a distinct (partition, serverKey). This mirrors
|
||||
// the real plugin discovery path where each goroutine owns one plugin/server.
|
||||
//
|
||||
// Each call serializes to its own "<key>.json.tmp" file followed by a
|
||||
// rename(2) to the final path, so concurrent writers targeting distinct keys
|
||||
// never collide. The invariant asserted here: after N parallel writes, the
|
||||
// Store returns each written snapshot intact under LoadTools.
|
||||
func TestSharedCacheStoreConcurrentSaveTools(t *testing.T) {
|
||||
const (
|
||||
partition = "default/default"
|
||||
writers = 16
|
||||
)
|
||||
|
||||
store := cache.NewStore(t.TempDir())
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < writers; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", idx)
|
||||
if err := store.SaveTools(partition, key, cache.ToolsSnapshot{
|
||||
ServerKey: key,
|
||||
}); err != nil {
|
||||
t.Errorf("SaveTools(%s): %v", key, err)
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < writers; i++ {
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", i)
|
||||
snapshot, _, err := store.LoadTools(partition, key)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTools(%s): %v", key, err)
|
||||
}
|
||||
if snapshot.ServerKey != key {
|
||||
t.Errorf("LoadTools(%s) returned ServerKey %q", key, snapshot.ServerKey)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppendDynamicServerConcurrent exercises the dynamicMu mutex on the
|
||||
// write path by spraying distinct server descriptors in parallel. Afterwards
|
||||
// every injected product ID must be resolvable — a missing entry would
|
||||
// indicate a lost write through an un-synchronized map update.
|
||||
func TestAppendDynamicServerConcurrent(t *testing.T) {
|
||||
dynamicMu.Lock()
|
||||
prev := struct {
|
||||
endpoints map[string]string
|
||||
products map[string]bool
|
||||
aliases map[string]string
|
||||
toolEndpoints map[string]string
|
||||
}{dynamicEndpoints, dynamicProducts, dynamicAliases, dynamicToolEndpoints}
|
||||
dynamicEndpoints = nil
|
||||
dynamicProducts = nil
|
||||
dynamicAliases = nil
|
||||
dynamicToolEndpoints = nil
|
||||
dynamicMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
dynamicMu.Lock()
|
||||
dynamicEndpoints = prev.endpoints
|
||||
dynamicProducts = prev.products
|
||||
dynamicAliases = prev.aliases
|
||||
dynamicToolEndpoints = prev.toolEndpoints
|
||||
dynamicMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
id := fmt.Sprintf("plugin-id-%d", idx)
|
||||
endpoint := fmt.Sprintf("https://example.test/%d", idx)
|
||||
AppendDynamicServer(market.ServerDescriptor{
|
||||
Endpoint: endpoint,
|
||||
CLI: market.CLIOverlay{
|
||||
ID: id,
|
||||
Command: id,
|
||||
},
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
id := fmt.Sprintf("plugin-id-%d", i)
|
||||
if endpoint, ok := directRuntimeEndpoint(id, ""); !ok || endpoint == "" {
|
||||
t.Errorf("directRuntimeEndpoint(%q) = (%q, %v), want non-empty", id, endpoint, ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegisterStdioClientConcurrent verifies the stdioMu-protected registry
|
||||
// survives concurrent writers — every registered client must be looked up
|
||||
// afterwards. Uses nil client pointers since LookupStdioClient only compares
|
||||
// keys, not values.
|
||||
func TestRegisterStdioClientConcurrent(t *testing.T) {
|
||||
stdioMu.Lock()
|
||||
prev := stdioClients
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
stdioMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
stdioMu.Lock()
|
||||
stdioClients = prev
|
||||
stdioMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
RegisterStdioClient(fmt.Sprintf("plugin/%d", idx), nil)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
key := fmt.Sprintf("plugin/%d", i)
|
||||
if _, ok := LookupStdioClient(key); !ok {
|
||||
t.Errorf("LookupStdioClient(%q) missing after concurrent registration", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolvePluginColdTimeouts covers the three code paths of the env
|
||||
// parser: unset (defaults), valid duration (applied to all three slots),
|
||||
// and invalid duration (logged and ignored, defaults returned).
|
||||
func TestResolvePluginColdTimeouts(t *testing.T) {
|
||||
t.Run("defaults when env unset", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("httpNoAuth = %v, want 1s", got.httpNoAuth)
|
||||
}
|
||||
if got.httpAuth != 1500*time.Millisecond {
|
||||
t.Errorf("httpAuth = %v, want 1.5s", got.httpAuth)
|
||||
}
|
||||
if got.stdio != 2*time.Second {
|
||||
t.Errorf("stdio = %v, want 2s", got.stdio)
|
||||
}
|
||||
})
|
||||
t.Run("env override applies to all slots", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "3500ms")
|
||||
got := resolvePluginColdTimeouts()
|
||||
want := 3500 * time.Millisecond
|
||||
if got.httpNoAuth != want || got.httpAuth != want || got.stdio != want {
|
||||
t.Errorf("override not propagated: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("invalid env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "not-a-duration")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second || got.stdio != 2*time.Second {
|
||||
t.Errorf("invalid env should not override defaults: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("non-positive env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "0")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("zero duration should not override defaults, got %v", got.httpNoAuth)
|
||||
}
|
||||
})
|
||||
}
|
||||
+764
-36
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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, °raded) {
|
||||
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
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -25,12 +26,61 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
func setupRuntimeCommandTest(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
discoverySrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_ = json.NewEncoder(w).Encode(contactDiscoveryResponse())
|
||||
}))
|
||||
t.Cleanup(func() { discoverySrv.Close() })
|
||||
SetDiscoveryBaseURL(discoverySrv.URL)
|
||||
t.Cleanup(func() { SetDiscoveryBaseURL("") })
|
||||
}
|
||||
|
||||
func contactDiscoveryResponse() map[string]any {
|
||||
return map[string]any{
|
||||
"metadata": map[string]any{"count": 1, "nextCursor": ""},
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"server": map[string]any{
|
||||
"name": "Contact",
|
||||
"description": "通讯录",
|
||||
"remotes": []any{
|
||||
map[string]any{
|
||||
"type": "streamable-http",
|
||||
"url": "https://mcp.dingtalk.com/contact/v1",
|
||||
},
|
||||
},
|
||||
},
|
||||
"_meta": map[string]any{
|
||||
"com.dingtalk.mcp.registry/metadata": map[string]any{
|
||||
"status": "active", "isLatest": true,
|
||||
},
|
||||
"com.dingtalk.mcp.registry/cli": map[string]any{
|
||||
"id": "contact",
|
||||
"command": "contact",
|
||||
"groups": map[string]any{
|
||||
"user": map[string]any{
|
||||
"description": "用户管理",
|
||||
},
|
||||
},
|
||||
"toolOverrides": map[string]any{
|
||||
"get_current_user_profile": map[string]any{
|
||||
"cliName": "get-self",
|
||||
"group": "user",
|
||||
"flags": map[string]any{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
|
||||
@@ -596,6 +646,95 @@ func contentScanServer() *mockmcp.Server {
|
||||
return mockmcp.MustNewServer(fixture)
|
||||
}
|
||||
|
||||
func TestClassifyToolResultHookPreemptsBusinessError(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
t.Setenv("DWS_TRUSTED_DOMAINS", "*")
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, "bad request", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
method, _ := req["method"].(string)
|
||||
switch method {
|
||||
case "initialize":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": map[string]any{"tools": map[string]any{"listChanged": false}},
|
||||
"serverInfo": map[string]any{"name": "doc", "version": "1.0.0"},
|
||||
},
|
||||
})
|
||||
case "notifications/initialized":
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
case "tools/list":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"tools": []map[string]any{{
|
||||
"name": "search_documents",
|
||||
"title": "Search",
|
||||
"description": "Search documents",
|
||||
"inputSchema": map[string]any{"type": "object"},
|
||||
}},
|
||||
},
|
||||
})
|
||||
case "tools/call":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"content": map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_LOW_RISK_NO_PERMISSION",
|
||||
"data": map[string]any{"requiredScopes": []any{}},
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
hookCalled := false
|
||||
sentinelMsg := "hook-intercepted-PAT"
|
||||
edition.Override(&edition.Hooks{
|
||||
ClassifyToolResult: func(content map[string]any) error {
|
||||
if code, ok := content["code"].(string); ok && strings.Contains(code, "PAT") {
|
||||
hookCalled = true
|
||||
return fmt.Errorf("%s", sentinelMsg)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(func() { edition.Override(&edition.Hooks{}) })
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.URL, false))
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetOut(&bytes.Buffer{})
|
||||
cmd.SetErr(&bytes.Buffer{})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want hook sentinel error")
|
||||
}
|
||||
if !hookCalled {
|
||||
t.Fatal("ClassifyToolResult hook was not called")
|
||||
}
|
||||
if !strings.Contains(err.Error(), sentinelMsg) {
|
||||
t.Fatalf("error = %q, want hook sentinel %q (not generic business error)", err.Error(), sentinelMsg)
|
||||
}
|
||||
if strings.Contains(err.Error(), "business_error") || strings.Contains(err.Error(), "mcp_tool_error") {
|
||||
t.Fatalf("error = %q, should NOT contain generic framework error category", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
|
||||
+273
-18
@@ -19,7 +19,9 @@ import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
@@ -27,10 +29,24 @@ import (
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILL_API_HOST",
|
||||
Category: configmeta.CategoryNetwork,
|
||||
Description: "覆盖 Skill API 地址",
|
||||
DefaultValue: "https://mcp.dingtalk.com",
|
||||
Example: "https://custom-mcp.example.com",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// legacySkillAPIHost is the legacy skill market host used by the old cli.
|
||||
legacySkillAPIHost = "https://mcp.dingtalk.com"
|
||||
// skillDownloadEndpoint is the API endpoint for downloading skills.
|
||||
skillDownloadEndpoint = "https://aihub.dingtalk.com/cli/download"
|
||||
// skillDownloadTimeout is the timeout for skill download operations.
|
||||
@@ -51,6 +67,22 @@ type downloadSkillResult struct {
|
||||
FileName string `json:"fileName"`
|
||||
}
|
||||
|
||||
// findSkillsResponse represents the legacy skill search API response.
|
||||
type findSkillsResponse struct {
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result []CliSkillDTO `json:"result,omitempty"`
|
||||
}
|
||||
|
||||
// CliSkillDTO mirrors the old cli response payload for `skill search`.
|
||||
type CliSkillDTO struct {
|
||||
SkillID string `json:"skillId"`
|
||||
Name string `json:"name"`
|
||||
Desc string `json:"desc"`
|
||||
Icon string `json:"icon"`
|
||||
}
|
||||
|
||||
// agentSkillPaths maps target names to their relative skill installation paths.
|
||||
// These paths are relative to the user's home directory.
|
||||
var agentSkillPaths = map[string]string{
|
||||
@@ -75,7 +107,7 @@ func buildSkillCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "skill",
|
||||
Short: "技能管理",
|
||||
Long: "管理钉钉技能市场的技能。支持下载和安装技能到指定 Agent 目录。",
|
||||
Long: "管理钉钉技能市场的技能。支持搜索、下载与安装到指定 Agent 目录。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
@@ -84,13 +116,61 @@ func buildSkillCommand() *cobra.Command {
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(newSkillAddCommand())
|
||||
cmd.AddCommand(
|
||||
newSkillInstallCommand(),
|
||||
newSkillGetCommand(),
|
||||
newSkillSearchCommand(),
|
||||
newSkillFindHintCommand(),
|
||||
newSkillAddHintCommand(),
|
||||
)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddCommand() *cobra.Command {
|
||||
func newSkillGetCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "add <skillId> <target>",
|
||||
Use: "get",
|
||||
Short: "获取技能压缩文件",
|
||||
Long: "从服务端下载技能包到本地临时目录。命令执行成功后会输出临时目录路径,供调用方使用。",
|
||||
Example: " dws skill get --skill-id <skillId>",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillGet,
|
||||
}
|
||||
cmd.Flags().String("skill-id", "", "技能 ID(必填)")
|
||||
_ = cmd.MarkFlagRequired("skill-id")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillSearchCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "search",
|
||||
Short: "从钉钉技能市场搜索技能",
|
||||
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
|
||||
Example: " dws skill search --query 关键词",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillFind,
|
||||
}
|
||||
cmd.Flags().String("query", "", "搜索关键词(必填)")
|
||||
_ = cmd.MarkFlagRequired("query")
|
||||
cmd.Flags().String("scopes", "", "查询范围,空格分隔。备选值:DingtalkMarket(钉钉市场)、OrgInternal(企业内部)。为空默认查市场技能")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillFindHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "find",
|
||||
Short: "兼容旧用法,提示使用 skill search",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill search --query <关键词>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newSkillInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install <skillId> <target>",
|
||||
Short: "下载并安装技能到指定目录",
|
||||
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
|
||||
|
||||
@@ -107,9 +187,9 @@ func newSkillAddCommand() *cobra.Command {
|
||||
. -> 当前目录
|
||||
|
||||
示例:
|
||||
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
dws skill install skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill install skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill install skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillAdd,
|
||||
@@ -118,6 +198,96 @@ func newSkillAddCommand() *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "add",
|
||||
Short: "兼容旧用法,提示使用 skill install",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill install <skillId> <target>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/install?skillId=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(skillID)))
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "⬇️ 下载技能包...")
|
||||
|
||||
tmpDir, err := downloadSkillToTmpDir(cmd.Context(), apiURL, accessToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), tmpDir)
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
keyword, _ := cmd.Flags().GetString("query")
|
||||
scopes, _ := cmd.Flags().GetString("scopes")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
|
||||
if scopes != "" {
|
||||
apiURL += "&scopes=" + url.QueryEscape(scopes)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
var result findSkillsResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to parse search response: %v", err))
|
||||
}
|
||||
if !result.Success {
|
||||
errMsg := strings.TrimSpace(result.ErrorMsg)
|
||||
if errMsg == "" {
|
||||
errMsg = strings.TrimSpace(result.ErrorCode)
|
||||
}
|
||||
if errMsg == "" {
|
||||
errMsg = "unknown error"
|
||||
}
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %s", errMsg))
|
||||
}
|
||||
|
||||
if len(result.Result) == 0 {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "未找到匹配的技能")
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, skill := range result.Result {
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "SkillID: %s\n", skill.SkillID)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Name: %s\n", skill.Name)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Desc: %s\n", skill.Desc)
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "---")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
skillID := strings.TrimSpace(args[0])
|
||||
target := strings.TrimSpace(args[1])
|
||||
@@ -132,13 +302,9 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
|
||||
}
|
||||
|
||||
// Load auth token
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(cmd.Context(), skillDownloadTimeout)
|
||||
@@ -148,7 +314,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
|
||||
// Step 1: Get download URL from API
|
||||
fmt.Fprintf(w, "正在获取技能信息...\n")
|
||||
downloadResp, err := fetchSkillDownloadInfo(ctx, tokenData.AccessToken, skillID)
|
||||
downloadResp, err := fetchSkillDownloadInfo(ctx, accessToken, skillID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -189,6 +355,33 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSkillAccessToken() (string, error) {
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return "", skillAuthError()
|
||||
}
|
||||
return tokenData.AccessToken, nil
|
||||
}
|
||||
|
||||
func skillAuthError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return apperrors.NewAuth("认证信息已失效",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("请先完成钉钉账号登录后重试"))
|
||||
}
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
}
|
||||
|
||||
func skillAPIHost() string {
|
||||
if override := strings.TrimSpace(os.Getenv("DWS_SKILL_API_HOST")); override != "" {
|
||||
return strings.TrimRight(override, "/")
|
||||
}
|
||||
return legacySkillAPIHost
|
||||
}
|
||||
|
||||
// resolveSkillTargetPath resolves the target argument to an absolute path.
|
||||
func resolveSkillTargetPath(target string) (string, error) {
|
||||
target = strings.TrimSpace(target)
|
||||
@@ -236,9 +429,7 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return nil, apperrors.NewAuth("authentication failed. Please run 'dws auth login' to refresh your token",
|
||||
apperrors.WithHint("请执行 'dws auth login' 重新登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
return nil, skillAuthError()
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
@@ -259,6 +450,70 @@ func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func downloadSkillToTmpDir(ctx context.Context, apiURL, accessToken string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: skillDownloadTimeout}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill package: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
tmpDir, err := os.MkdirTemp("", "dws-skill-*")
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp dir: %v", err))
|
||||
}
|
||||
|
||||
filename := filenameFromDisposition(resp.Header.Get("Content-Disposition"))
|
||||
destPath := filepath.Join(tmpDir, filename)
|
||||
file, err := os.Create(destPath)
|
||||
if err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
if _, err := io.Copy(file, resp.Body); err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
|
||||
}
|
||||
return tmpDir, nil
|
||||
}
|
||||
|
||||
func filenameFromDisposition(cd string) string {
|
||||
if cd != "" {
|
||||
if _, params, err := mime.ParseMediaType(cd); err == nil {
|
||||
if name := strings.TrimSpace(params["filename"]); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
return "skill.zip"
|
||||
}
|
||||
|
||||
func parseLegacySkillAPIError(resp *http.Response) error {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusUnauthorized:
|
||||
return skillAuthError()
|
||||
case http.StatusBadRequest:
|
||||
return apperrors.NewValidation("request parameters are invalid")
|
||||
case http.StatusNotFound:
|
||||
return apperrors.NewValidation("skill does not exist or corresponding file was not found")
|
||||
default:
|
||||
return apperrors.NewAPI(fmt.Sprintf("skill API returned HTTP %d", resp.StatusCode),
|
||||
apperrors.WithRetryable(resp.StatusCode >= 500))
|
||||
}
|
||||
}
|
||||
|
||||
// downloadSkillFile downloads the skill zip file to a temporary location.
|
||||
func downloadSkillFile(ctx context.Context, downloadURL, fileName string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
|
||||
|
||||
@@ -319,7 +319,7 @@ func TestExtractSkillZipPreventZipSlip(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandValidation(t *testing.T) {
|
||||
func TestSkillInstallCommandValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
@@ -328,19 +328,19 @@ func TestSkillAddCommandValidation(t *testing.T) {
|
||||
}{
|
||||
{
|
||||
name: "missing arguments",
|
||||
args: []string{"skill", "add"},
|
||||
args: []string{"skill", "install"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "missing target",
|
||||
args: []string{"skill", "add", "skill-123"},
|
||||
args: []string{"skill", "install", "skill-123"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "too many arguments",
|
||||
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
|
||||
args: []string{"skill", "install", "skill-123", "qoder", "extra"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
@@ -366,7 +366,7 @@ func TestSkillAddCommandValidation(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
func TestSkillInstallInvalidTarget(t *testing.T) {
|
||||
// Setup: Create config directory with valid token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
@@ -380,11 +380,11 @@ func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to save token data: %v", err)
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
|
||||
cmd.SetArgs([]string{"skill", "install", "skill-123", "invalid-target"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -399,7 +399,7 @@ func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
func TestSkillInstallRequiresAuth(t *testing.T) {
|
||||
// Setup: Create config directory without token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
@@ -411,7 +411,7 @@ func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
|
||||
cmd.SetArgs([]string{"skill", "install", "skill-123", "qoder"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -563,14 +563,16 @@ func TestSkillCommandHelp(t *testing.T) {
|
||||
if !strings.Contains(output, "技能") {
|
||||
t.Errorf("help should mention '技能', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "add") {
|
||||
t.Errorf("help should mention 'add' subcommand, got: %s", output)
|
||||
for _, subcmd := range []string{"install", "search", "get"} {
|
||||
if !strings.Contains(output, subcmd) {
|
||||
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandHelp(t *testing.T) {
|
||||
func TestSkillInstallCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "--help"})
|
||||
cmd.SetArgs([]string{"skill", "install", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -590,6 +592,56 @@ func TestSkillAddCommandHelp(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillGetCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "get"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillSearchCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "search"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillFindHintCommand(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "find"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), "dws skill search --query") {
|
||||
t.Fatalf("output = %q, want legacy hint", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSkillFileSuccess(t *testing.T) {
|
||||
// Create a mock server that returns a zip file
|
||||
expectedContent := []byte("fake zip content")
|
||||
|
||||
@@ -0,0 +1,119 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
const stdioEndpointScheme = "stdio://"
|
||||
|
||||
var (
|
||||
stdioMu sync.RWMutex
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
)
|
||||
|
||||
// RegisterStdioClient stores a StdioClient keyed by its canonical product ID
|
||||
// (the CLI.ID used in the server descriptor). The runner looks up this client
|
||||
// when a stdio:// endpoint is resolved at execution time.
|
||||
func RegisterStdioClient(productID string, client *transport.StdioClient) {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
stdioClients[productID] = client
|
||||
}
|
||||
|
||||
// LookupStdioClient returns the StdioClient registered for the given product ID.
|
||||
// The productID can be either the full key (pluginName/serverKey) or just the serverKey.
|
||||
// This supports backward compatibility with existing CanonicalProduct values.
|
||||
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
|
||||
stdioMu.RLock()
|
||||
defer stdioMu.RUnlock()
|
||||
// Try exact match first
|
||||
if c, ok := stdioClients[productID]; ok {
|
||||
return c, true
|
||||
}
|
||||
// If not found, try matching by serverKey suffix (for backward compatibility)
|
||||
for id, c := range stdioClients {
|
||||
if idx := strings.LastIndex(id, "/"); idx >= 0 {
|
||||
if id[idx+1:] == productID {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// StdioEndpoint returns a virtual endpoint URL for a stdio-based MCP server.
|
||||
// Format: stdio://{pluginName}/{serverKey}
|
||||
func StdioEndpoint(pluginName, serverKey string) string {
|
||||
return stdioEndpointScheme + pluginName + "/" + serverKey
|
||||
}
|
||||
|
||||
// IsStdioEndpoint returns true if the endpoint uses the stdio:// scheme.
|
||||
func IsStdioEndpoint(endpoint string) bool {
|
||||
return strings.HasPrefix(endpoint, stdioEndpointScheme)
|
||||
}
|
||||
|
||||
// StopAllStdioClients stops all registered stdio clients.
|
||||
// This should be called on program exit to terminate child processes.
|
||||
func StopAllStdioClients() {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
for id, client := range stdioClients {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
}
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
}
|
||||
|
||||
// StopStdioClient stops a specific stdio client by product ID.
|
||||
// Returns true if the client was found and stopped, false otherwise.
|
||||
func StopStdioClient(productID string) bool {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
client, ok := stdioClients[productID]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", productID, "error", err)
|
||||
}
|
||||
delete(stdioClients, productID)
|
||||
return true
|
||||
}
|
||||
|
||||
// StopStdioClientsByPlugin stops all stdio clients belonging to a plugin.
|
||||
// The productID format is "pluginName/serverKey". This function stops all
|
||||
// clients whose productID has the given pluginName prefix.
|
||||
func StopStdioClientsByPlugin(pluginName string) int {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
prefix := pluginName + "/"
|
||||
count := 0
|
||||
for id, client := range stdioClients {
|
||||
if len(id) > len(prefix) && id[:len(prefix)] == prefix {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
delete(stdioClients, id)
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestStdioEndpoint(t *testing.T) {
|
||||
endpoint := StdioEndpoint("hello-plugin", "hello")
|
||||
want := "stdio://hello-plugin/hello"
|
||||
if endpoint != want {
|
||||
t.Errorf("StdioEndpoint() = %q, want %q", endpoint, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsStdioEndpoint(t *testing.T) {
|
||||
tests := []struct {
|
||||
endpoint string
|
||||
want bool
|
||||
}{
|
||||
{"stdio://hello-plugin/hello", true},
|
||||
{"stdio://conference/local", true},
|
||||
{"https://mcp.dingtalk.com", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := IsStdioEndpoint(tt.endpoint); got != tt.want {
|
||||
t.Errorf("IsStdioEndpoint(%q) = %v, want %v", tt.endpoint, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioClientRegistry(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
stdioMu.Lock()
|
||||
delete(stdioClients, "test-product")
|
||||
stdioMu.Unlock()
|
||||
}()
|
||||
|
||||
// Initially not found
|
||||
if _, ok := LookupStdioClient("test-product"); ok {
|
||||
t.Error("expected LookupStdioClient to return false for unregistered product")
|
||||
}
|
||||
|
||||
// Register a client
|
||||
client := transport.NewStdioClient("echo", nil, nil)
|
||||
RegisterStdioClient("test-product", client)
|
||||
|
||||
// Now should be found
|
||||
got, ok := LookupStdioClient("test-product")
|
||||
if !ok {
|
||||
t.Fatal("expected LookupStdioClient to return true after registration")
|
||||
}
|
||||
if got != client {
|
||||
t.Error("LookupStdioClient returned different client instance")
|
||||
}
|
||||
}
|
||||
+215
-11
@@ -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
@@ -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 '框架开销'")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 ""
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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 }
|
||||
|
||||
@@ -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
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -18,8 +18,34 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_ID",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppKey (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CHANNEL",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "第三方渠道编码 (channelCode),如 Qoderwork",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// AuthorizeURL is the DingTalk OAuth authorization page.
|
||||
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
|
||||
@@ -58,6 +84,14 @@ const (
|
||||
// DeviceGrantType is the grant_type value defined by RFC 8628.
|
||||
DeviceGrantType = "urn:ietf:params:oauth:grant-type:device_code"
|
||||
|
||||
// Terminal API base URL for developer settings page.
|
||||
DefaultTerminalBaseURL = "https://open-dev.dingtalk.com"
|
||||
// DevicePollPath is the device flow polling path (used with MCP base URL).
|
||||
DevicePollPath = "/cli/oauth/device/poll"
|
||||
|
||||
// DeveloperSettingsPath is the path to the organization developer settings page.
|
||||
DeveloperSettingsPath = "/fe/old#/developerSettings"
|
||||
|
||||
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
|
||||
LogoutContinueURL = "https://login.dingtalk.com"
|
||||
|
||||
@@ -74,6 +108,19 @@ const (
|
||||
MCPRevokeTokenPath = "/oauth2/revokeToken"
|
||||
)
|
||||
|
||||
// GetTerminalBaseURL returns the terminal base URL with priority:
|
||||
// 1. ~/.dws/terminal_url file content (for pre-release environment)
|
||||
// 2. Default value (https://open-dev.dingtalk.com)
|
||||
func GetTerminalBaseURL() string {
|
||||
return config.GetTerminalBaseURL()
|
||||
}
|
||||
|
||||
// GetDeveloperSettingsURL returns the full URL to the organization developer
|
||||
// settings page, derived from the terminal base URL.
|
||||
func GetDeveloperSettingsURL() string {
|
||||
return config.GetDeveloperSettingsURL()
|
||||
}
|
||||
|
||||
// GetMCPBaseURL returns the MCP base URL with priority:
|
||||
// 1. ~/.dws/mcp_url file content (for pre-release environment)
|
||||
// 2. Default value (https://mcp.dingtalk.com)
|
||||
@@ -110,7 +157,7 @@ func SetClientIDFromMCP(id string) {
|
||||
func IsClientIDFromMCP() bool {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return clientIDFromMCP
|
||||
return clientIDFromMCP || edition.Get().AuthClientFromMCP
|
||||
}
|
||||
|
||||
// GetUserAccessTokenURL returns the appropriate token exchange URL.
|
||||
@@ -189,6 +236,9 @@ func ClientID() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
if id := edition.Get().AuthClientID; id != "" {
|
||||
return id
|
||||
}
|
||||
// Try loading from persisted app config
|
||||
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
|
||||
return id
|
||||
|
||||
@@ -21,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 (
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
+952
-190
File diff suppressed because it is too large
Load Diff
@@ -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()
|
||||
|
||||
|
||||
@@ -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
@@ -20,7 +20,11 @@ import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// TokenData holds the OAuth token set persisted to disk.
|
||||
@@ -61,44 +65,88 @@ func (t *TokenData) HasPersistentCode() bool {
|
||||
return t != nil && t.PersistentCode != ""
|
||||
}
|
||||
|
||||
// SaveTokenData saves TokenData to the platform keychain.
|
||||
// Uses the new keychain-based storage with random master key for better security.
|
||||
const tokenJSONFile = "token.json"
|
||||
|
||||
// TokenMarker is a lightweight file the host application reads to detect
|
||||
// whether the CLI has a valid token without accessing the keychain.
|
||||
type TokenMarker struct {
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// WriteTokenMarker writes a token.json marker containing only an updated_at
|
||||
// timestamp. The host application uses this file's presence and mtime to
|
||||
// decide whether it needs to trigger a new auth exchange.
|
||||
func WriteTokenMarker(configDir string) error {
|
||||
marker := TokenMarker{UpdatedAt: time.Now().Format(time.RFC3339)}
|
||||
data, _ := json.MarshalIndent(marker, "", " ")
|
||||
if err := os.MkdirAll(configDir, 0o700); err != nil {
|
||||
return err
|
||||
}
|
||||
tmp := filepath.Join(configDir, tokenJSONFile+".tmp")
|
||||
if err := os.WriteFile(tmp, data, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmp, filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// DeleteTokenMarker removes the token.json marker file.
|
||||
func DeleteTokenMarker(configDir string) error {
|
||||
return os.Remove(filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to the default keychain-based storage.
|
||||
func SaveTokenData(configDir string, data *TokenData) error {
|
||||
if h := edition.Get(); h.SaveToken != nil {
|
||||
jsonData, err := json.MarshalIndent(data, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling token data for hook: %w", err)
|
||||
}
|
||||
return h.SaveToken(configDir, jsonData)
|
||||
}
|
||||
return SaveTokenDataKeychain(data)
|
||||
}
|
||||
|
||||
// LoadTokenData reads TokenData from the platform keychain.
|
||||
// On first call, it attempts to migrate legacy .data file if present.
|
||||
// LoadTokenData reads TokenData. When an edition hook (LoadToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain with legacy .data migration.
|
||||
func LoadTokenData(configDir string) (*TokenData, error) {
|
||||
// Try loading from new keychain first
|
||||
if h := edition.Get(); h.LoadToken != nil {
|
||||
jsonData, err := h.LoadToken(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var td TokenData
|
||||
if err := json.Unmarshal(jsonData, &td); err != nil {
|
||||
return nil, fmt.Errorf("parsing token data from hook: %w", err)
|
||||
}
|
||||
return &td, nil
|
||||
}
|
||||
|
||||
// Default: keychain with legacy .data migration
|
||||
if TokenDataExistsKeychain() {
|
||||
return LoadTokenDataKeychain()
|
||||
}
|
||||
|
||||
// Fallback: try legacy .data file and migrate
|
||||
data, err := LoadSecureTokenData(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Migrate to keychain for future use
|
||||
if err := SaveTokenDataKeychain(data); err == nil {
|
||||
// Successfully migrated, delete legacy file
|
||||
_ = DeleteSecureData(configDir)
|
||||
}
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes token data from both keychain and legacy storage.
|
||||
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain + legacy cleanup.
|
||||
func DeleteTokenData(configDir string) error {
|
||||
// Delete from keychain
|
||||
if h := edition.Get(); h.DeleteToken != nil {
|
||||
return h.DeleteToken(configDir)
|
||||
}
|
||||
keychainErr := DeleteTokenDataKeychain()
|
||||
|
||||
// Also clean up any legacy .data file
|
||||
legacyErr := DeleteSecureData(configDir)
|
||||
|
||||
// Return keychain error if any, otherwise legacy error
|
||||
if keychainErr != nil {
|
||||
return keychainErr
|
||||
}
|
||||
|
||||
@@ -20,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, °raded) {
|
||||
fmt.Fprintf(cmd.ErrOrStderr(), "hint: %s\n", degraded.Hint)
|
||||
payload := map[string]any{
|
||||
"kind": "schema",
|
||||
"count": 0,
|
||||
"products": []any{},
|
||||
"degraded": true,
|
||||
"reason": string(degraded.Reason),
|
||||
"hint": degraded.Hint,
|
||||
}
|
||||
return output.WriteFiltered(
|
||||
cmd.OutOrStdout(),
|
||||
output.ResolveFormat(cmd, output.FormatJSON),
|
||||
payload,
|
||||
output.ResolveFields(cmd),
|
||||
output.ResolveJQ(cmd),
|
||||
)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -213,6 +233,27 @@ func newProductCommand(product ir.CanonicalProduct, runner executor.Runner, engi
|
||||
for _, tool := range product.Tools {
|
||||
cmd.AddCommand(newToolCommand(product, tool, runner, engine))
|
||||
}
|
||||
|
||||
// Register phase: notify the pipeline that a product and its
|
||||
// tools have been added to the command tree. This runs once at
|
||||
// startup (not per-request) and enables handlers to inspect or
|
||||
// enrich the registered command surface.
|
||||
if engine != nil && engine.HasHandlers(pipeline.Register) {
|
||||
pctx := &pipeline.Context{
|
||||
Command: product.ID,
|
||||
}
|
||||
// Best-effort — registration errors are logged but do not
|
||||
// prevent the CLI from starting.
|
||||
if pipeErr := engine.RunPhase(pipeline.Register, pctx); pipeErr != nil {
|
||||
slog.Debug("pipeline register phase", "product", product.ID, "error", pipeErr)
|
||||
} else {
|
||||
slog.Debug("pipeline register",
|
||||
"product", product.ID,
|
||||
"tool_count", len(product.Tools),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -368,6 +409,16 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
params = pctx.Params
|
||||
for _, c := range pctx.Corrections {
|
||||
slog.Debug("pipeline correction",
|
||||
"phase", "post-parse",
|
||||
"handler", c.Handler,
|
||||
"kind", c.Kind,
|
||||
"field", c.Field,
|
||||
"original", c.Original,
|
||||
"corrected", c.Corrected,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
if err := ValidateInputSchema(params, tool.InputSchema); err != nil {
|
||||
@@ -392,6 +443,10 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
params = pctx.Params
|
||||
slog.Debug("pipeline pre-request",
|
||||
"command", tool.CanonicalPath,
|
||||
"param_count", len(params),
|
||||
)
|
||||
}
|
||||
|
||||
invocation := executor.NewInvocation(product, tool, params)
|
||||
@@ -414,6 +469,10 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
return pipeErr
|
||||
}
|
||||
result.Response = pctx.Response
|
||||
slog.Debug("pipeline post-response",
|
||||
"command", tool.CanonicalPath,
|
||||
"has_response", result.Response != nil,
|
||||
)
|
||||
}
|
||||
|
||||
if warning := lifecycleWarning(product); warning != "" {
|
||||
|
||||
@@ -1033,6 +1033,83 @@ func newTestMCPCommand(t *testing.T, catalog ir.Catalog, runner executor.Runner)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func TestSchemaCommandOutputsDegradedOnUnauthenticated(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
degradedErr := &CatalogDegraded{
|
||||
Reason: DegradedUnauthenticated,
|
||||
Hint: "未登录,无法发现 MCP 服务。请先执行: dws auth login",
|
||||
}
|
||||
cmd := NewSchemaCommand(errorLoader{err: degradedErr})
|
||||
|
||||
var out, errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v, want nil (degraded handled gracefully)", err)
|
||||
}
|
||||
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if payload["degraded"] != true {
|
||||
t.Fatalf("payload[degraded] = %v, want true", payload["degraded"])
|
||||
}
|
||||
if payload["reason"] != "unauthenticated" {
|
||||
t.Fatalf("payload[reason] = %v, want unauthenticated", payload["reason"])
|
||||
}
|
||||
if payload["count"] != float64(0) {
|
||||
t.Fatalf("payload[count] = %v, want 0", payload["count"])
|
||||
}
|
||||
if !strings.Contains(errOut.String(), "hint:") {
|
||||
t.Fatalf("stderr = %q, want hint message", errOut.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaCommandOutputsDegradedOnMarketUnreachable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
degradedErr := &CatalogDegraded{
|
||||
Reason: DegradedMarketUnreachable,
|
||||
Hint: "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络",
|
||||
}
|
||||
cmd := NewSchemaCommand(errorLoader{err: degradedErr})
|
||||
|
||||
var out, errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v, want nil", err)
|
||||
}
|
||||
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &payload); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if payload["reason"] != "market_unreachable" {
|
||||
t.Fatalf("payload[reason] = %v, want market_unreachable", payload["reason"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaCommandPropagatesNonDegradedError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
wantErr := errors.New("unexpected failure")
|
||||
cmd := NewSchemaCommand(errorLoader{err: wantErr})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if !errors.Is(err, wantErr) {
|
||||
t.Fatalf("Execute() error = %v, want %v", err, wantErr)
|
||||
}
|
||||
}
|
||||
|
||||
type errorLoader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
+102
-9
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -20,6 +20,8 @@ import (
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
// Category represents a stable error class with a documented exit code.
|
||||
@@ -198,12 +200,31 @@ func NewInternal(message string, opts ...Option) error {
|
||||
return newError(CategoryInternal, message, opts...)
|
||||
}
|
||||
|
||||
// ExitCoder is implemented by errors that provide their own exit code.
|
||||
// Edition-specific error types (e.g. PATError, CLIError) implement this
|
||||
// so the framework can resolve exit codes without importing edition packages.
|
||||
type ExitCoder interface {
|
||||
ExitCode() int
|
||||
}
|
||||
|
||||
// RawStderrError is implemented by errors that must output raw content
|
||||
// directly to stderr, bypassing all CLI formatting (e.g. "Error:" prefix).
|
||||
// PAT authorization errors use this to pass JSON through to the desktop runtime.
|
||||
type RawStderrError interface {
|
||||
error
|
||||
RawStderr() string
|
||||
}
|
||||
|
||||
// ExitCode maps any error to a stable exit code.
|
||||
func ExitCode(err error) int {
|
||||
var typed *Error
|
||||
if stderrors.As(err, &typed) {
|
||||
return typed.ExitCode()
|
||||
}
|
||||
var ec ExitCoder
|
||||
if stderrors.As(err, &ec) {
|
||||
return ec.ExitCode()
|
||||
}
|
||||
return 5
|
||||
}
|
||||
|
||||
@@ -257,7 +278,7 @@ func PrintJSON(w io.Writer, err error) error {
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
|
||||
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
|
||||
errorPayload["action_url"] = config.GetDeveloperSettingsURL()
|
||||
}
|
||||
}
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
@@ -323,7 +344,7 @@ func PrintHumanAt(w io.Writer, err error, v Verbosity) error {
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
|
||||
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
|
||||
lines = append(lines, "Action: 开启地址: "+config.GetDeveloperSettingsURL())
|
||||
}
|
||||
|
||||
if len(typed.Actions) > 0 {
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package errors
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type stubExitCoder struct{ code int }
|
||||
|
||||
func (s *stubExitCoder) Error() string { return "stub" }
|
||||
func (s *stubExitCoder) ExitCode() int { return s.code }
|
||||
|
||||
type stubRawStderr struct{ raw string }
|
||||
|
||||
func (s *stubRawStderr) Error() string { return s.raw }
|
||||
func (s *stubRawStderr) RawStderr() string { return s.raw }
|
||||
|
||||
func TestExitCode_ExitCoderInterface(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
err error
|
||||
want int
|
||||
}{
|
||||
{"exit code 4 via interface", &stubExitCoder{code: 4}, 4},
|
||||
{"exit code 1 via interface", &stubExitCoder{code: 1}, 1},
|
||||
{"framework Error takes precedence", NewAPI("api"), 1},
|
||||
{"plain error falls back to 5", stderrors.New("plain"), 5},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
if got := ExitCode(tc.err); got != tc.want {
|
||||
t.Errorf("ExitCode() = %d, want %d", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExitCode_WrappedExitCoder(t *testing.T) {
|
||||
t.Parallel()
|
||||
wrapped := stderrors.Join(stderrors.New("context"), &stubExitCoder{code: 4})
|
||||
if got := ExitCode(wrapped); got != 4 {
|
||||
t.Errorf("ExitCode(wrapped) = %d, want 4", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRawStderrError_Interface(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &stubRawStderr{raw: `{"code":"PAT_LOW_RISK_NO_PERMISSION"}`}
|
||||
var raw RawStderrError
|
||||
if !stderrors.As(err, &raw) {
|
||||
t.Fatal("expected errors.As to match RawStderrError")
|
||||
}
|
||||
if !strings.Contains(raw.RawStderr(), "PAT_LOW_RISK_NO_PERMISSION") {
|
||||
t.Errorf("RawStderr() = %q, want PAT code", raw.RawStderr())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package errors
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ExitCodePermission is the process exit code for PAT authorisation failures.
|
||||
const ExitCodePermission = 4
|
||||
|
||||
// PATError represents a PAT (Personal Action Token) authorization failure
|
||||
// that should be passed through to stderr as raw JSON without any CLI-layer
|
||||
// wrapping. The host application (e.g. RewindDesktop) parses the JSON to
|
||||
// display its own authorisation UI.
|
||||
type PATError struct {
|
||||
RawJSON string
|
||||
}
|
||||
|
||||
func (e *PATError) Error() string { return e.RawJSON }
|
||||
|
||||
// ExitCode returns the documented exit code for PAT permission errors (4).
|
||||
func (e *PATError) ExitCode() int { return ExitCodePermission }
|
||||
|
||||
// RawStderr returns the raw JSON to be written directly to stderr.
|
||||
func (e *PATError) RawStderr() string { return e.RawJSON }
|
||||
|
||||
// patNoPermissionCodes are PAT error codes that should be passed through
|
||||
// as transparent PATError without CLI-level wrapping.
|
||||
var patNoPermissionCodes = map[string]bool{
|
||||
"PAT_NO_PERMISSION": true,
|
||||
"PAT_LOW_RISK_NO_PERMISSION": true,
|
||||
"PAT_MEDIUM_RISK_NO_PERMISSION": true,
|
||||
"PAT_HIGH_RISK_NO_PERMISSION": true,
|
||||
}
|
||||
|
||||
// patAuthRequiredCodes are error codes that trigger the PAT authorization
|
||||
// flow (e.g. the server auto-created a CLI app and returned auth details).
|
||||
var patAuthRequiredCodes = map[string]bool{
|
||||
"AGENT_CODE_NOT_EXISTS": true,
|
||||
}
|
||||
|
||||
// IsPATError reports whether err is a *PATError.
|
||||
func IsPATError(err error) bool {
|
||||
_, ok := err.(*PATError)
|
||||
return ok
|
||||
}
|
||||
|
||||
// IsPATNoPermissionCode reports whether code is a known PAT permission error code.
|
||||
func IsPATNoPermissionCode(code string) bool {
|
||||
return patNoPermissionCodes[code]
|
||||
}
|
||||
|
||||
// ---- DWS gateway auth errors (shared between PAT & general auth) ----------
|
||||
|
||||
// dwsGatewayErrors is the set of DWS gateway-level auth error codes.
|
||||
var dwsGatewayErrors = map[string]bool{
|
||||
"DWS_SERVICE_UNAUTHORIZED": true,
|
||||
"DWS_AUTH_SERVICE_FAILED": true,
|
||||
}
|
||||
|
||||
// getDWSGatewayErrorCode extracts a DWS gateway error code from errBody
|
||||
// (supports both errorCode and error_code field names).
|
||||
func getDWSGatewayErrorCode(errBody map[string]any) (string, bool) {
|
||||
for _, key := range []string{"errorCode", "error_code"} {
|
||||
if code, ok := errBody[key].(string); ok && dwsGatewayErrors[code] {
|
||||
return code, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// isNotLoggedInError checks if the error body indicates missing authentication.
|
||||
func isNotLoggedInError(body map[string]any) bool {
|
||||
if errMsg, ok := body["error"].(string); ok {
|
||||
if strings.Contains(errMsg, "Missing service_id or access_key") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isBusinessError checks if a parsed JSON body represents a business-level error.
|
||||
func isBusinessError(body map[string]any) bool {
|
||||
if _, ok := body["error"].(string); ok {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(bool); ok && !v {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(string); ok && strings.EqualFold(v, "false") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ---- Classification functions -----------------------------------------------
|
||||
|
||||
// ClassifyToolResultContent checks a raw MCP tool result content map for
|
||||
// DWS gateway auth errors and PAT permission error codes. This is intended
|
||||
// for use as the edition.Hooks.ClassifyToolResult callback so the framework's
|
||||
// runner returns a typed error before its generic business-error classification.
|
||||
//
|
||||
// Check order: DWS gateway auth > PAT permission.
|
||||
func ClassifyToolResultContent(content map[string]any) error {
|
||||
if _, ok := getDWSGatewayErrorCode(content); ok {
|
||||
raw, _ := json.Marshal(content)
|
||||
return NewAuth(string(raw),
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClassifyMCPResponseText classifies a text response returned by an MCP tool call.
|
||||
// Returns a typed error for known gateway auth failures, PAT interceptions,
|
||||
// and business-level errors embedded in HTTP-200 JSON bodies.
|
||||
//
|
||||
// Check order: DWS gateway > PAT permission > generic business error.
|
||||
func ClassifyMCPResponseText(text string) error {
|
||||
var body map[string]any
|
||||
if json.Unmarshal([]byte(text), &body) != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, ok := getDWSGatewayErrorCode(body); ok {
|
||||
return NewAuth(text,
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
|
||||
if isNotLoggedInError(body) {
|
||||
return NewAuth("当前未登录",
|
||||
WithReason("not_configured"),
|
||||
WithHint(notLoggedInHint()),
|
||||
WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := body[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(body, code)}
|
||||
}
|
||||
}
|
||||
|
||||
if isBusinessError(body) {
|
||||
return NewAPI(text,
|
||||
WithReason("business_error"),
|
||||
WithHint(suggestForBusinessErrorText(body)),
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ---- Hints -----------------------------------------------------------------
|
||||
|
||||
func authExpiredHint() string {
|
||||
return "Re-authenticate: dws auth login"
|
||||
}
|
||||
|
||||
func notLoggedInHint() string {
|
||||
return "请先登录:dws auth login"
|
||||
}
|
||||
|
||||
func suggestForBusinessErrorText(body map[string]any) string {
|
||||
msg := ""
|
||||
if v, ok := body["errorMsg"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["message"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["error"].(string); ok {
|
||||
msg = v
|
||||
}
|
||||
switch {
|
||||
case strings.Contains(msg, "搜索内容不能为空"):
|
||||
return "请提供非空搜索关键词: dws doc search --query \"关键词\""
|
||||
case strings.Contains(msg, "User has no permission to access this email"):
|
||||
return "请确认邮箱地址正确,查看可用邮箱: dws mail mailbox list"
|
||||
case strings.Contains(msg, "频率超限") || strings.Contains(msg, "rate limit"):
|
||||
return "API rate limit exceeded, wait a moment and retry"
|
||||
case strings.Contains(msg, "参数错误") || strings.Contains(msg, "param error"):
|
||||
return "Check input parameters. Use --help for available flags"
|
||||
default:
|
||||
return "MCP tool returned a business error; check parameters and refer to skill documentation."
|
||||
}
|
||||
}
|
||||
|
||||
// ---- PAT JSON helpers ------------------------------------------------------
|
||||
|
||||
var patTopLevelStrip = map[string]bool{
|
||||
"success": true, "code": true, "errorCode": true, "error_code": true,
|
||||
"message": true, "error": true, "trace_id": true, "class": true,
|
||||
}
|
||||
|
||||
func cleanPATJSON(body map[string]any, code string) string {
|
||||
out := map[string]any{
|
||||
"success": false,
|
||||
"code": code,
|
||||
}
|
||||
if data, ok := body["data"]; ok {
|
||||
out["data"] = stripClassFields(data)
|
||||
} else {
|
||||
fallback := map[string]any{}
|
||||
for k, v := range body {
|
||||
if !patTopLevelStrip[k] {
|
||||
fallback[k] = v
|
||||
}
|
||||
}
|
||||
if len(fallback) > 0 {
|
||||
out["data"] = stripClassFields(fallback)
|
||||
}
|
||||
}
|
||||
b, err := json.MarshalIndent(out, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Sprintf(`{"success":false,"code":"%s"}`, code)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// ---- Runner adapter functions ------------------------------------------------
|
||||
// These match the function signatures referenced by runner.go's PAT check
|
||||
// framework (ClassifyPatAuthCheck / AsPatAuthCheckError).
|
||||
|
||||
// ClassifyPatAuthCheck is the open-source fallback that checks a tool-call
|
||||
// Content map for PAT permission codes and auth-required codes. Returns a
|
||||
// non-nil *PATError when the content carries a recognised PAT/auth error.
|
||||
func ClassifyPatAuthCheck(content map[string]any) *PATError {
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok {
|
||||
if patNoPermissionCodes[code] || patAuthRequiredCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AsPatAuthCheckError extracts a *PATError from an error chain.
|
||||
func AsPatAuthCheckError(err error) *PATError {
|
||||
var patErr *PATError
|
||||
if stderrors.As(err, &patErr) {
|
||||
return patErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func stripClassFields(v any) any {
|
||||
switch val := v.(type) {
|
||||
case map[string]any:
|
||||
clean := make(map[string]any, len(val))
|
||||
for k, item := range val {
|
||||
if k == "class" {
|
||||
continue
|
||||
}
|
||||
clean[k] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
case []any:
|
||||
clean := make([]any, len(val))
|
||||
for i, item := range val {
|
||||
clean[i] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,521 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package errors
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// PATError basic behaviour
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPATError_Implements(t *testing.T) {
|
||||
t.Parallel()
|
||||
raw := `{"success":false,"code":"PAT_NO_PERMISSION"}`
|
||||
pe := &PATError{RawJSON: raw}
|
||||
|
||||
if pe.Error() != raw {
|
||||
t.Errorf("Error() = %q, want %q", pe.Error(), raw)
|
||||
}
|
||||
if pe.ExitCode() != ExitCodePermission {
|
||||
t.Errorf("ExitCode() = %d, want %d", pe.ExitCode(), ExitCodePermission)
|
||||
}
|
||||
if pe.RawStderr() != raw {
|
||||
t.Errorf("RawStderr() = %q, want %q", pe.RawStderr(), raw)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPATError_True(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PATError{RawJSON: "{}"}
|
||||
if !IsPATError(err) {
|
||||
t.Fatal("expected IsPATError to return true for *PATError")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPATError_False(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := stderrors.New("some other error")
|
||||
if IsPATError(err) {
|
||||
t.Fatal("expected IsPATError to return false for non-PATError")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// IsPATNoPermissionCode
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsPATNoPermissionCode(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
code string
|
||||
want bool
|
||||
}{
|
||||
{"PAT_NO_PERMISSION", true},
|
||||
{"PAT_LOW_RISK_NO_PERMISSION", true},
|
||||
{"PAT_MEDIUM_RISK_NO_PERMISSION", true},
|
||||
{"PAT_HIGH_RISK_NO_PERMISSION", true},
|
||||
{"AGENT_CODE_NOT_EXISTS", false},
|
||||
{"UNKNOWN_CODE", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
if got := IsPATNoPermissionCode(tc.code); got != tc.want {
|
||||
t.Errorf("IsPATNoPermissionCode(%q) = %v, want %v", tc.code, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// getDWSGatewayErrorCode
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetDWSGatewayErrorCode_ErrorCode(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED"}
|
||||
code, ok := getDWSGatewayErrorCode(body)
|
||||
if !ok || code != "DWS_SERVICE_UNAUTHORIZED" {
|
||||
t.Errorf("got (%q, %v), want (DWS_SERVICE_UNAUTHORIZED, true)", code, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_ErrorCodeUnderscore(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error_code": "DWS_AUTH_SERVICE_FAILED"}
|
||||
code, ok := getDWSGatewayErrorCode(body)
|
||||
if !ok || code != "DWS_AUTH_SERVICE_FAILED" {
|
||||
t.Errorf("got (%q, %v), want (DWS_AUTH_SERVICE_FAILED, true)", code, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_Unknown(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"errorCode": "SOME_OTHER_ERROR"}
|
||||
_, ok := getDWSGatewayErrorCode(body)
|
||||
if ok {
|
||||
t.Fatal("expected ok=false for unknown error code")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetDWSGatewayErrorCode_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{}
|
||||
_, ok := getDWSGatewayErrorCode(body)
|
||||
if ok {
|
||||
t.Fatal("expected ok=false for empty body")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// isNotLoggedInError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsNotLoggedInError_True(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "Missing service_id or access_key in request headers"}
|
||||
if !isNotLoggedInError(body) {
|
||||
t.Fatal("expected true for Missing service_id message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNotLoggedInError_False(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "something else happened"}
|
||||
if isNotLoggedInError(body) {
|
||||
t.Fatal("expected false for unrelated error message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsNotLoggedInError_NoErrorField(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"message": "Missing service_id or access_key"}
|
||||
if isNotLoggedInError(body) {
|
||||
t.Fatal("expected false when error field is absent")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// isBusinessError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsBusinessError_ErrorField(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"error": "some error message"}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when 'error' field is present")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessBoolFalse(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": false}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when success=false (bool)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessStringFalse(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": "False"}
|
||||
if !isBusinessError(body) {
|
||||
t.Fatal("expected true when success=\"False\" (string)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_SuccessTrue(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"success": true, "data": "ok"}
|
||||
if isBusinessError(body) {
|
||||
t.Fatal("expected false when success=true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsBusinessError_EmptyBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{"data": "hello"}
|
||||
if isBusinessError(body) {
|
||||
t.Fatal("expected false for body without error indicators")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyToolResultContent
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyToolResultContent_GatewayAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"errorCode": "DWS_SERVICE_UNAUTHORIZED", "message": "expired"}
|
||||
err := ClassifyToolResultContent(content)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error for gateway auth")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Category != CategoryAuth {
|
||||
t.Errorf("Category = %v, want %v", typed.Category, CategoryAuth)
|
||||
}
|
||||
if typed.Reason != "gateway_auth_expired" {
|
||||
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyToolResultContent_PATPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"data": map[string]any{"desc": "需要授权"},
|
||||
}
|
||||
err := ClassifyToolResultContent(content)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error for PAT permission")
|
||||
}
|
||||
var patErr *PATError
|
||||
if !stderrors.As(err, &patErr) {
|
||||
t.Fatalf("expected *PATError, got %T", err)
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain PAT_NO_PERMISSION, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyToolResultContent_NoError(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"success": true, "data": "ok"}
|
||||
if err := ClassifyToolResultContent(content); err != nil {
|
||||
t.Fatalf("expected nil error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyMCPResponseText
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyMCPResponseText_GatewayAuth(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"errorCode":"DWS_SERVICE_UNAUTHORIZED","message":"token expired"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "gateway_auth_expired" {
|
||||
t.Errorf("Reason = %q, want gateway_auth_expired", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_NotLoggedIn(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"error":"Missing service_id or access_key in headers"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "not_configured" {
|
||||
t.Errorf("Reason = %q, want not_configured", typed.Reason)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_PATPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"code":"PAT_HIGH_RISK_NO_PERMISSION","data":{"desc":"high risk"}}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var patErr *PATError
|
||||
if !stderrors.As(err, &patErr) {
|
||||
t.Fatalf("expected *PATError, got %T", err)
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_HIGH_RISK_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_BusinessError(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"success":false,"errorMsg":"搜索内容不能为空"}`
|
||||
err := ClassifyMCPResponseText(text)
|
||||
if err == nil {
|
||||
t.Fatal("expected non-nil error")
|
||||
}
|
||||
var typed *Error
|
||||
if !stderrors.As(err, &typed) {
|
||||
t.Fatalf("expected *Error, got %T", err)
|
||||
}
|
||||
if typed.Reason != "business_error" {
|
||||
t.Errorf("Reason = %q, want business_error", typed.Reason)
|
||||
}
|
||||
if !strings.Contains(typed.Hint, "搜索关键词") {
|
||||
t.Errorf("Hint should contain search suggestion, got: %s", typed.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_InvalidJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := "not json at all"
|
||||
if err := ClassifyMCPResponseText(text); err != nil {
|
||||
t.Fatalf("expected nil for invalid JSON, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyMCPResponseText_NoError(t *testing.T) {
|
||||
t.Parallel()
|
||||
text := `{"success":true,"data":"hello"}`
|
||||
if err := ClassifyMCPResponseText(text); err != nil {
|
||||
t.Fatalf("expected nil for success response, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ClassifyPatAuthCheck
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestClassifyPatAuthCheck_PATNoPermission(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"code": "PAT_NO_PERMISSION", "data": map[string]any{"flowId": "f1"}}
|
||||
patErr := ClassifyPatAuthCheck(content)
|
||||
if patErr == nil {
|
||||
t.Fatal("expected non-nil *PATError")
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("RawJSON should contain code, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_AgentCodeNotExists(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"errorCode": "AGENT_CODE_NOT_EXISTS", "data": map[string]any{"clientId": "c1"}}
|
||||
patErr := ClassifyPatAuthCheck(content)
|
||||
if patErr == nil {
|
||||
t.Fatal("expected non-nil *PATError for AGENT_CODE_NOT_EXISTS")
|
||||
}
|
||||
if !strings.Contains(patErr.RawJSON, "AGENT_CODE_NOT_EXISTS") {
|
||||
t.Errorf("RawJSON should contain AGENT_CODE_NOT_EXISTS, got: %s", patErr.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_NoMatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{"code": "SOME_BUSINESS_ERROR", "message": "oops"}
|
||||
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
|
||||
t.Fatalf("expected nil, got %v", patErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClassifyPatAuthCheck_EmptyContent(t *testing.T) {
|
||||
t.Parallel()
|
||||
content := map[string]any{}
|
||||
if patErr := ClassifyPatAuthCheck(content); patErr != nil {
|
||||
t.Fatalf("expected nil for empty content, got %v", patErr)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AsPatAuthCheckError
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAsPatAuthCheckError_Wrapped(t *testing.T) {
|
||||
t.Parallel()
|
||||
inner := &PATError{RawJSON: `{"code":"PAT_NO_PERMISSION"}`}
|
||||
wrapped := stderrors.Join(stderrors.New("context"), inner)
|
||||
got := AsPatAuthCheckError(wrapped)
|
||||
if got == nil {
|
||||
t.Fatal("expected non-nil *PATError from wrapped error")
|
||||
}
|
||||
if got.RawJSON != inner.RawJSON {
|
||||
t.Errorf("RawJSON = %q, want %q", got.RawJSON, inner.RawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAsPatAuthCheckError_NotPAT(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := stderrors.New("just a plain error")
|
||||
if got := AsPatAuthCheckError(err); got != nil {
|
||||
t.Fatalf("expected nil for non-PAT error, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// cleanPATJSON
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCleanPATJSON_WithData(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"data": map[string]any{
|
||||
"desc": "需要授权",
|
||||
"flowId": "f123",
|
||||
"class": "com.foo.Bar",
|
||||
},
|
||||
}
|
||||
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
|
||||
if !strings.Contains(result, "PAT_NO_PERMISSION") {
|
||||
t.Errorf("expected code in output, got: %s", result)
|
||||
}
|
||||
if !strings.Contains(result, "flowId") {
|
||||
t.Errorf("expected flowId in data, got: %s", result)
|
||||
}
|
||||
if strings.Contains(result, "class") {
|
||||
t.Errorf("expected class field to be stripped, got: %s", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanPATJSON_WithoutData(t *testing.T) {
|
||||
t.Parallel()
|
||||
body := map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_NO_PERMISSION",
|
||||
"message": "no permission",
|
||||
"extra": "value",
|
||||
}
|
||||
result := cleanPATJSON(body, "PAT_NO_PERMISSION")
|
||||
if !strings.Contains(result, "extra") {
|
||||
t.Errorf("expected extra field in fallback data, got: %s", result)
|
||||
}
|
||||
// Top-level stripped fields should not appear
|
||||
if strings.Contains(result, `"message"`) {
|
||||
t.Errorf("expected message to be stripped from top level, got: %s", result)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// stripClassFields
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStripClassFields_Map(t *testing.T) {
|
||||
t.Parallel()
|
||||
input := map[string]any{
|
||||
"name": "test",
|
||||
"class": "com.foo.Bar",
|
||||
"nested": map[string]any{
|
||||
"value": 42,
|
||||
"class": "com.baz.Qux",
|
||||
},
|
||||
}
|
||||
result := stripClassFields(input).(map[string]any)
|
||||
if _, ok := result["class"]; ok {
|
||||
t.Error("top-level class should be removed")
|
||||
}
|
||||
nested := result["nested"].(map[string]any)
|
||||
if _, ok := nested["class"]; ok {
|
||||
t.Error("nested class should be removed")
|
||||
}
|
||||
if nested["value"] != 42 {
|
||||
t.Errorf("nested value should be preserved, got %v", nested["value"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripClassFields_Array(t *testing.T) {
|
||||
t.Parallel()
|
||||
input := []any{
|
||||
map[string]any{"id": 1, "class": "Foo"},
|
||||
map[string]any{"id": 2},
|
||||
}
|
||||
result := stripClassFields(input).([]any)
|
||||
first := result[0].(map[string]any)
|
||||
if _, ok := first["class"]; ok {
|
||||
t.Error("class in array element should be removed")
|
||||
}
|
||||
if first["id"] != 1 {
|
||||
t.Error("other fields in array element should be preserved")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripClassFields_Scalar(t *testing.T) {
|
||||
t.Parallel()
|
||||
if stripClassFields("hello") != "hello" {
|
||||
t.Error("scalar string should pass through unchanged")
|
||||
}
|
||||
if stripClassFields(42) != 42 {
|
||||
t.Error("scalar int should pass through unchanged")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// suggestForBusinessErrorText
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSuggestForBusinessErrorText(t *testing.T) {
|
||||
t.Parallel()
|
||||
cases := []struct {
|
||||
body map[string]any
|
||||
contains string
|
||||
}{
|
||||
{map[string]any{"errorMsg": "搜索内容不能为空"}, "搜索关键词"},
|
||||
{map[string]any{"message": "User has no permission to access this email"}, "邮箱"},
|
||||
{map[string]any{"error": "频率超限"}, "rate limit"},
|
||||
{map[string]any{"errorMsg": "参数错误"}, "parameters"},
|
||||
{map[string]any{"error": "unknown"}, "business error"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
hint := suggestForBusinessErrorText(tc.body)
|
||||
if !strings.Contains(strings.ToLower(hint), strings.ToLower(tc.contains)) {
|
||||
t.Errorf("suggestForBusinessErrorText(%v) = %q, want to contain %q", tc.body, hint, tc.contains)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -20,9 +20,27 @@ import (
|
||||
"strings"
|
||||
|
||||
registryassets "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/registry"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILLS_PERSONAS_FILE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "覆盖内置 personas.yaml 的本地文件路径",
|
||||
Example: "/path/to/personas.yaml",
|
||||
Hidden: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILLS_RECIPES_FILE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "覆盖内置 recipes.yaml 的本地文件路径",
|
||||
Example: "/path/to/recipes.yaml",
|
||||
Hidden: true,
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
PersonaRegistryPathEnv = "DWS_SKILLS_PERSONAS_FILE"
|
||||
RecipeRegistryPathEnv = "DWS_SKILLS_RECIPES_FILE"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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."
|
||||
}
|
||||
|
||||
@@ -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 [命令路径] 查看完整说明。"
|
||||
}
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -0,0 +1,137 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package pat
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
|
||||
"github.com/fatih/color"
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
var validGrantTypes = map[string]bool{
|
||||
"once": true,
|
||||
"session": true,
|
||||
"permanent": true,
|
||||
}
|
||||
|
||||
func newChmodCommand(caller edition.ToolCaller) *cobra.Command {
|
||||
chmodCmd := &cobra.Command{
|
||||
Use: "chmod <scope>...",
|
||||
Short: "授予指定权限",
|
||||
Long: `授予指定 scope 的操作权限。
|
||||
|
||||
scope 格式: <product>.<entity>:<permission>
|
||||
例: aitable.record:read chat.group:write calendar.event:read
|
||||
|
||||
grantType 规则:
|
||||
once 一次性,执行一次后自动失效
|
||||
session 当前会话有效(默认),需要 --session-id
|
||||
permanent 永久有效`,
|
||||
Args: cobra.MinimumNArgs(1),
|
||||
Example: ` dws pat chmod aitable.record:read --agentCode agt-xxxx --grant-type session --session-id session-xxx
|
||||
dws pat chmod chat.message:list --grant-type once --agentCode agt-xxxx
|
||||
dws pat chmod aitable.record:read aitable.record:write --agentCode agt-xxxx --grant-type permanent`,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
agentCode, _ := cmd.Flags().GetString("agentCode")
|
||||
if agentCode == "" {
|
||||
return fmt.Errorf("flag --agentCode is required\n hint: dws pat chmod <scope>... --agentCode <id>")
|
||||
}
|
||||
scopes := args
|
||||
grantType, _ := cmd.Flags().GetString("grant-type")
|
||||
sessionID, _ := cmd.Flags().GetString("session-id")
|
||||
|
||||
if !validGrantTypes[grantType] {
|
||||
return fmt.Errorf("invalid --grant-type %q, must be one of: once, session, permanent", grantType)
|
||||
}
|
||||
|
||||
if grantType == "session" && sessionID == "" && os.Getenv("DWS_SESSION_ID") == "" {
|
||||
return fmt.Errorf("--session-id is required when --grant-type is session\n hint: dws pat chmod <scope> --agentCode <id> --grant-type session --session-id <id>")
|
||||
}
|
||||
|
||||
if caller != nil && caller.DryRun() {
|
||||
bold := color.New(color.FgYellow, color.Bold)
|
||||
bold.Println("[DRY-RUN] Preview only, not executed:")
|
||||
fmt.Printf("%-16s%s\n", "Tool:", "个人授权")
|
||||
fmt.Printf("%-16s%s\n", "AgentCode:", agentCode)
|
||||
fmt.Printf("%-16s%v\n", "Scope:", scopes)
|
||||
fmt.Printf("%-16s%s\n", "GrantType:", grantType)
|
||||
if sessionID != "" {
|
||||
fmt.Printf("%-16s%s\n", "SessionID:", sessionID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if caller == nil {
|
||||
return fmt.Errorf("internal error: tool runtime not initialized")
|
||||
}
|
||||
|
||||
toolArgs := map[string]any{
|
||||
"agentCode": agentCode,
|
||||
"scope": scopes,
|
||||
"grantType": grantType,
|
||||
}
|
||||
if sessionID == "" {
|
||||
sessionID = os.Getenv("DWS_SESSION_ID")
|
||||
}
|
||||
if sessionID != "" {
|
||||
toolArgs["sessionId"] = sessionID
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
result, err := caller.CallTool(ctx, "pat", "个人授权", toolArgs)
|
||||
if err != nil {
|
||||
return fmt.Errorf("pat chmod failed: %w", err)
|
||||
}
|
||||
|
||||
return handleToolResult(result)
|
||||
},
|
||||
}
|
||||
|
||||
chmodCmd.Flags().String("agentCode", "", "Agent 唯一标识(必填)")
|
||||
_ = chmodCmd.MarkFlagRequired("agentCode")
|
||||
chmodCmd.Flags().String("grant-type", "session", "授权策略: once|session|permanent")
|
||||
chmodCmd.Flags().String("session-id", "", "会话标识(session 模式下必填)")
|
||||
|
||||
return chmodCmd
|
||||
}
|
||||
|
||||
// handleToolResult processes a ToolResult and writes output to stdout.
|
||||
func handleToolResult(result *edition.ToolResult) error {
|
||||
if result == nil {
|
||||
return fmt.Errorf("empty tool result")
|
||||
}
|
||||
for _, c := range result.Content {
|
||||
if c.Type != "text" || c.Text == "" {
|
||||
continue
|
||||
}
|
||||
if respErr := apperrors.ClassifyMCPResponseText(c.Text); respErr != nil {
|
||||
return respErr
|
||||
}
|
||||
fmt.Println(c.Text)
|
||||
return nil
|
||||
}
|
||||
data, err := json.MarshalIndent(result, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to marshal result: %w", err)
|
||||
}
|
||||
fmt.Println(string(data))
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
// Package pat implements the "dws pat" command group for PAT (Personal Action
|
||||
// Token) authorization management.
|
||||
package pat
|
||||
|
||||
import (
|
||||
"github.com/spf13/cobra"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// RegisterCommands adds the pat command tree to rootCmd.
|
||||
func RegisterCommands(root *cobra.Command, c edition.ToolCaller) {
|
||||
patCmd := &cobra.Command{
|
||||
Use: "pat",
|
||||
Short: "行为授权管理",
|
||||
Long: `管理行为授权(PAT)。
|
||||
|
||||
命令结构:
|
||||
dws pat chmod <scope>... 授予指定权限`,
|
||||
RunE: cmdutil.GroupRunE,
|
||||
}
|
||||
|
||||
patCmd.AddCommand(newChmodCommand(c))
|
||||
root.AddCommand(patCmd)
|
||||
}
|
||||
@@ -244,6 +244,194 @@ func TestFullPipelineEndToEnd(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestFullFivePhasePipeline exercises all five phases in order:
|
||||
// Register → PreParse → PostParse → PreRequest → PostResponse,
|
||||
// simulating a complete command lifecycle from registration through
|
||||
// response output.
|
||||
func TestFullFivePhasePipeline(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
RegisterHandler{},
|
||||
AliasHandler{},
|
||||
StickyHandler{},
|
||||
ParamNameHandler{},
|
||||
ParamValueHandler{},
|
||||
PreRequestHandler{},
|
||||
PostResponseHandler{},
|
||||
)
|
||||
|
||||
// Verify all five phases have handlers.
|
||||
for _, phase := range []pipeline.Phase{
|
||||
pipeline.Register,
|
||||
pipeline.PreParse,
|
||||
pipeline.PostParse,
|
||||
pipeline.PreRequest,
|
||||
pipeline.PostResponse,
|
||||
} {
|
||||
if !engine.HasHandlers(phase) {
|
||||
t.Fatalf("engine missing handlers for phase %v", phase)
|
||||
}
|
||||
}
|
||||
|
||||
// Phase 1: Register — command tree being built.
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable",
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.Register, ctx); err != nil {
|
||||
t.Fatalf("Register error: %v", err)
|
||||
}
|
||||
|
||||
// Phase 2: PreParse — fix raw argv.
|
||||
ctx.Args = []string{
|
||||
"--userId", "u001",
|
||||
"--pageSize50",
|
||||
"--verbosetrue",
|
||||
}
|
||||
ctx.FlagSpecs = flagSpecs("user-id", "page-size", "verbose")
|
||||
|
||||
if err := engine.RunPhase(pipeline.PreParse, ctx); err != nil {
|
||||
t.Fatalf("PreParse error: %v", err)
|
||||
}
|
||||
|
||||
want := "--user-id u001 --page-size 50 --verbose true"
|
||||
got := strings.Join(ctx.Args, " ")
|
||||
if got != want {
|
||||
t.Errorf("after PreParse: Args = %q, want %q", got, want)
|
||||
}
|
||||
preParseCorrections := len(ctx.Corrections)
|
||||
|
||||
// Phase 3: PostParse — simulate Cobra having parsed the corrected
|
||||
// args into structured params, then normalise values.
|
||||
ctx.Command = "aitable.query_records"
|
||||
ctx.Params = map[string]any{
|
||||
"user_id": "u001",
|
||||
"page_size": "1,000",
|
||||
"verbose": "yes",
|
||||
}
|
||||
ctx.Schema = map[string]any{
|
||||
"properties": map[string]any{
|
||||
"user_id": map[string]any{"type": "string"},
|
||||
"page_size": map[string]any{"type": "integer"},
|
||||
"verbose": map[string]any{"type": "boolean"},
|
||||
},
|
||||
}
|
||||
|
||||
if err := engine.RunPhase(pipeline.PostParse, ctx); err != nil {
|
||||
t.Fatalf("PostParse error: %v", err)
|
||||
}
|
||||
|
||||
if got := ctx.Params["verbose"]; got != true {
|
||||
t.Errorf("verbose = %v (%T), want true (bool)", got, got)
|
||||
}
|
||||
if got := ctx.Params["page_size"]; got != int64(1000) {
|
||||
t.Errorf("page_size = %v, want 1000", got)
|
||||
}
|
||||
postParseCorrections := len(ctx.Corrections) - preParseCorrections
|
||||
if postParseCorrections != 2 {
|
||||
t.Errorf("PostParse corrections = %d, want 2", postParseCorrections)
|
||||
}
|
||||
|
||||
// Phase 4: PreRequest — inspect final payload before dispatch.
|
||||
ctx.Payload = ctx.Params
|
||||
if err := engine.RunPhase(pipeline.PreRequest, ctx); err != nil {
|
||||
t.Fatalf("PreRequest error: %v", err)
|
||||
}
|
||||
// Verify payload was not corrupted.
|
||||
if ctx.Payload["user_id"] != "u001" {
|
||||
t.Error("PreRequest corrupted Payload")
|
||||
}
|
||||
|
||||
// Phase 5: PostResponse — process response before output.
|
||||
ctx.Response = map[string]any{
|
||||
"records": []any{
|
||||
map[string]any{"id": "rec001", "fields": map[string]any{"name": "test"}},
|
||||
},
|
||||
"total": 1,
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.PostResponse, ctx); err != nil {
|
||||
t.Fatalf("PostResponse error: %v", err)
|
||||
}
|
||||
// Verify response was not corrupted.
|
||||
if ctx.Response["total"] != 1 {
|
||||
t.Error("PostResponse corrupted Response")
|
||||
}
|
||||
}
|
||||
|
||||
// TestFullFivePhasePipelineWithEngineRun exercises all five phases
|
||||
// using Engine.Run (single shot) to verify the ordering is correct
|
||||
// end-to-end.
|
||||
func TestFullFivePhasePipelineWithEngineRun(t *testing.T) {
|
||||
var seq []string
|
||||
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
&phaseTracker{name: "reg", phase: pipeline.Register, seq: &seq},
|
||||
&phaseTracker{name: "pre-parse", phase: pipeline.PreParse, seq: &seq},
|
||||
&phaseTracker{name: "post-parse", phase: pipeline.PostParse, seq: &seq},
|
||||
&phaseTracker{name: "pre-req", phase: pipeline.PreRequest, seq: &seq},
|
||||
&phaseTracker{name: "post-resp", phase: pipeline.PostResponse, seq: &seq},
|
||||
)
|
||||
|
||||
ctx := &pipeline.Context{Command: "test.tool"}
|
||||
if err := engine.Run(ctx); err != nil {
|
||||
t.Fatalf("Engine.Run error: %v", err)
|
||||
}
|
||||
|
||||
want := "reg,pre-parse,post-parse,pre-req,post-resp"
|
||||
got := strings.Join(seq, ",")
|
||||
if got != want {
|
||||
t.Errorf("phase execution order = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFivePhasePipelineCorrectHandlerCounts verifies that the
|
||||
// production-equivalent engine has the expected handler distribution.
|
||||
func TestFivePhasePipelineCorrectHandlerCounts(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
RegisterHandler{},
|
||||
AliasHandler{},
|
||||
StickyHandler{},
|
||||
ParamNameHandler{},
|
||||
ParamValueHandler{},
|
||||
PreRequestHandler{},
|
||||
PostResponseHandler{},
|
||||
)
|
||||
|
||||
tests := []struct {
|
||||
phase pipeline.Phase
|
||||
want int
|
||||
}{
|
||||
{pipeline.Register, 1},
|
||||
{pipeline.PreParse, 3},
|
||||
{pipeline.PostParse, 1},
|
||||
{pipeline.PreRequest, 1},
|
||||
{pipeline.PostResponse, 1},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := len(engine.Handlers(tt.phase)); got != tt.want {
|
||||
t.Errorf("Handlers(%v) = %d, want %d", tt.phase, got, tt.want)
|
||||
}
|
||||
}
|
||||
if got := engine.HandlerCount(); got != 7 {
|
||||
t.Errorf("HandlerCount = %d, want 7", got)
|
||||
}
|
||||
}
|
||||
|
||||
// phaseTracker is a test helper that records its name when Handle is called.
|
||||
type phaseTracker struct {
|
||||
name string
|
||||
phase pipeline.Phase
|
||||
seq *[]string
|
||||
}
|
||||
|
||||
func (h *phaseTracker) Name() string { return h.name }
|
||||
func (h *phaseTracker) Phase() pipeline.Phase { return h.phase }
|
||||
func (h *phaseTracker) Handle(_ *pipeline.Context) error {
|
||||
*h.seq = append(*h.seq, h.name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// TestPreParseDoesNotBreakValidArgs verifies that valid, correctly
|
||||
// formatted args pass through the pipeline without modification.
|
||||
func TestPreParseDoesNotBreakValidArgs(t *testing.T) {
|
||||
|
||||
@@ -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]
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
// PreRequestHandler runs in the PreRequest phase — after parameter
|
||||
// validation succeeds and just before the JSON-RPC call is dispatched.
|
||||
// It receives the final payload and can inspect or mutate it.
|
||||
//
|
||||
// Default behaviour: no-op pass-through. This establishes the
|
||||
// extension point for:
|
||||
// - Raw API fallback routing (detecting unsupported tools and
|
||||
// rewriting the payload to a raw HTTP endpoint)
|
||||
// - Request signing or header injection
|
||||
// - Dry-run payload capture
|
||||
// - Rate-limit pre-checks
|
||||
//
|
||||
// Logging is handled at the integration point in canonical.go,
|
||||
// consistent with how other phases log at their call sites.
|
||||
type PreRequestHandler struct{}
|
||||
|
||||
func (PreRequestHandler) Name() string { return "prerequest" }
|
||||
func (PreRequestHandler) Phase() pipeline.Phase { return pipeline.PreRequest }
|
||||
|
||||
func (PreRequestHandler) Handle(ctx *pipeline.Context) error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
func TestPreRequestHandlerMeta(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
if got := h.Name(); got != "prerequest" {
|
||||
t.Errorf("Name() = %q, want %q", got, "prerequest")
|
||||
}
|
||||
if got := h.Phase(); got != pipeline.PreRequest {
|
||||
t.Errorf("Phase() = %v, want %v", got, pipeline.PreRequest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerEmptyContext(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerNoSideEffects(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable.query_records",
|
||||
Params: map[string]any{
|
||||
"spaceId": "sp001",
|
||||
"datasheetId": "ds001",
|
||||
},
|
||||
Payload: map[string]any{
|
||||
"spaceId": "sp001",
|
||||
"datasheetId": "ds001",
|
||||
},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
if ctx.Params["spaceId"] != "sp001" {
|
||||
t.Error("PreRequestHandler should not mutate Params")
|
||||
}
|
||||
if ctx.Payload["spaceId"] != "sp001" {
|
||||
t.Error("PreRequestHandler should not mutate Payload")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerNilPayload(t *testing.T) {
|
||||
h := PreRequestHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "chat.send_message",
|
||||
Params: map[string]any{"userId": "u001"},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPreRequestHandlerInEngine(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.Register(PreRequestHandler{})
|
||||
|
||||
if !engine.HasHandlers(pipeline.PreRequest) {
|
||||
t.Fatal("engine should have PreRequest handler")
|
||||
}
|
||||
|
||||
ctx := &pipeline.Context{
|
||||
Command: "todo.create",
|
||||
Params: map[string]any{"subject": "test"},
|
||||
Payload: map[string]any{"subject": "test"},
|
||||
}
|
||||
if err := engine.RunPhase(pipeline.PreRequest, ctx); err != nil {
|
||||
t.Fatalf("RunPhase(PreRequest) returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
// RegisterHandler runs during the Register phase — the first stage
|
||||
// in the pipeline, executed while the Cobra command tree is being
|
||||
// built. It validates that the registration context carries a
|
||||
// non-empty command identifier.
|
||||
//
|
||||
// The handler is intentionally lightweight and side-effect free.
|
||||
// This provides the structural hook for future extensions (e.g.
|
||||
// dynamic command injection, feature gating, or Raw API fallback
|
||||
// command registration) without adding any runtime overhead to
|
||||
// the default path. Logging is handled at the call site in
|
||||
// canonical.go, consistent with how PreParse logging is done
|
||||
// in cobra.go.
|
||||
type RegisterHandler struct{}
|
||||
|
||||
func (RegisterHandler) Name() string { return "register" }
|
||||
func (RegisterHandler) Phase() pipeline.Phase { return pipeline.Register }
|
||||
|
||||
func (RegisterHandler) Handle(ctx *pipeline.Context) error {
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,84 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
func TestRegisterHandlerMeta(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
if got := h.Name(); got != "register" {
|
||||
t.Errorf("Name() = %q, want %q", got, "register")
|
||||
}
|
||||
if got := h.Phase(); got != pipeline.Register {
|
||||
t.Errorf("Phase() = %v, want %v", got, pipeline.Register)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerEmptyContext(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerWithCommand(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "aitable",
|
||||
Schema: map[string]any{
|
||||
"properties": map[string]any{
|
||||
"spaceId": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerNoSideEffects(t *testing.T) {
|
||||
h := RegisterHandler{}
|
||||
ctx := &pipeline.Context{
|
||||
Command: "todo",
|
||||
Params: map[string]any{"key": "value"},
|
||||
}
|
||||
if err := h.Handle(ctx); err != nil {
|
||||
t.Fatalf("Handle returned error: %v", err)
|
||||
}
|
||||
if ctx.Params["key"] != "value" {
|
||||
t.Error("RegisterHandler should not mutate Params")
|
||||
}
|
||||
if ctx.Command != "todo" {
|
||||
t.Error("RegisterHandler should not mutate Command")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterHandlerInEngine(t *testing.T) {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.Register(RegisterHandler{})
|
||||
|
||||
if !engine.HasHandlers(pipeline.Register) {
|
||||
t.Fatal("engine should have Register handler")
|
||||
}
|
||||
|
||||
ctx := &pipeline.Context{Command: "calendar"}
|
||||
if err := engine.RunPhase(pipeline.Register, ctx); err != nil {
|
||||
t.Fatalf("RunPhase(Register) returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package plugin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
)
|
||||
|
||||
const defaultHookTimeout = 30 * time.Second
|
||||
|
||||
// HookAdapter wraps a plugin hook entry as a pipeline.Handler.
|
||||
type HookAdapter struct {
|
||||
pluginName string
|
||||
entry HookEntry
|
||||
phase pipeline.Phase
|
||||
timeout time.Duration
|
||||
}
|
||||
|
||||
// NewHookAdapter creates a pipeline handler from a plugin hook entry.
|
||||
func NewHookAdapter(pluginName string, entry HookEntry) *HookAdapter {
|
||||
phase := parsePhase(entry.Phase)
|
||||
timeout := defaultHookTimeout
|
||||
if entry.Timeout > 0 {
|
||||
timeout = time.Duration(entry.Timeout) * time.Second
|
||||
}
|
||||
return &HookAdapter{
|
||||
pluginName: pluginName,
|
||||
entry: entry,
|
||||
phase: phase,
|
||||
timeout: timeout,
|
||||
}
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Name() string {
|
||||
return fmt.Sprintf("plugin-hook:%s/%s", h.pluginName, h.entry.Phase)
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Phase() pipeline.Phase {
|
||||
return h.phase
|
||||
}
|
||||
|
||||
func (h *HookAdapter) Handle(ctx *pipeline.Context) error {
|
||||
// Check matcher: if set, only run for matching commands.
|
||||
if h.entry.Matcher != "" {
|
||||
matched, err := filepath.Match(h.entry.Matcher, ctx.Command)
|
||||
if err != nil || !matched {
|
||||
return nil // skip silently
|
||||
}
|
||||
}
|
||||
|
||||
// Serialize context to JSON for the hook's stdin.
|
||||
input, err := json.Marshal(map[string]any{
|
||||
"command": ctx.Command,
|
||||
"params": ctx.Params,
|
||||
"args": ctx.Args,
|
||||
})
|
||||
if err != nil {
|
||||
slog.Warn("plugin hook: failed to serialize context",
|
||||
"plugin", h.pluginName, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
timeoutCtx, cancel := context.WithTimeout(context.Background(), h.timeout)
|
||||
defer cancel()
|
||||
|
||||
cmd := exec.CommandContext(timeoutCtx, "sh", "-c", h.entry.Command)
|
||||
cmd.Stdin = strings.NewReader(string(input))
|
||||
output, err := cmd.CombinedOutput()
|
||||
|
||||
if err != nil {
|
||||
if exitErr, ok := err.(*exec.ExitError); ok {
|
||||
code := exitErr.ExitCode()
|
||||
if code == 2 {
|
||||
// Exit 2 = abort pipeline.
|
||||
return fmt.Errorf("plugin hook %s/%s aborted: %s",
|
||||
h.pluginName, h.entry.Phase, strings.TrimSpace(string(output)))
|
||||
}
|
||||
}
|
||||
slog.Warn("plugin hook failed",
|
||||
"plugin", h.pluginName,
|
||||
"phase", h.entry.Phase,
|
||||
"error", err,
|
||||
"output", string(output),
|
||||
)
|
||||
return nil // non-fatal: log warning and continue
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func parsePhase(s string) pipeline.Phase {
|
||||
switch strings.TrimSpace(strings.ToLower(s)) {
|
||||
case "pre-parse":
|
||||
return pipeline.PreParse
|
||||
case "post-parse":
|
||||
return pipeline.PostParse
|
||||
case "pre-request":
|
||||
return pipeline.PreRequest
|
||||
case "post-response":
|
||||
return pipeline.PostResponse
|
||||
default:
|
||||
return pipeline.PreRequest
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,193 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package plugin
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSetAndGetPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Initially empty.
|
||||
val, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if ok {
|
||||
t.Errorf("expected not found, got %q", val)
|
||||
}
|
||||
|
||||
// Set a value.
|
||||
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
|
||||
|
||||
// Read it back.
|
||||
val, ok = loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if !ok {
|
||||
t.Fatal("expected to find config after set")
|
||||
}
|
||||
if val != "sk-test-12345" {
|
||||
t.Errorf("got %q, want sk-test-12345", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetPluginConfigMultipleKeys(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("my-plugin", "API_KEY", "key-1")
|
||||
loader.SetPluginConfig("my-plugin", "API_ENDPOINT", "https://example.com")
|
||||
loader.SetPluginConfig("other-plugin", "TOKEN", "tok-abc")
|
||||
|
||||
val, ok := loader.GetPluginConfig("my-plugin", "API_KEY")
|
||||
if !ok || val != "key-1" {
|
||||
t.Errorf("API_KEY = %q (ok=%v), want key-1", val, ok)
|
||||
}
|
||||
|
||||
val, ok = loader.GetPluginConfig("my-plugin", "API_ENDPOINT")
|
||||
if !ok || val != "https://example.com" {
|
||||
t.Errorf("API_ENDPOINT = %q (ok=%v), want https://example.com", val, ok)
|
||||
}
|
||||
|
||||
val, ok = loader.GetPluginConfig("other-plugin", "TOKEN")
|
||||
if !ok || val != "tok-abc" {
|
||||
t.Errorf("TOKEN = %q (ok=%v), want tok-abc", val, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsetPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Unset on empty returns false.
|
||||
if loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
|
||||
t.Error("expected false for unset on empty config")
|
||||
}
|
||||
|
||||
// Set then unset.
|
||||
loader.SetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY", "sk-test-12345")
|
||||
if !loader.UnsetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY") {
|
||||
t.Error("expected true for unset of existing key")
|
||||
}
|
||||
|
||||
// Verify it's gone.
|
||||
_, ok := loader.GetPluginConfig("demo-devtool", "DASHSCOPE_API_KEY")
|
||||
if ok {
|
||||
t.Error("expected not found after unset")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsetPluginConfigCleansEmptyMap(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", "KEY1", "val1")
|
||||
loader.UnsetPluginConfig("demo-devtool", "KEY1")
|
||||
|
||||
// After removing the last key, the plugin entry should be cleaned up.
|
||||
configs := loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 0 {
|
||||
t.Errorf("expected empty config map after removing last key, got %v", configs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestListPluginConfig(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Empty list.
|
||||
configs := loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 0 {
|
||||
t.Errorf("expected empty, got %v", configs)
|
||||
}
|
||||
|
||||
// Set some values.
|
||||
loader.SetPluginConfig("demo-devtool", "KEY_A", "val-a")
|
||||
loader.SetPluginConfig("demo-devtool", "KEY_B", "val-b")
|
||||
|
||||
configs = loader.ListPluginConfig("demo-devtool")
|
||||
if len(configs) != 2 {
|
||||
t.Fatalf("expected 2 configs, got %d", len(configs))
|
||||
}
|
||||
if configs["KEY_A"] != "val-a" {
|
||||
t.Errorf("KEY_A = %q, want val-a", configs["KEY_A"])
|
||||
}
|
||||
if configs["KEY_B"] != "val-b" {
|
||||
t.Errorf("KEY_B = %q, want val-b", configs["KEY_B"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectPluginConfigEnv(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
// Use a unique env var name to avoid test pollution.
|
||||
envKey := "DWS_TEST_INJECT_CONFIG_" + t.Name()
|
||||
t.Cleanup(func() { os.Unsetenv(envKey) })
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", envKey, "injected-value")
|
||||
|
||||
// Ensure it's not already set.
|
||||
os.Unsetenv(envKey)
|
||||
|
||||
loader.InjectPluginConfigEnv()
|
||||
|
||||
got := os.Getenv(envKey)
|
||||
if got != "injected-value" {
|
||||
t.Errorf("env %s = %q, want injected-value", envKey, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInjectPluginConfigEnvDoesNotOverride(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
envKey := "DWS_TEST_INJECT_NOOVERRIDE_" + t.Name()
|
||||
t.Cleanup(func() { os.Unsetenv(envKey) })
|
||||
|
||||
// Pre-set the env var.
|
||||
os.Setenv(envKey, "user-value")
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", envKey, "config-value")
|
||||
loader.InjectPluginConfigEnv()
|
||||
|
||||
got := os.Getenv(envKey)
|
||||
if got != "user-value" {
|
||||
t.Errorf("env %s = %q, want user-value (should not be overridden)", envKey, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetPluginConfigOverwritesExisting(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("demo-devtool", "KEY", "old-value")
|
||||
loader.SetPluginConfig("demo-devtool", "KEY", "new-value")
|
||||
|
||||
val, ok := loader.GetPluginConfig("demo-devtool", "KEY")
|
||||
if !ok || val != "new-value" {
|
||||
t.Errorf("got %q (ok=%v), want new-value", val, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPluginConfigWrongPlugin(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
loader := &Loader{PluginsDir: dir, CLIVersion: "1.0.0"}
|
||||
|
||||
loader.SetPluginConfig("plugin-a", "KEY", "value")
|
||||
|
||||
_, ok := loader.GetPluginConfig("plugin-b", "KEY")
|
||||
if ok {
|
||||
t.Error("expected not found for different plugin name")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,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
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user