Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9a119fbd64 | ||
|
|
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 | ||
|
|
28b775198d | ||
|
|
e976bd5fc9 | ||
|
|
86355413fc | ||
|
|
c650afa6eb | ||
|
|
14f558facf | ||
|
|
0257d1f084 | ||
|
|
b1b4730536 | ||
|
|
f732dcd2ba | ||
|
|
8406355e7f | ||
|
|
3c5f40648e | ||
|
|
794e168008 | ||
|
|
35548e4780 | ||
|
|
2a056cc5d0 | ||
|
|
3baadb99ce | ||
|
|
fcb8b2c782 | ||
|
|
c74f1eeb26 | ||
|
|
933615003c | ||
|
|
fc49f3dc7a | ||
|
|
8fb0ecb9ba | ||
|
|
ce6c32bf53 | ||
|
|
7dbef98dd0 | ||
|
|
773804ee80 | ||
|
|
1df56cc99b | ||
|
|
0606762c29 | ||
|
|
e7677df541 | ||
|
|
dda4dacb1c | ||
|
|
345156c605 | ||
|
|
cc4dd1e87b | ||
|
|
3c75c66d4d | ||
|
|
93d6fdb17e | ||
|
|
75b873d231 | ||
|
|
a912cbc52b | ||
|
|
5a99b84c25 | ||
|
|
7fda120d5a | ||
|
|
3b8233e6ba | ||
|
|
675ce53c06 | ||
|
|
d51c1ff131 | ||
|
|
8b423b97e8 | ||
|
|
1a6d129fe3 | ||
|
|
a0bf715ddf | ||
|
|
110f887181 | ||
|
|
d2c5a027d4 | ||
|
|
077665e27b | ||
|
|
5c41d2b8f4 | ||
|
|
7d9e592f84 | ||
|
|
11199e6848 | ||
|
|
19c38d4b94 | ||
|
|
851cf43180 | ||
|
|
e004df38c7 | ||
|
|
3bb504bb1f | ||
|
|
26263f8a17 | ||
|
|
fc22f53b92 | ||
|
|
654bcc4ecb | ||
|
|
2175f2fe59 | ||
|
|
81ca098db1 | ||
|
|
25118d1ec7 | ||
|
|
3e3c17d686 | ||
|
|
964855373e | ||
|
|
3c83c0cff2 | ||
|
|
4b555abefe | ||
|
|
4742112000 | ||
|
|
7f0567aa39 | ||
|
|
8fb1dcda86 | ||
|
|
94deafbaa9 | ||
|
|
2555447c7b | ||
|
|
6e91b2d142 | ||
|
|
54145b65ec | ||
|
|
2fad9c95db |
@@ -1 +1 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 53.9%"><title>coverage: 53.9%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">53.9%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">53.9%</text></g></svg>
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 48.8%"><title>coverage: 48.8%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">48.8%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">48.8%</text></g></svg>
|
||||
|
Before Width: | Height: | Size: 1.1 KiB After Width: | Height: | Size: 1.1 KiB |
@@ -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
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
# Issue 变更推送到 Webhook
|
||||
# 当有 Issue 变更时,发送指定格式的数据到 webhook
|
||||
name: 📤 Issue Webhook Notification
|
||||
|
||||
on:
|
||||
issues:
|
||||
types: [opened, reopened, closed, edited, labeled, unlabeled]
|
||||
|
||||
jobs:
|
||||
notify:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: 📬 Send Issue to Webhook
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const webhook = process.env.ISSUE_WEBHOOK_URL;
|
||||
if (!webhook) {
|
||||
console.log('⚠️ ISSUE_WEBHOOK_URL not set, skipping notification');
|
||||
return;
|
||||
}
|
||||
|
||||
const payload = context.payload;
|
||||
const issue = payload.issue;
|
||||
const action = payload.action;
|
||||
|
||||
// 构建指定格式的数据
|
||||
const webhookPayload = {
|
||||
action: action,
|
||||
issue: {
|
||||
id: issue.id,
|
||||
number: issue.number,
|
||||
title: issue.title,
|
||||
body: issue.body,
|
||||
state: issue.state,
|
||||
html_url: issue.html_url,
|
||||
labels: (issue.labels || []).map(label => label.name).join(', ') || '无标签'
|
||||
}
|
||||
};
|
||||
|
||||
const response = await fetch(webhook, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify(webhookPayload)
|
||||
});
|
||||
|
||||
if (response.ok) {
|
||||
console.log('✅ Webhook notification sent successfully');
|
||||
} else {
|
||||
console.log('❌ Failed to send webhook notification:', response.status, response.statusText);
|
||||
}
|
||||
env:
|
||||
ISSUE_WEBHOOK_URL: ${{ secrets.DINGTALK_AI_TABLE_WEBHOOK }}
|
||||
@@ -0,0 +1,50 @@
|
||||
# Issue 自动同步到钉钉群
|
||||
# 当有新 Issue 时,自动发送到钉钉群(不包括 comment)
|
||||
name: 🔔 Issue Notification to DingTalk
|
||||
|
||||
on:
|
||||
issues:
|
||||
types: [opened, reopened, closed, labeled]
|
||||
|
||||
jobs:
|
||||
notify:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: 📬 Send Issue to DingTalk
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const webhook = process.env.DINGTALK_WEBHOOK;
|
||||
if (!webhook) {
|
||||
console.log('⚠️ DINGTALK_WEBHOOK not set, skipping notification');
|
||||
return;
|
||||
}
|
||||
|
||||
const payload = context.payload;
|
||||
const issue = payload.issue;
|
||||
const action = payload.action;
|
||||
|
||||
// 构建消息标题和内容(确保包含关键字 "issue" 以支持 Custom Keywords 模式)
|
||||
const title = `[${action.toUpperCase()}] Issue #${issue.number}: ${issue.title}`;
|
||||
const content = issue.body?.substring(0, 500) || 'No description';
|
||||
const url = issue.html_url;
|
||||
const labelsText = (issue.labels || []).map(label => label.name).join(', ') || '无标签';
|
||||
|
||||
// 消息内容必须包含关键字(如 "issue")以支持 Custom Keywords 安全模式
|
||||
const message = {
|
||||
msgtype: 'markdown',
|
||||
markdown: {
|
||||
title: 'GitHub Issue 通知',
|
||||
text: `## 🔔 GitHub Issue 通知\n\n**${title}**\n\n🏷️ **Labels**: ${labelsText}\n\n${content}${content.length >= 500 ? '...' : ''}\n\n[点击查看详情](${url})\n\n---\n📦 ${context.repo.owner}/${context.repo.repo}\n\n**关键词**: issue`
|
||||
}
|
||||
};
|
||||
|
||||
await fetch(webhook, {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify(message)
|
||||
});
|
||||
|
||||
console.log('✅ DingTalk notification sent');
|
||||
env:
|
||||
DINGTALK_WEBHOOK: ${{ secrets.DINGTALK_WEBHOOK }}
|
||||
@@ -37,6 +37,14 @@ jobs:
|
||||
|
||||
- name: Post-release packaging
|
||||
run: ./scripts/release/post-goreleaser.sh
|
||||
env:
|
||||
DWS_PACKAGE_VERSION: ${{ github.ref_name }}
|
||||
|
||||
- name: Upload dws-skills.zip to release
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
gh release upload "${{ github.ref_name }}" dist/dws-skills.zip --clobber
|
||||
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@v4
|
||||
|
||||
@@ -27,3 +27,4 @@ test/cli_compat/testdata/
|
||||
credentials*
|
||||
plans
|
||||
_docs
|
||||
dws.zip
|
||||
|
||||
@@ -69,3 +69,4 @@ release:
|
||||
draft: false
|
||||
prerelease: auto
|
||||
name_template: "v{{.Version}}"
|
||||
mode: replace
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
<p align="center">
|
||||
<img src="https://img.shields.io/badge/Go-1.25+-green?logo=go&logoColor=white" alt="Go 1.25+">
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/blob/main/LICENSE"><img src="https://img.shields.io/badge/License-Apache_2.0-blue" alt="License Apache-2.0"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/badge/release-v1.1.0-red" alt="v1.1.0"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/github/v/release/DingTalk-Real-AI/dingtalk-workspace-cli?color=red&label=release" alt="Latest Release"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml"><img src="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml/badge.svg" alt="CI"></a>
|
||||
<img src=".github/badges/coverage.svg" alt="Coverage">
|
||||
</p>
|
||||
@@ -19,15 +19,16 @@
|
||||
</p>
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **Co-creation Phase**: This project accesses DingTalk enterprise data and requires enterprise admin authorization. Please join the DingTalk DWS co-creation group to complete whitelist configuration. See [Getting Started](#getting-started) below.
|
||||
> **Co-creation Phase**: This project accesses DingTalk enterprise data and requires enterprise admin authorization. Join the DingTalk DWS co-creation group for support and updates. See [Getting Started](#getting-started) below.
|
||||
>
|
||||
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i1/O1CN01ZqtgeV1cImFmTZPAH_!!6000000003578-2-tps-398-372.png" alt="DingTalk Group QR Code" width="150"></a>
|
||||
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i4/O1CN01Rijgk81gKqVSKMzdx_!!6000000004124-2-tps-654-644.png" alt="DingTalk Group QR Code" width="150"></a>
|
||||
|
||||
<details>
|
||||
<summary><strong>Table of Contents</strong></summary>
|
||||
|
||||
- [Why dws?](#why-dws)
|
||||
- [Installation](#installation)
|
||||
- [Upgrade](#upgrade)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Using with Agents](#using-with-agents)
|
||||
@@ -39,6 +40,7 @@
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
---
|
||||
|
||||
<h2 id="why-dws">Why dws?</h2>
|
||||
@@ -64,8 +66,19 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
|
||||
<details>
|
||||
<summary>Other install methods</summary>
|
||||
|
||||
**npm** (requires Node.js (npm/npx)):
|
||||
|
||||
```bash
|
||||
npm install -g dingtalk-workspace-cli
|
||||
```
|
||||
|
||||
**Pre-built binary**: download from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases).
|
||||
|
||||
> **macOS users**: If you see "cannot be opened because Apple cannot check it for malicious software", run:
|
||||
> ```bash
|
||||
> xattr -d com.apple.quarantine /path/to/dws
|
||||
> ```
|
||||
|
||||
**Build from source**:
|
||||
|
||||
```bash
|
||||
@@ -79,67 +92,95 @@ cp dws ~/.local/bin/ # install to PATH
|
||||
|
||||
</details>
|
||||
|
||||
## Upgrade
|
||||
|
||||
> Requires **v1.0.7** or later. For earlier versions, please re-run the [install script](#installation) to upgrade.
|
||||
|
||||
dws has built-in self-upgrade capability. Updates are pulled directly from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) with SHA256 integrity verification and automatic backup.
|
||||
|
||||
```bash
|
||||
dws upgrade # interactive upgrade to latest version
|
||||
dws upgrade --check # check for new versions without installing
|
||||
dws upgrade --list # list all available versions
|
||||
dws upgrade --version v1.0.7 # upgrade to a specific version
|
||||
dws upgrade --rollback # rollback to the previous version
|
||||
dws upgrade -y # skip confirmation prompt
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary><strong>How it works</strong></summary>
|
||||
|
||||
The upgrade process follows a two-phase atomic flow to ensure consistency:
|
||||
|
||||
1. **Prepare** — downloads the platform-specific binary and skill packages to a temporary directory, verifies SHA256 checksums, and extracts/validates all files. If any step fails, the upgrade aborts without modifying the existing installation.
|
||||
2. **Apply** — only after all preparations succeed, the binary is replaced and skill packages are installed to all detected agent directories (`~/.agents/skills/dws`, `~/.claude/skills/dws`, `~/.cursor/skills/dws`, etc.).
|
||||
|
||||
A backup of the current version is automatically created before each upgrade. Use `dws upgrade --rollback` to restore the previous version if needed.
|
||||
|
||||
| Flag | Description |
|
||||
|------|-------------|
|
||||
| `--check` | Check for updates without installing |
|
||||
| `--list` | List all available versions with changelogs |
|
||||
| `--version` | Upgrade to a specific version (e.g. `v1.0.7`) |
|
||||
| `--rollback` | Rollback to the previous backed-up version |
|
||||
| `--force` | Force reinstall even if already on the latest version |
|
||||
| `--skip-skills` | Skip skill package update |
|
||||
| `-y` | Skip confirmation prompt |
|
||||
|
||||
</details>
|
||||
|
||||
## Getting Started
|
||||
|
||||
### Step 1: Create a DingTalk Application
|
||||
|
||||
Go to the [Open Platform Console](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app). Under "Internal Enterprise Apps - DingTalk Apps", click **Create App**.
|
||||
|
||||
<details>
|
||||
<summary>View screenshot</summary>
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="Create Application" width="600">
|
||||
</p>
|
||||
</details>
|
||||
|
||||
### Step 2: Configure Redirect URL
|
||||
|
||||
Go to app settings → **Security Settings**. Add the following redirect URLs and save:
|
||||
|
||||
```
|
||||
http://127.0.0.1
|
||||
https://login.dingtalk.com
|
||||
```bash
|
||||
dws auth login # browser opens automatically
|
||||
dws auth login --device # for headless environments (Docker, SSH, CI)
|
||||
```
|
||||
|
||||
> `http://127.0.0.1` is for local browser login; `https://login.dingtalk.com` is for `--device` device-flow login (Docker containers, remote servers, and other headless environments). We recommend configuring both.
|
||||
Select your organization and authorize. That's it.
|
||||
|
||||
> If your organization hasn't enabled CLI access, you'll be prompted to send an access request to your admin. Once approved, re-run `dws auth login`.
|
||||
|
||||
<details>
|
||||
<summary>View screenshot</summary>
|
||||
<summary><strong>Organization hasn't enabled CLI access?</strong></summary>
|
||||
|
||||
1. After selecting your organization, click "Apply Now" to notify the admin
|
||||
2. The admin receives a request card and can approve with one click
|
||||
3. Once approved, re-run `dws auth login`
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="Configure Redirect URL" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="Apply for Access" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
### Step 3: Publish the Application
|
||||
|
||||
Click "App Release - Version Management & Release" to publish and go live.
|
||||
|
||||
<details>
|
||||
<summary>View screenshot</summary>
|
||||
<summary><strong>Admin: Enable CLI access for your organization</strong></summary>
|
||||
|
||||
Go to [Developer Platform](https://open-dev.dingtalk.com) → "CLI Access Management" → Enable.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="Publish Application" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI Access Management" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
### Step 4: Request Whitelist Access
|
||||
<details>
|
||||
<summary><strong>Custom App mode (CI/CD, ISV integration)</strong></summary>
|
||||
|
||||
Join the DingTalk DWS co-creation group and provide your **Client ID** and **admin confirmation** to complete whitelist setup.
|
||||
For enterprise-managed scenarios, create your own DingTalk app:
|
||||
|
||||
### Step 5: Authenticate
|
||||
1. [Open Platform Console](https://open-dev.dingtalk.com/fe/app#/corp/app) → Create App
|
||||
2. Security Settings → Add redirect URLs: `http://127.0.0.1,https://login.dingtalk.com`
|
||||
3. Publish the app
|
||||
4. Login:
|
||||
|
||||
```bash
|
||||
dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
|
||||
```
|
||||
|
||||
Or via environment variables:
|
||||
Credentials are securely persisted after first login (Keychain). Subsequent runs auto-refresh tokens.
|
||||
|
||||
```bash
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
```
|
||||
|
||||
> CLI flags take precedence over environment variables. Credentials are used for DingTalk's OAuth device flow.
|
||||
</details>
|
||||
|
||||
## Quick Start
|
||||
|
||||
@@ -154,13 +195,6 @@ dws todo task list --dry-run # preview without executing
|
||||
|
||||
dws is designed as an AI-native CLI. Complete [Installation](#installation) and [Getting Started](#getting-started) first, then configure your agent:
|
||||
|
||||
```bash
|
||||
# Configure auth via environment variables (recommended for agents, no interactive login)
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
```
|
||||
|
||||
### Agent Invocation Patterns
|
||||
|
||||
```bash
|
||||
@@ -183,7 +217,7 @@ Agents don't need pre-built knowledge of every command. Use `dws schema` to dyna
|
||||
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
|
||||
|
||||
# Step 2: Inspect target tool's parameter schema
|
||||
dws schema aitable.query_records --jq '.tool.input_schema'
|
||||
dws schema aitable.query_records --jq '.tool.parameters'
|
||||
|
||||
# Step 3: Construct the correct call
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
|
||||
@@ -191,7 +225,7 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
|
||||
|
||||
### Agent Skills
|
||||
|
||||
The repo ships Agent Skills (`SKILL.md` files) for every DingTalk product. After installing, tools like Claude Code / Cursor can use DingTalk capabilities directly:
|
||||
The repo ships a complete Agent Skill system (`skills/`). After installing, AI tools like Claude Code / Cursor can operate DingTalk directly through natural language:
|
||||
|
||||
```bash
|
||||
# Install skills into current project
|
||||
@@ -200,12 +234,45 @@ curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace
|
||||
|
||||
> `install.sh` installs to `$HOME/.agents/skills/dws` (global); `install-skills.sh` installs to `./.agents/skills/dws` (current project).
|
||||
|
||||
Author your own Agent Skills and orchestrate them with dws skills for cross-product workflows: **ISV Skill → dws Skill → DingTalk Open Platform API (enforced auth + full audit)**.
|
||||
**What's included:**
|
||||
|
||||
| Component | Path | Description |
|
||||
|-----------|------|-------------|
|
||||
| Master Skill | `SKILL.md` | Intent routing, decision tree, safety rules, error handling |
|
||||
| Product references | `references/products/*.md` | Per-product command reference (aitable, chat, calendar, etc.) |
|
||||
| Intent guide | `references/intent-guide.md` | Disambiguation for confusing scenarios (e.g. report vs todo) |
|
||||
| Global reference | `references/global-reference.md` | Auth, output formats, global flags |
|
||||
| Error codes | `references/error-codes.md` | Error codes + debugging workflows |
|
||||
| Recovery guide | `references/recovery-guide.md` | `RECOVERY_EVENT_ID` handling |
|
||||
| Ready-made scripts | `scripts/*.py` | 13 batch operation scripts (see below) |
|
||||
|
||||
<details>
|
||||
<summary><strong>Ready-made scripts</strong> — 13 Python scripts for common multi-step workflows</summary>
|
||||
|
||||
| Script | Description |
|
||||
|--------|-------------|
|
||||
| `calendar_schedule_meeting.py` | Create event + add participants + find & book available meeting room |
|
||||
| `calendar_free_slot_finder.py` | Find common free slots across multiple people, recommend best meeting time |
|
||||
| `calendar_today_agenda.py` | View today/tomorrow/this week's schedule |
|
||||
| `import_records.py` | Batch import records from CSV/JSON into AITable |
|
||||
| `bulk_add_fields.py` | Batch add fields to an AITable data table |
|
||||
| `upload_attachment.py` | Upload attachment to AITable attachment field |
|
||||
| `todo_batch_create.py` | Batch create todos from JSON (with priority, due date, executors) |
|
||||
| `todo_daily_summary.py` | Summarize today/this week's incomplete todos |
|
||||
| `todo_overdue_check.py` | Scan overdue todos and output overdue list |
|
||||
| `contact_dept_members.py` | Search department by name and list all members |
|
||||
| `attendance_my_record.py` | View my attendance records for today/this week/specific date |
|
||||
| `attendance_team_shift.py` | Query team shift schedules and attendance statistics |
|
||||
| `report_inbox_today.py` | View today's received reports with details |
|
||||
|
||||
</details>
|
||||
|
||||
**ISV Integration**: Author your own Agent Skills and orchestrate them with dws skills for cross-product workflows: **ISV Skill → dws Skill → DingTalk Open Platform API (enforced auth + full audit)**.
|
||||
|
||||
## Features
|
||||
|
||||
<details>
|
||||
<summary><strong>Smart Input Correction</strong> — auto-corrects common AI model parameter mistakes <code>v1.0.1</code></summary>
|
||||
<summary><strong>Smart Input Correction</strong> — auto-corrects common AI model parameter mistakes</summary>
|
||||
|
||||
Built-in pipeline engine that normalizes flag names, splits sticky arguments, and fuzzy-matches typos:
|
||||
|
||||
@@ -234,7 +301,7 @@ dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-i
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>jq Filtering & Field Selection</strong> — fine-grained output control to reduce token consumption <code>v1.0.1</code></summary>
|
||||
<summary><strong>jq Filtering & Field Selection</strong> — fine-grained output control to reduce token consumption</summary>
|
||||
|
||||
```bash
|
||||
# Built-in jq expressions
|
||||
@@ -248,19 +315,19 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocati
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Schema Introspection</strong> — query parameter schemas before making calls <code>v1.0.1</code></summary>
|
||||
<summary><strong>Schema Introspection</strong> — query parameter schemas before making calls</summary>
|
||||
|
||||
```bash
|
||||
dws schema # list all products and tools
|
||||
dws schema aitable.query_records # view parameter schema
|
||||
dws schema aitable.query_records --jq '.tool.input_schema.required' # view required fields
|
||||
dws schema aitable.query_records --jq '.tool.required' # view required fields
|
||||
dws schema --jq '.products[].id' # extract all product IDs
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Pipe & File Input</strong> — read flag values from files or stdin <code>v1.0.1</code></summary>
|
||||
<summary><strong>Pipe & File Input</strong> — read flag values from files or stdin</summary>
|
||||
|
||||
```bash
|
||||
# Read message body from a file
|
||||
@@ -280,21 +347,22 @@ dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
|
||||
## Key Services
|
||||
|
||||
| Service | Command | Description |
|
||||
|---------|---------|-------------|
|
||||
| Contact | `contact` | Users / departments |
|
||||
| Chat | `chat` | Group management / members / bot messaging / webhook |
|
||||
| Calendar | `calendar` | Events / meeting rooms / free-busy |
|
||||
| Todo | `todo` | Task management |
|
||||
| Approval | `oa` | Processes / forms / instances |
|
||||
| Attendance | `attendance` | Clock-in / shifts / statistics |
|
||||
| Ding | `ding` | DING messages / send / recall |
|
||||
| Report | `report` | Reports / templates / statistics |
|
||||
| AITable | `aitable` | AI table operations |
|
||||
| Workbench | `workbench` | App query |
|
||||
| DevDoc | `devdoc` | Open platform docs search |
|
||||
| Service | Command | Commands | Subcommands | Description |
|
||||
|---------|---------|:--------:|-------------|-------------|
|
||||
| Contact | `contact` | 6 | `user` `dept` | Search users by name/mobile, batch query, departments, current user profile |
|
||||
| Chat | `chat` | 10 | `message` `group` `search` | Group CRUD, member management, bot messaging, webhook |
|
||||
| Bot | `chat bot` | 6 | `bot` `group` `message` `search` | Robot creation/search, group/single messaging, webhook, message recall |
|
||||
| Calendar | `calendar` | 13 | `event` `room` `participant` `busy` | Events CRUD, meeting room booking, free-busy query, participant management |
|
||||
| Todo | `todo` | 6 | `task` | Create, list, update, done, get detail, delete |
|
||||
| Approval | `oa` | 9 | `approval` | Approve/reject/revoke, pending tasks, initiated instances, process list |
|
||||
| Attendance | `attendance` | 4 | `record` `shift` `summary` `rules` | Clock-in records, shift schedules, attendance summary, group rules |
|
||||
| Ding | `ding` | 2 | `message` | Send/recall DING messages |
|
||||
| Report | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | Create reports, sent/received list, templates, statistics |
|
||||
| AITable | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | Full CRUD for bases/tables/records/fields, templates |
|
||||
| Workbench | `workbench` | 2 | `app` | Batch query app details |
|
||||
| DevDoc | `devdoc` | 1 | `article` | Search platform docs and error codes |
|
||||
|
||||
Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
|
||||
> 86 commands across 12 products. Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
|
||||
|
||||
<details>
|
||||
<summary>Coming soon</summary>
|
||||
|
||||
+138
-70
@@ -9,7 +9,7 @@
|
||||
<p align="center">
|
||||
<img src="https://img.shields.io/badge/Go-1.25+-green?logo=go&logoColor=white" alt="Go 1.25+">
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/blob/main/LICENSE"><img src="https://img.shields.io/badge/License-Apache_2.0-blue" alt="License Apache-2.0"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/badge/release-v1.1.0-red" alt="v1.1.0"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases"><img src="https://img.shields.io/github/v/release/DingTalk-Real-AI/dingtalk-workspace-cli?color=red&label=release" alt="Latest Release"></a>
|
||||
<a href="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml"><img src="https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/actions/workflows/ci.yml/badge.svg" alt="CI"></a>
|
||||
<img src=".github/badges/coverage.svg" alt="Coverage">
|
||||
</p>
|
||||
@@ -19,15 +19,16 @@
|
||||
</p>
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。当前为灰度共创阶段,请加入钉钉 DWS 共创群完成白名单配置。详见下方 [开始使用](#开始使用)。
|
||||
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。欢迎加入钉钉 DWS 共创群获取支持与最新动态。详见下方 [开始使用](#开始使用)。
|
||||
>
|
||||
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i1/O1CN01ZqtgeV1cImFmTZPAH_!!6000000003578-2-tps-398-372.png" alt="DingTalk Group QR Code" width="150"></a>
|
||||
> <a href="https://qr.dingtalk.com/action/joingroup?code=v1,k1,v9/YMJG9qXhvFk5juktYnQziN70rF7QHebC/JLztTVRuRVJIwrSsXmL8oFqU5ajJ&_dt_no_comment=1&origin=11"><img src="https://img.alicdn.com/imgextra/i4/O1CN01Rijgk81gKqVSKMzdx_!!6000000004124-2-tps-654-644.png" alt="DingTalk Group QR Code" width="150"></a>
|
||||
|
||||
<details>
|
||||
<summary><strong>目录</strong></summary>
|
||||
|
||||
- [为什么选择 dws?](#why-dws)
|
||||
- [安装](#安装)
|
||||
- [升级](#升级)
|
||||
- [开始使用](#开始使用)
|
||||
- [快速开始](#快速开始)
|
||||
- [在 Agent 中使用](#在-agent-中使用)
|
||||
@@ -39,6 +40,7 @@
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
---
|
||||
|
||||
<h2 id="why-dws">为什么选择 dws?</h2>
|
||||
@@ -64,8 +66,19 @@ irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/ma
|
||||
<details>
|
||||
<summary>其他安装方式</summary>
|
||||
|
||||
**npm**(需要 Node.js(npm/npx)):
|
||||
|
||||
```bash
|
||||
npm install -g dingtalk-workspace-cli
|
||||
```
|
||||
|
||||
**预编译二进制文件**:从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 下载。
|
||||
|
||||
> **macOS 用户注意**:如果提示“无法打开,因为 Apple 无法检查其是否包含恶意软件”,请执行:
|
||||
> ```bash
|
||||
> xattr -d com.apple.quarantine /path/to/dws
|
||||
> ```
|
||||
|
||||
**从源码构建**:
|
||||
|
||||
```bash
|
||||
@@ -79,67 +92,95 @@ cp dws ~/.local/bin/ # 安装到 PATH
|
||||
|
||||
</details>
|
||||
|
||||
## 升级
|
||||
|
||||
> 需要 **v1.0.7** 及以上版本。更早版本请重新执行[安装脚本](#安装)进行升级。
|
||||
|
||||
dws 内置自升级能力,直接从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 拉取更新,支持 SHA256 完整性校验和自动备份。
|
||||
|
||||
```bash
|
||||
dws upgrade # 交互式升级到最新版本
|
||||
dws upgrade --check # 仅检查是否有新版本
|
||||
dws upgrade --list # 列出所有可用版本
|
||||
dws upgrade --version v1.0.7 # 升级到指定版本
|
||||
dws upgrade --rollback # 回滚到上一版本
|
||||
dws upgrade -y # 跳过确认直接升级
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary><strong>工作原理</strong></summary>
|
||||
|
||||
升级过程采用两阶段原子流程,确保一致性:
|
||||
|
||||
1. **准备阶段** — 将平台对应的二进制文件和技能包下载到临时目录,校验 SHA256 校验和,解压并验证所有文件。任何步骤失败则立即中止,不会修改现有安装。
|
||||
2. **执行阶段** — 仅在所有准备工作成功后,替换二进制文件并将技能包安装到所有已检测到的 Agent 目录(`~/.agents/skills/dws`、`~/.claude/skills/dws`、`~/.cursor/skills/dws` 等)。
|
||||
|
||||
每次升级前自动备份当前版本,可通过 `dws upgrade --rollback` 随时回滚。
|
||||
|
||||
| Flag | 说明 |
|
||||
|------|------|
|
||||
| `--check` | 仅检查更新,不安装 |
|
||||
| `--list` | 列出所有可用版本及更新日志 |
|
||||
| `--version` | 升级到指定版本(如 `v1.0.7`) |
|
||||
| `--rollback` | 回滚到上一个备份版本 |
|
||||
| `--force` | 强制重新安装,即使已是最新版本 |
|
||||
| `--skip-skills` | 跳过技能包更新 |
|
||||
| `-y` | 跳过确认提示 |
|
||||
|
||||
</details>
|
||||
|
||||
## 开始使用
|
||||
|
||||
### 步骤 1:创建钉钉应用
|
||||
|
||||
进入 [开放平台应用开发后台](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app),在「企业内部应用 - 钉钉应用」点击**创建应用**。
|
||||
|
||||
<details>
|
||||
<summary>查看截图</summary>
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="创建应用" width="600">
|
||||
</p>
|
||||
</details>
|
||||
|
||||
### 步骤 2:配置重定向 URL
|
||||
|
||||
进入应用 → **安全设置**,在「重定向 URL」中添加以下地址并保存:
|
||||
|
||||
```
|
||||
http://127.0.0.1
|
||||
https://login.dingtalk.com
|
||||
```bash
|
||||
dws auth login # 自动唤起浏览器
|
||||
dws auth login --device # 无浏览器环境(Docker、SSH、CI)
|
||||
```
|
||||
|
||||
> `http://127.0.0.1` 用于本地浏览器登录;`https://login.dingtalk.com` 用于 `--device` 设备流登录(Docker 容器、远程服务器等无浏览器环境)。建议两个都配置。
|
||||
选择组织并授权即可。
|
||||
|
||||
> 如果组织尚未开启 CLI 访问权限,系统会引导你向管理员发送申请。审批通过后重新执行 `dws auth login` 即可。
|
||||
|
||||
<details>
|
||||
<summary>查看截图</summary>
|
||||
<summary><strong>组织未开启 CLI 访问权限?</strong></summary>
|
||||
|
||||
1. 选择组织后,点击「立即申请」通知管理员
|
||||
2. 管理员收到申请卡片,一键审批
|
||||
3. 审批通过后,重新执行 `dws auth login`
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="配置重定向URL" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="申请权限" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
### 步骤 3:发布应用
|
||||
|
||||
点击「应用发布 - 版本管理与发布」,发布版本使应用上线。
|
||||
|
||||
<details>
|
||||
<summary>查看截图</summary>
|
||||
<summary><strong>管理员:为组织开启 CLI 访问权限</strong></summary>
|
||||
|
||||
进入 [开发者平台](https://open-dev.dingtalk.com) →「CLI 访问管理」→ 开启。
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="发布应用" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI访问管理" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
### 步骤 4:申请白名单
|
||||
<details>
|
||||
<summary><strong>自建应用模式(CI/CD、ISV 集成)</strong></summary>
|
||||
|
||||
加入钉钉 DWS 共创群,提供 **Client ID** 和**管理员确认凭证**完成白名单配置。
|
||||
企业自主管控场景,可创建自有钉钉应用:
|
||||
|
||||
### 步骤 5:登录认证
|
||||
1. [开放平台应用开发后台](https://open-dev.dingtalk.com/fe/app#/corp/app) → 创建应用
|
||||
2. 安全设置 → 添加重定向 URL:`http://127.0.0.1,https://login.dingtalk.com`
|
||||
3. 发布应用
|
||||
4. 登录:
|
||||
|
||||
```bash
|
||||
dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
|
||||
```
|
||||
|
||||
或通过环境变量:
|
||||
首次登录后凭证安全存储(Keychain),后续自动刷新 Token。
|
||||
|
||||
```bash
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
```
|
||||
|
||||
> CLI 参数优先于环境变量。凭证用于钉钉 OAuth 设备流认证。
|
||||
</details>
|
||||
|
||||
## 快速开始
|
||||
|
||||
@@ -154,13 +195,6 @@ dws todo task list --dry-run # 预览操作但不执行
|
||||
|
||||
dws 是为 AI Agent 设计的 CLI 工具。请先完成[安装](#安装)和[开始使用](#开始使用),然后配置 Agent 环境:
|
||||
|
||||
```bash
|
||||
# 通过环境变量配置认证(Agent 推荐方式,无需交互式登录)
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
```
|
||||
|
||||
### Agent 调用模式
|
||||
|
||||
```bash
|
||||
@@ -183,7 +217,7 @@ Agent 无需预置所有命令知识,通过 `dws schema` 动态发现可用能
|
||||
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
|
||||
|
||||
# 第二步:查看目标工具的参数结构
|
||||
dws schema aitable.query_records --jq '.tool.input_schema'
|
||||
dws schema aitable.query_records --jq '.tool.parameters'
|
||||
|
||||
# 第三步:构造正确的调用
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
|
||||
@@ -191,7 +225,7 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
|
||||
|
||||
### Agent Skills
|
||||
|
||||
仓库为每个钉钉产品提供 Agent Skill(`SKILL.md`),安装后 Claude Code / Cursor 等工具可直接使用钉钉能力:
|
||||
仓库内置完整的 Agent Skill 体系(`skills/`),安装后 Claude Code / Cursor 等 AI 工具可通过自然语言直接操作钉钉:
|
||||
|
||||
```bash
|
||||
# 安装 skills 到当前项目
|
||||
@@ -200,12 +234,45 @@ curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace
|
||||
|
||||
> `install.sh` 安装到 `$HOME/.agents/skills/dws`(全局);`install-skills.sh` 安装到 `./.agents/skills/dws`(当前项目)。
|
||||
|
||||
编写您自己的 Agent Skill,与 dws 内置 Skill 搭配构建跨产品工作流:**ISV Skill → dws Skill → 钉钉开放平台 API(强制鉴权 + 全链路审计)**。
|
||||
**包含内容:**
|
||||
|
||||
| 组件 | 路径 | 说明 |
|
||||
|------|------|------|
|
||||
| 主 Skill | `SKILL.md` | 意图路由、决策树、安全规则、错误处理 |
|
||||
| 产品参考 | `references/products/*.md` | 各产品命令详细参考(aitable、chat、calendar 等) |
|
||||
| 意图指南 | `references/intent-guide.md` | 易混淆场景消歧(如 report vs todo) |
|
||||
| 全局参考 | `references/global-reference.md` | 认证、输出格式、全局 flag |
|
||||
| 错误码 | `references/error-codes.md` | 错误码 + 调试流程 |
|
||||
| Recovery 指南 | `references/recovery-guide.md` | `RECOVERY_EVENT_ID` 处理 |
|
||||
| 现成脚本 | `scripts/*.py` | 13 个批量操作脚本(见下方) |
|
||||
|
||||
<details>
|
||||
<summary><strong>现成脚本</strong> — 13 个 Python 脚本,覆盖常见多步工作流</summary>
|
||||
|
||||
| 脚本 | 说明 |
|
||||
|------|------|
|
||||
| `calendar_schedule_meeting.py` | 一键创建日程 + 添加参与者 + 搜索并预定空闲会议室 |
|
||||
| `calendar_free_slot_finder.py` | 查询多人共同空闲时段,推荐最佳会议时间 |
|
||||
| `calendar_today_agenda.py` | 查看今天/明天/本周的日程安排 |
|
||||
| `import_records.py` | 从 CSV/JSON 批量导入记录到 AI 表格 |
|
||||
| `bulk_add_fields.py` | 批量添加字段到 AI 表格数据表 |
|
||||
| `upload_attachment.py` | 上传附件到 AI 表格 attachment 字段 |
|
||||
| `todo_batch_create.py` | 从 JSON 文件批量创建待办(含优先级、截止时间、执行者) |
|
||||
| `todo_daily_summary.py` | 汇总今天/本周未完成的待办 |
|
||||
| `todo_overdue_check.py` | 扫描已过截止时间但未完成的待办,输出逾期清单 |
|
||||
| `contact_dept_members.py` | 按部门名称搜索并列出所有成员 |
|
||||
| `attendance_my_record.py` | 查看我今天/本周/指定日期的考勤记录 |
|
||||
| `attendance_team_shift.py` | 查询团队成员本周排班和出勤统计 |
|
||||
| `report_inbox_today.py` | 查看今天收到的日志列表及详情 |
|
||||
|
||||
</details>
|
||||
|
||||
**ISV 集成**:编写您自己的 Agent Skill,与 dws 内置 Skill 搭配构建跨产品工作流:**ISV Skill → dws Skill → 钉钉开放平台 API(强制鉴权 + 全链路审计)**。
|
||||
|
||||
## 功能特性
|
||||
|
||||
<details>
|
||||
<summary><strong>智能输入纠错</strong> — 自动修正 AI 模型常见的参数错误 <code>v1.0.1</code></summary>
|
||||
<summary><strong>智能输入纠错</strong> — 自动修正 AI 模型常见的参数错误</summary>
|
||||
|
||||
内置 Pipeline 纠错引擎,支持命名风格转换、粘连参数拆分、拼写模糊匹配:
|
||||
|
||||
@@ -234,7 +301,7 @@ dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-i
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>jq 过滤 & 字段筛选</strong> — 精确控制输出,减少 token 消耗 <code>v1.0.1</code></summary>
|
||||
<summary><strong>jq 过滤 & 字段筛选</strong> — 精确控制输出,减少 token 消耗</summary>
|
||||
|
||||
```bash
|
||||
# 内置 jq 表达式
|
||||
@@ -248,19 +315,19 @@ dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocati
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Schema 自省</strong> — 调用前查询任意工具的参数结构 <code>v1.0.1</code></summary>
|
||||
<summary><strong>Schema 自省</strong> — 调用前查询任意工具的参数结构</summary>
|
||||
|
||||
```bash
|
||||
dws schema # 列出所有产品和工具
|
||||
dws schema aitable.query_records # 查看参数 Schema
|
||||
dws schema aitable.query_records --jq '.tool.input_schema.required' # 查看必填字段
|
||||
dws schema aitable.query_records --jq '.tool.required' # 查看必填字段
|
||||
dws schema --jq '.products[].id' # 提取所有产品 ID
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>管道 & 文件输入</strong> — 从文件或 stdin 读取 flag 值 <code>v1.0.1</code></summary>
|
||||
<summary><strong>管道 & 文件输入</strong> — 从文件或 stdin 读取 flag 值</summary>
|
||||
|
||||
```bash
|
||||
# 从文件读取消息内容
|
||||
@@ -280,21 +347,22 @@ dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
|
||||
## 核心服务
|
||||
|
||||
| 服务 | 命令 | 描述 |
|
||||
|---------|---------|-------------|
|
||||
| 通讯录 | `contact` | 用户 / 部门 |
|
||||
| 群聊 | `chat` | 群管理 / 群成员 / 机器人消息 / Webhook |
|
||||
| 日历 | `calendar` | 日程 / 会议室 / 闲忙 |
|
||||
| 待办 | `todo` | 任务管理 |
|
||||
| 审批 | `oa` | 流程 / 表单 / 实例 |
|
||||
| 考勤 | `attendance` | 打卡 / 排班 / 统计 |
|
||||
| DING | `ding` | DING 消息 / 发送 / 撤回 |
|
||||
| 日志 | `report` | 日志 / 模版 / 统计 |
|
||||
| 智能表格 | `aitable` | AI 表格操作 |
|
||||
| 工作台 | `workbench` | 应用查询 |
|
||||
| 开发者文档 | `devdoc` | 开放平台文档搜索 |
|
||||
| 服务 | 命令 | 命令数 | 子命令 | 描述 |
|
||||
|------|------|:------:|--------|------|
|
||||
| 通讯录 | `contact` | 6 | `user` `dept` | 按姓名/手机号搜索、批量查询、部门树、当前用户信息 |
|
||||
| 群聊 | `chat` | 10 | `message` `group` `search` | 群增删改查、成员管理、机器人消息、Webhook |
|
||||
| 机器人 | `chat bot` | 6 | `bot` `group` `message` `search` | 机器人创建/搜索、群聊/单聊消息、Webhook、消息撤回 |
|
||||
| 日历 | `calendar` | 13 | `event` `room` `participant` `busy` | 日程增删改查、会议室预订、闲忙查询、参与者管理 |
|
||||
| 待办 | `todo` | 6 | `task` | 创建、列表、修改、完成、详情、删除 |
|
||||
| 审批 | `oa` | 9 | `approval` | 同意/拒绝/撤销、待我审批、我发起的、流程列表 |
|
||||
| 考勤 | `attendance` | 4 | `record` `shift` `summary` `rules` | 打卡记录、排班查询、考勤摘要、考勤组规则 |
|
||||
| DING | `ding` | 2 | `message` | 发送/撤回 DING 消息 |
|
||||
| 日志 | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | 创建日志、收发列表、模版、统计 |
|
||||
| 智能表格 | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | 多维表/数据表/记录/字段全量 CRUD、模板 |
|
||||
| 工作台 | `workbench` | 2 | `app` | 批量查询应用详情 |
|
||||
| 开发者文档 | `devdoc` | 1 | `article` | 搜索开放平台文档与错误码 |
|
||||
|
||||
运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看子命令。
|
||||
> 12 个产品,86 个命令。运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看子命令。
|
||||
|
||||
<details>
|
||||
<summary>即将推出</summary>
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -5,6 +5,7 @@ go 1.25.8
|
||||
require (
|
||||
github.com/fatih/color v1.18.0
|
||||
github.com/google/uuid v1.6.0
|
||||
github.com/itchyny/gojq v0.12.18
|
||||
github.com/spf13/cobra v1.10.2
|
||||
github.com/zalando/go-keyring v0.2.8
|
||||
golang.org/x/crypto v0.49.0
|
||||
@@ -15,7 +16,6 @@ require (
|
||||
require (
|
||||
github.com/danieljoos/wincred v1.2.3 // indirect
|
||||
github.com/godbus/dbus/v5 v5.2.2 // indirect
|
||||
github.com/itchyny/gojq v0.12.18 // indirect
|
||||
github.com/itchyny/timefmt-go v0.1.7 // indirect
|
||||
github.com/mattn/go-colorable v0.1.13 // indirect
|
||||
github.com/mattn/go-isatty v0.0.20 // indirect
|
||||
|
||||
@@ -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()
|
||||
@@ -185,15 +189,31 @@ func newAuthLogoutCommand() *cobra.Command {
|
||||
defer cancel()
|
||||
_ = authpkg.RevokeTokenRemote(revokeCtx)
|
||||
|
||||
// Load token data to get associated clientId before deletion
|
||||
var storedClientID string
|
||||
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
|
||||
storedClientID = tokenData.ClientID
|
||||
}
|
||||
|
||||
if err := authpkg.DeleteTokenData(configDir); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to clear token data: %v", err))
|
||||
}
|
||||
// Clean up associated client secret from keychain
|
||||
if storedClientID != "" {
|
||||
_ = authpkg.DeleteClientSecret(storedClientID)
|
||||
}
|
||||
// Clean up app credentials (app.json + keychain secret)
|
||||
_ = authpkg.DeleteAppConfig(configDir)
|
||||
_ = os.Remove(filepath.Join(configDir, "mcp_url"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token.json"))
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintln(w, "[OK] 已清除所有认证信息")
|
||||
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
|
||||
if !edition.Get().IsEmbedded {
|
||||
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
@@ -223,6 +243,8 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
tokenData = updatedData
|
||||
refreshed = true
|
||||
}
|
||||
} else if edition.Get().AutoPurgeToken {
|
||||
_ = authpkg.DeleteTokenData(configDir)
|
||||
}
|
||||
}
|
||||
if authStatusAuthenticated(tokenData) {
|
||||
@@ -250,7 +272,9 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
}
|
||||
} else {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "状态:", "未登录")
|
||||
fmt.Fprintln(w, "运行 dws auth login 进行登录")
|
||||
if !edition.Get().IsEmbedded {
|
||||
fmt.Fprintln(w, "运行 dws auth login 进行登录")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
@@ -286,6 +310,7 @@ func newAuthExchangeCommand() *cobra.Command {
|
||||
if err != nil {
|
||||
return apperrors.NewAuth(fmt.Sprintf("failed to exchange authorization code: %v", err))
|
||||
}
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
@@ -325,10 +350,13 @@ func newAuthResetCommand() *cobra.Command {
|
||||
}
|
||||
_ = os.Remove(filepath.Join(configDir, "mcp_url"))
|
||||
_ = os.Remove(filepath.Join(configDir, "token"))
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintln(w, "[OK] 认证信息已重置")
|
||||
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
|
||||
if !edition.Get().IsEmbedded {
|
||||
fmt.Fprintln(w, "请运行 dws auth login 重新登录")
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
+43
-7
@@ -29,16 +29,25 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/compat"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/helpers"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cobra.Command {
|
||||
if fn := edition.Get().StaticServers; fn != nil {
|
||||
injectStaticServers(fn())
|
||||
// Static servers provided by the edition hook — skip Market discovery
|
||||
// entirely. The overlay registers its own product commands via
|
||||
// RegisterExtraCommands; we only add the open-source helpers here.
|
||||
commands := helpers.NewPublicCommands(runner)
|
||||
return mergeTopLevelCommands(commands)
|
||||
}
|
||||
|
||||
var commands []*cobra.Command
|
||||
// Generate commands dynamically from the market discovery API.
|
||||
if dynamicCmds := loadDynamicCommands(ctx, runner); len(dynamicCmds) > 0 {
|
||||
commands = append(commands, dynamicCmds...)
|
||||
}
|
||||
@@ -46,6 +55,26 @@ func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cob
|
||||
return mergeTopLevelCommands(commands)
|
||||
}
|
||||
|
||||
// injectStaticServers converts edition.ServerInfo entries into
|
||||
// market.ServerDescriptor and feeds them into SetDynamicServers so the
|
||||
// direct-runtime endpoint resolver can find them.
|
||||
func injectStaticServers(servers []edition.ServerInfo) {
|
||||
descriptors := make([]market.ServerDescriptor, 0, len(servers))
|
||||
for _, s := range servers {
|
||||
descriptors = append(descriptors, market.ServerDescriptor{
|
||||
Key: s.ID,
|
||||
DisplayName: s.Name,
|
||||
Endpoint: s.Endpoint,
|
||||
CLI: market.CLIOverlay{
|
||||
ID: s.ID,
|
||||
Command: s.ID,
|
||||
Prefixes: s.Prefixes,
|
||||
},
|
||||
})
|
||||
}
|
||||
SetDynamicServers(descriptors)
|
||||
}
|
||||
|
||||
// loadDynamicCommands loads the server registry and generates CLI commands
|
||||
// dynamically from CLIOverlay metadata. It consults the disk cache first.
|
||||
// Within the short revalidation window it uses the cached registry directly;
|
||||
@@ -59,7 +88,6 @@ func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cob
|
||||
// Tests may override discoveryBaseURLOverride to redirect to a local server;
|
||||
// in that case the registry cache is always bypassed.
|
||||
func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.Command {
|
||||
|
||||
store := cacheStoreFromEnv()
|
||||
partition := config.DefaultPartition
|
||||
|
||||
@@ -70,13 +98,15 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
|
||||
useCache := strings.TrimSpace(os.Getenv(cli.CatalogFixtureEnv)) == ""
|
||||
|
||||
// --- Cache-first server registry ---
|
||||
cacheLoadStart := time.Now()
|
||||
snapshot, freshness, cacheErr := store.LoadRegistry(partition)
|
||||
RecordTiming(ctx, "registry_cache", time.Since(cacheLoadStart))
|
||||
|
||||
var servers []market.ServerDescriptor
|
||||
now := store.Now().UTC()
|
||||
usingCachedRegistry := useCache && cacheErr == nil && len(snapshot.Servers) > 0
|
||||
|
||||
if usingCachedRegistry {
|
||||
slog.Debug("loadDynamicCommands: using cached registry", "servers", len(snapshot.Servers), "freshness", freshness)
|
||||
servers = snapshot.Servers
|
||||
// Only trigger async revalidation in production (no URL override).
|
||||
// Tests set discoveryBaseURLOverride and control cache expiry directly,
|
||||
@@ -92,9 +122,10 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
|
||||
if discoveryBaseURLOverride != "" {
|
||||
baseURL = discoveryBaseURLOverride
|
||||
}
|
||||
slog.Debug("loadDynamicCommands: fetching servers from market API", "base_url", baseURL)
|
||||
fetchStart := time.Now()
|
||||
client := market.NewClient(baseURL, ipv4OnlyHTTPClient())
|
||||
resp, fetchErr := client.FetchServers(ctx, config.DefaultFetchServersLimit)
|
||||
RecordTiming(ctx, "market_fetch", time.Since(fetchStart))
|
||||
if fetchErr != nil {
|
||||
slog.Debug("loadDynamicCommands: market API fetch failed", "error", fetchErr)
|
||||
// Degrade to stale cache if available (production only).
|
||||
@@ -106,12 +137,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
|
||||
}
|
||||
} else {
|
||||
servers = market.NormalizeServers(resp, "market")
|
||||
slog.Debug("loadDynamicCommands: normalized servers", "count", len(servers))
|
||||
// Persist fresh data (only in non-test mode).
|
||||
if useCache {
|
||||
saveStart := time.Now()
|
||||
if saveErr := store.SaveRegistry(partition, cache.RegistrySnapshot{Servers: servers}); saveErr != nil {
|
||||
slog.Debug("loadDynamicCommands: failed to save registry cache", "error", saveErr)
|
||||
}
|
||||
RecordTiming(ctx, "cache_save", time.Since(saveStart))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -122,9 +154,13 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
|
||||
// Inject dynamic server data for endpoint resolution
|
||||
SetDynamicServers(servers)
|
||||
|
||||
detailStart := time.Now()
|
||||
detailsByID := loadCachedDetailsFast(store, servers)
|
||||
RecordTiming(ctx, "tool_metadata", time.Since(detailStart))
|
||||
|
||||
buildStart := time.Now()
|
||||
cmds := compat.BuildDynamicCommands(servers, runner, detailsByID)
|
||||
slog.Debug("loadDynamicCommands: built dynamic commands", "commands", len(cmds))
|
||||
RecordTiming(ctx, "build_commands", time.Since(buildStart))
|
||||
|
||||
return cmds
|
||||
}
|
||||
|
||||
@@ -0,0 +1,669 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newPluginCommand() *cobra.Command {
|
||||
pluginCmd := newPlaceholderParent("plugin", "Manage plugins")
|
||||
|
||||
pluginCmd.AddCommand(
|
||||
newPluginListCommand(),
|
||||
newPluginInstallCommand(),
|
||||
newPluginInfoCommand(),
|
||||
newPluginEnableCommand(),
|
||||
newPluginDisableCommand(),
|
||||
newPluginRemoveCommand(),
|
||||
newPluginValidateCommand(),
|
||||
newPluginCreateCommand(),
|
||||
newPluginDevCommand(),
|
||||
newPluginConfigCommand(),
|
||||
newPluginBuildCommand(),
|
||||
)
|
||||
|
||||
return pluginCmd
|
||||
}
|
||||
|
||||
func newPluginListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "List installed plugins",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
if wantJSON {
|
||||
return output.WriteJSON(cmd.OutOrStdout(), plugins)
|
||||
}
|
||||
|
||||
if len(plugins) == 0 {
|
||||
fmt.Fprintln(cmd.OutOrStdout(), "No plugins installed.")
|
||||
return nil
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
"NAME", "VERSION", "TYPE", "STATUS", "DESCRIPTION")
|
||||
fmt.Fprintln(w, strings.Repeat("-", 85))
|
||||
for _, p := range plugins {
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
p.Name, p.Version, p.Type, statusStr(p.Enabled), p.Description)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install",
|
||||
Short: "Install a plugin",
|
||||
Example: ` dws plugin install --dir ./conference
|
||||
dws plugin install --git https://github.com/DingTalk-Real-AI/conference.git`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dirPath, _ := cmd.Flags().GetString("dir")
|
||||
gitURL, _ := cmd.Flags().GetString("git")
|
||||
|
||||
if dirPath == "" && gitURL == "" {
|
||||
return apperrors.NewValidation("specify install source: --dir <path> or --git <url>")
|
||||
}
|
||||
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if gitURL != "" {
|
||||
p, err := loader.InstallFromGit(gitURL)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
}
|
||||
|
||||
p, err := loader.InstallFromDir(dirPath)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("dir", "", "Install from a local directory")
|
||||
cmd.Flags().String("git", "", "Install from a Git repository")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInfoCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "info <name>",
|
||||
Short: "Show plugin details",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
for _, p := range plugins {
|
||||
if p.Name == name {
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Name: %s\n", p.Name)
|
||||
fmt.Fprintf(w, "Version: %s\n", p.Version)
|
||||
fmt.Fprintf(w, "Type: %s\n", p.Type)
|
||||
fmt.Fprintf(w, "Status: %s\n", statusStr(p.Enabled))
|
||||
fmt.Fprintf(w, "Path: %s\n", p.Path)
|
||||
if p.Description != "" {
|
||||
fmt.Fprintf(w, "Description: %s\n", p.Description)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found", name))
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginEnableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "enable <name>",
|
||||
Short: "Enable a plugin",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], true); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s enabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginDisableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "disable <name>",
|
||||
Short: "Disable a plugin (managed plugins can be disabled but not removed)",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], false); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s disabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginRemoveCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "remove <name>",
|
||||
Short: "Remove a user plugin (managed plugins cannot be removed)",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
keepData, _ := cmd.Flags().GetBool("keep-data")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.RemovePlugin(args[0], keepData); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s removed.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("keep-data", false, "Keep plugin data directory")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginValidateCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "validate <dir>",
|
||||
Short: "Validate a plugin.json",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
m, err := plugin.ParseManifest(dir + "/plugin.json")
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("parse failed: %v", err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Valid: %s (%s)\n", m.Name, m.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginCreateCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create <name>",
|
||||
Short: "Scaffold a new plugin directory",
|
||||
Example: ` dws plugin create my-tool
|
||||
dws plugin create my-tool --type managed --description "My awesome tool"`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
desc, _ := cmd.Flags().GetString("description")
|
||||
pluginType, _ := cmd.Flags().GetString("type")
|
||||
|
||||
if pluginType == "" {
|
||||
pluginType = "user"
|
||||
}
|
||||
if pluginType != "managed" && pluginType != "user" {
|
||||
return apperrors.NewValidation("type must be 'managed' or 'user'")
|
||||
}
|
||||
|
||||
// Validate name format
|
||||
m := &plugin.Manifest{Name: name, Version: "0.1.0", Type: pluginType}
|
||||
if err := m.Validate(""); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin name: %v", err))
|
||||
}
|
||||
|
||||
dir := filepath.Join(".", name)
|
||||
if _, err := os.Stat(dir); err == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("directory %q already exists", dir))
|
||||
}
|
||||
|
||||
// Create directory structure
|
||||
dirs := []string{
|
||||
dir,
|
||||
filepath.Join(dir, "skills", name),
|
||||
filepath.Join(dir, "hooks"),
|
||||
}
|
||||
for _, d := range dirs {
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
|
||||
}
|
||||
}
|
||||
|
||||
// Write plugin.json
|
||||
pluginJSON := fmt.Sprintf(`{
|
||||
"name": %q,
|
||||
"version": "0.1.0",
|
||||
"description": %q,
|
||||
"type": %q,
|
||||
"minCLIVersion": %q,
|
||||
"mcpServers": {
|
||||
%q: {
|
||||
"type": "stdio",
|
||||
"command": "${DWS_PLUGIN_ROOT}/bin/server",
|
||||
"args": []
|
||||
}
|
||||
},
|
||||
"build": {
|
||||
"command": "echo 'TODO: replace with your build command, e.g.: bun build --compile src/server.ts --outfile bin/server'",
|
||||
"output": "bin/server"
|
||||
},
|
||||
"skills": "./skills/",
|
||||
"hooks": "./hooks/hooks.json"
|
||||
}
|
||||
`, name, desc, pluginType, RawVersion(), name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(pluginJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write plugin.json: %v", err))
|
||||
}
|
||||
|
||||
// Write SKILL.md template
|
||||
skillMD := fmt.Sprintf(`---
|
||||
name: %s
|
||||
description: %s
|
||||
cli_version: ">=%s"
|
||||
---
|
||||
|
||||
# %s
|
||||
|
||||
## Intent Recognition
|
||||
|
||||
Use this skill when the user mentions:
|
||||
- TODO: add your intent keywords here
|
||||
|
||||
## Command Decision Tree
|
||||
|
||||
| User Intent | Command | Required Parameters |
|
||||
|-------------|---------|---------------------|
|
||||
| TODO | `+"`dws %s <sub-command>`"+` | `+"`--param`"+` |
|
||||
|
||||
## Parameter Rules
|
||||
|
||||
### TODO: parameter type
|
||||
- Format description
|
||||
- Conversion rules
|
||||
`, name, desc, RawVersion(), name, name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "skills", name, "SKILL.md"), []byte(skillMD), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write SKILL.md: %v", err))
|
||||
}
|
||||
|
||||
// Write hooks.json template
|
||||
hooksJSON := `{
|
||||
"hooks": []
|
||||
}
|
||||
`
|
||||
if err := os.WriteFile(filepath.Join(dir, "hooks", "hooks.json"), []byte(hooksJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write hooks.json: %v", err))
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Created plugin scaffold at ./%s/\n", name)
|
||||
fmt.Fprintf(w, " %s/\n", name)
|
||||
fmt.Fprintf(w, " ├── plugin.json\n")
|
||||
fmt.Fprintf(w, " ├── skills/%s/SKILL.md\n", name)
|
||||
fmt.Fprintf(w, " └── hooks/hooks.json\n")
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "Next steps:\n")
|
||||
fmt.Fprintf(w, " 1. Edit plugin.json to configure your MCP servers\n")
|
||||
fmt.Fprintf(w, " 2. Edit skills/%s/SKILL.md to describe your commands\n", name)
|
||||
fmt.Fprintf(w, " 3. Run: dws plugin validate ./%s\n", name)
|
||||
fmt.Fprintf(w, " 4. Run: dws plugin dev ./%s\n", name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("description", "", "Plugin description")
|
||||
cmd.Flags().String("type", "user", "Plugin type: managed or user")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginDevCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "dev <dir>",
|
||||
Short: "Register a local directory as a dev plugin",
|
||||
Long: `Registers a plugin from a local source directory for development.
|
||||
The plugin is loaded directly from the source directory on next CLI invocation,
|
||||
without copying files to ~/.dws/plugins/. Use 'dws plugin dev --off <name>'
|
||||
to unregister.`,
|
||||
Example: ` dws plugin dev ./my-tool
|
||||
dws plugin dev --off my-tool`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
off, _ := cmd.Flags().GetBool("off")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if off {
|
||||
// Unregister dev plugin
|
||||
name := args[0]
|
||||
if err := loader.UnregisterDevPlugin(name); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q unregistered.\n", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Register dev plugin
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
// Validate the plugin first
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
|
||||
if err := loader.RegisterDevPlugin(m.Name, absDir); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to register: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q registered from %s\n", m.Name, absDir)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "It will be loaded on next dws invocation.\n")
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "To unregister: dws plugin dev --off %s\n", m.Name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("off", false, "Unregister a dev plugin")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigCommand() *cobra.Command {
|
||||
configCmd := newPlaceholderParent("config", "Manage plugin configuration")
|
||||
configCmd.AddCommand(
|
||||
newPluginConfigSetCommand(),
|
||||
newPluginConfigGetCommand(),
|
||||
newPluginConfigListCommand(),
|
||||
newPluginConfigUnsetCommand(),
|
||||
)
|
||||
return configCmd
|
||||
}
|
||||
|
||||
func newPluginConfigSetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "set <plugin-name> <key> <value>",
|
||||
Short: "Set a plugin config value",
|
||||
Long: `Persistently set a configuration value for a plugin.
|
||||
The value is stored in ~/.dws/settings.json and automatically injected
|
||||
as an environment variable when the plugin is loaded.
|
||||
|
||||
Environment variables set by the user (e.g. via export) take precedence
|
||||
over values stored in settings.json.`,
|
||||
Example: ` dws plugin config set demo-devtool DASHSCOPE_API_KEY sk-xxx
|
||||
dws plugin config set my-plugin API_ENDPOINT https://api.example.com`,
|
||||
Args: cobra.ExactArgs(3),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key, value := args[0], args[1], args[2]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// Validate that the plugin exists.
|
||||
plugins := loader.ListInstalled()
|
||||
found := false
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found; use 'dws plugin list' to see installed plugins", pluginName))
|
||||
}
|
||||
|
||||
loader.SetPluginConfig(pluginName, key, value)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config saved: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigGetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "get <plugin-name> <key>",
|
||||
Short: "Get a plugin config value",
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
val, ok := loader.GetPluginConfig(pluginName, key)
|
||||
if !ok {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintln(cmd.OutOrStdout(), val)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list <plugin-name>",
|
||||
Short: "List all config values for a plugin",
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
configs := loader.ListPluginConfig(pluginName)
|
||||
|
||||
// Also load the plugin manifest to show declared userConfig keys.
|
||||
declaredKeys := loadDeclaredUserConfig(loader, pluginName)
|
||||
|
||||
if wantJSON {
|
||||
result := make(map[string]any)
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
if sensitive {
|
||||
result[k] = maskSensitiveValue(v)
|
||||
} else {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
// Include declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
entry := map[string]any{
|
||||
"value": nil,
|
||||
"description": ci.Description,
|
||||
"required": ci.Default == "",
|
||||
}
|
||||
result[k] = entry
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "plugin_config",
|
||||
"plugin": pluginName,
|
||||
"config": result,
|
||||
})
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
if len(configs) == 0 && len(declaredKeys) == 0 {
|
||||
fmt.Fprintf(w, "No configuration for plugin %q.\n", pluginName)
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "Configuration for %s:\n\n", pluginName)
|
||||
|
||||
// Show set values.
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
displayVal := v
|
||||
if sensitive {
|
||||
displayVal = maskSensitiveValue(v)
|
||||
}
|
||||
fmt.Fprintf(w, " %s = %s\n", k, displayVal)
|
||||
}
|
||||
|
||||
// Show declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
desc := ""
|
||||
if ci.Description != "" {
|
||||
desc = " # " + ci.Description
|
||||
}
|
||||
fmt.Fprintf(w, " %s = (not set)%s\n", k, desc)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigUnsetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "unset <plugin-name> <key>",
|
||||
Short: "Remove a plugin config value",
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if !loader.UnsetPluginConfig(pluginName, key) {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config removed: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// loadDeclaredUserConfig loads the userConfig section from a plugin's manifest.
|
||||
func loadDeclaredUserConfig(loader *plugin.Loader, pluginName string) map[string]plugin.ConfigItem {
|
||||
plugins := loader.ListInstalled()
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
m, err := plugin.ParseManifest(filepath.Join(p.Path, "plugin.json"))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return m.UserConfig
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// maskSensitiveValue masks a sensitive value, showing only the first 4
|
||||
// and last 2 characters for values longer than 8 characters.
|
||||
func maskSensitiveValue(value string) string {
|
||||
if len(value) <= 8 {
|
||||
return strings.Repeat("*", len(value))
|
||||
}
|
||||
return value[:4] + strings.Repeat("*", len(value)-6) + value[len(value)-2:]
|
||||
}
|
||||
|
||||
func newPluginBuildCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "build <dir>",
|
||||
Short: "Build plugin's stdio server into a native binary",
|
||||
Long: `Runs the build command declared in plugin.json to compile the
|
||||
plugin's server into a single executable. This ensures plugin users
|
||||
don't need any language runtime (Node.js, Python, etc.) installed.
|
||||
|
||||
The build configuration is read from the "build" field in plugin.json:
|
||||
|
||||
{
|
||||
"build": {
|
||||
"command": "bun build --compile src/server.ts --outfile bin/server",
|
||||
"output": "bin/server"
|
||||
}
|
||||
}`,
|
||||
Example: ` dws plugin build ./my-plugin
|
||||
dws plugin build .`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
|
||||
if m.Build == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf(
|
||||
"plugin %q has no \"build\" field in plugin.json.\n"+
|
||||
"Add a build config, e.g.:\n\n"+
|
||||
" \"build\": {\n"+
|
||||
" \"command\": \"bun build --compile src/server.js --outfile bin/server\",\n"+
|
||||
" \"output\": \"bin/server\"\n"+
|
||||
" }", m.Name))
|
||||
}
|
||||
|
||||
if err := plugin.BuildPlugin(absDir); err != nil {
|
||||
return apperrors.NewInternal(err.Error())
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Build succeeded: %s\n", m.Build.Output)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func statusStr(enabled bool) string {
|
||||
if enabled {
|
||||
return "enabled"
|
||||
}
|
||||
return "disabled"
|
||||
}
|
||||
@@ -261,7 +261,7 @@ func TestExecuteWritesRecoveryEventIDToStderrOnCapturedFailure(t *testing.T) {
|
||||
|
||||
oldArgs := os.Args
|
||||
defer func() { os.Args = oldArgs }()
|
||||
os.Args = []string{"dws", "mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`}
|
||||
os.Args = []string{"dws", "mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"}
|
||||
|
||||
stdoutR, stdoutW, err := os.Pipe()
|
||||
if err != nil {
|
||||
|
||||
+650
-20
@@ -15,21 +15,25 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/signal"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/compat"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
@@ -39,8 +43,11 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline/handlers"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
)
|
||||
@@ -51,12 +58,23 @@ const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
|
||||
|
||||
// Execute runs the root command and returns the process exit code.
|
||||
func Execute() int {
|
||||
timing := NewTimingCollector()
|
||||
defer func() {
|
||||
timing.PrintIfEnabled()
|
||||
timing.WriteReportIfEnabled(RawVersion(), SanitizeCommand(os.Args))
|
||||
}()
|
||||
|
||||
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||
defer cancel()
|
||||
|
||||
// Attach timing collector to context for use by child components
|
||||
ctx = WithTimingCollector(ctx, timing)
|
||||
|
||||
initStart := time.Now()
|
||||
recovery.ResetRuntimeState()
|
||||
engine := newPipelineEngine()
|
||||
root := NewRootCommandWithEngine(ctx, engine)
|
||||
timing.Record("cmd_init", time.Since(initStart))
|
||||
|
||||
// Run PreParse handlers on raw argv before Cobra parses flags.
|
||||
// This corrects model-generated errors like --userId → --user-id
|
||||
@@ -114,10 +132,29 @@ func flagErrorWithSuggestions(cmd *cobra.Command, err error) error {
|
||||
}
|
||||
|
||||
func printExecutionError(root *cobra.Command, stdout, stderr io.Writer, err error) error {
|
||||
var raw apperrors.RawStderrError
|
||||
if stderrors.As(err, &raw) {
|
||||
_, writeErr := fmt.Fprintln(stderr, raw.RawStderr())
|
||||
return writeErr
|
||||
}
|
||||
if wantsJSONErrors(root) {
|
||||
return apperrors.PrintJSON(stdout, err)
|
||||
}
|
||||
return apperrors.PrintHuman(stderr, err)
|
||||
return apperrors.PrintHumanAt(stderr, err, resolveVerbosity(root))
|
||||
}
|
||||
|
||||
// resolveVerbosity derives the error verbosity level from the root command's flags.
|
||||
func resolveVerbosity(cmd *cobra.Command) apperrors.Verbosity {
|
||||
if cmd == nil {
|
||||
return apperrors.VerbosityNormal
|
||||
}
|
||||
if debug, err := cmd.Flags().GetBool("debug"); err == nil && debug {
|
||||
return apperrors.VerbosityDebug
|
||||
}
|
||||
if verbose, err := cmd.Flags().GetBool("verbose"); err == nil && verbose {
|
||||
return apperrors.VerbosityVerbose
|
||||
}
|
||||
return apperrors.VerbosityNormal
|
||||
}
|
||||
|
||||
func wantsJSONErrors(root *cobra.Command) bool {
|
||||
@@ -192,6 +229,7 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
AuthTokenFunc: func(ctx context.Context) string {
|
||||
return resolveRuntimeAuthToken(ctx, "")
|
||||
},
|
||||
LoggerFunc: FileLoggerInstance,
|
||||
}
|
||||
runner := newCommandRunnerWithFlags(loader, flags)
|
||||
|
||||
@@ -218,7 +256,13 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
// Configure global slog level based on --debug / --verbose flags.
|
||||
configureLogLevel(flags)
|
||||
|
||||
return configureOutputSink(cmd)
|
||||
if err := configureOutputSink(cmd); err != nil {
|
||||
return err
|
||||
}
|
||||
if fn := edition.Get().AfterPersistentPreRun; fn != nil {
|
||||
return fn(cmd, args)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
PersistentPostRunE: func(cmd *cobra.Command, args []string) error {
|
||||
CloseFileLogger()
|
||||
@@ -236,18 +280,38 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
|
||||
utilityCommands := []*cobra.Command{
|
||||
newAuthCommand(),
|
||||
newSkillCommand(),
|
||||
newCacheCommand(),
|
||||
newConfigCommand(),
|
||||
newDoctorCommand(),
|
||||
newCompletionCommand(root),
|
||||
newRecoveryCommand(rootCtx, loader, flags),
|
||||
newUpgradeCommand(),
|
||||
newVersionCommand(),
|
||||
newPluginCommand(),
|
||||
schemaCmd,
|
||||
genSkillsCmd,
|
||||
mcpCmd,
|
||||
}
|
||||
root.AddCommand(utilityCommands...)
|
||||
|
||||
root.AddCommand(newLegacyPublicCommands(rootCtx, runner)...)
|
||||
root.AddCommand(newLegacyHiddenCommands(runner)...)
|
||||
|
||||
// --- Plugin loading: runs AFTER legacy commands so that
|
||||
// AppendDynamicServer adds plugin endpoints on top of Market
|
||||
// endpoints (SetDynamicServers is called inside loadDynamicCommands).
|
||||
pluginCmds := loadPlugins(engine, runner)
|
||||
if len(pluginCmds) > 0 {
|
||||
addPluginCommandsSafe(root, pluginCmds)
|
||||
}
|
||||
|
||||
if fn := edition.Get().RegisterExtraCommands; fn != nil {
|
||||
caller := newToolCallerAdapter(runner, flags)
|
||||
fn(root, caller)
|
||||
deduplicateCommands(root)
|
||||
}
|
||||
|
||||
hideNonDirectRuntimeCommands(root)
|
||||
configureRootHelp(root)
|
||||
// Set custom flag error handler for better UX
|
||||
@@ -261,6 +325,10 @@ func newAuthCommand() *cobra.Command {
|
||||
return buildAuthCommand()
|
||||
}
|
||||
|
||||
func newSkillCommand() *cobra.Command {
|
||||
return buildSkillCommand()
|
||||
}
|
||||
|
||||
func newCacheCommand() *cobra.Command {
|
||||
cacheCmd := newPlaceholderParent("cache", "缓存管理")
|
||||
|
||||
@@ -422,24 +490,51 @@ func newVersionCommand() *cobra.Command {
|
||||
Example: " dws version\n dws version --format json",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
format, err := cmd.Flags().GetString("format")
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read format flag")
|
||||
wantJSON := cmd.Flags().Changed("format")
|
||||
if wantJSON {
|
||||
format, _ := cmd.Flags().GetString("format")
|
||||
wantJSON = (format == "json")
|
||||
}
|
||||
payload := map[string]any{
|
||||
"version": Version(),
|
||||
"go": "1.24+",
|
||||
|
||||
editionName := edition.Get().Name
|
||||
if editionName == "" {
|
||||
editionName = "open"
|
||||
}
|
||||
if format == "json" {
|
||||
ver := RawVersion()
|
||||
bt := BuildTime()
|
||||
gc := GitCommit()
|
||||
goVer := "1.24+"
|
||||
|
||||
arch := "MCP Dynamic Aggregation"
|
||||
|
||||
if wantJSON {
|
||||
payload := map[string]any{
|
||||
"version": ver,
|
||||
"edition": editionName,
|
||||
"architecture": arch,
|
||||
"go": goVer,
|
||||
}
|
||||
if bt != "unknown" {
|
||||
payload["build"] = bt
|
||||
}
|
||||
if gc != "unknown" {
|
||||
payload["commit"] = gc
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), payload)
|
||||
}
|
||||
_, err = fmt.Fprintf(
|
||||
cmd.OutOrStdout(),
|
||||
"版本: %s\nGo: %s\n",
|
||||
Version(),
|
||||
"1.24+",
|
||||
)
|
||||
return err
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Version:", ver)
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Edition:", editionName)
|
||||
if bt != "unknown" {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Build:", bt)
|
||||
}
|
||||
if gc != "unknown" {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Commit:", gc)
|
||||
}
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Architecture:", arch)
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Go:", goVer)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -536,17 +631,34 @@ func newMCPCommand(ctx context.Context, loader cli.CatalogLoader, runner executo
|
||||
}
|
||||
|
||||
// hideNonDirectRuntimeCommands marks top-level product commands as hidden
|
||||
// unless they correspond to a product discovered via dynamic server discovery.
|
||||
// unless they correspond to a product discovered via dynamic server discovery
|
||||
// or listed in the edition's VisibleProducts hook.
|
||||
// Public utility commands (auth, cache, completion, version) are always kept
|
||||
// visible; explicitly hidden commands stay hidden.
|
||||
func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
allowedProducts := DirectRuntimeProductIDs()
|
||||
var allowedProducts map[string]bool
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
products := fn()
|
||||
allowedProducts = make(map[string]bool, len(products))
|
||||
for _, p := range products {
|
||||
allowedProducts[p] = true
|
||||
}
|
||||
} else {
|
||||
allowedProducts = DirectRuntimeProductIDs()
|
||||
}
|
||||
staticCommands := map[string]bool{
|
||||
"auth": true,
|
||||
"cache": true,
|
||||
"config": true,
|
||||
"doctor": true,
|
||||
"completion": true,
|
||||
"skill": true,
|
||||
"plugin": true,
|
||||
"version": true,
|
||||
"help": true,
|
||||
"recovery": true,
|
||||
"schema": true,
|
||||
"mcp": true,
|
||||
}
|
||||
for _, cmd := range root.Commands() {
|
||||
name := cmd.Name()
|
||||
@@ -563,6 +675,84 @@ func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
}
|
||||
}
|
||||
|
||||
// reservedCommands is the set of built-in command names that plugins must
|
||||
// not override. This protects core CLI functionality from being hijacked
|
||||
// by a malicious or misconfigured plugin.
|
||||
var reservedCommands = map[string]bool{
|
||||
"auth": true, "login": true, "logout": true,
|
||||
"plugin": true, "skill": true, "cache": true,
|
||||
"config": true, "doctor": true, "completion": true,
|
||||
"recovery": true, "upgrade": true, "version": true,
|
||||
"schema": true, "mcp": true, "help": true,
|
||||
}
|
||||
|
||||
// addPluginCommandsSafe registers plugin commands with conflict detection.
|
||||
//
|
||||
// Rules:
|
||||
// - Plugin vs reserved (auth/plugin/cache/...) → reject, warn
|
||||
// - Plugin vs plugin (same name) → reject later one, warn
|
||||
// - Plugin vs Market dynamic command → allow, plugin wins
|
||||
func addPluginCommandsSafe(root *cobra.Command, pluginCmds []*cobra.Command) {
|
||||
// Build index of existing commands before plugin registration.
|
||||
existing := make(map[string]bool)
|
||||
for _, cmd := range root.Commands() {
|
||||
existing[cmd.Name()] = true
|
||||
}
|
||||
|
||||
pluginSeen := make(map[string]bool)
|
||||
|
||||
for _, cmd := range pluginCmds {
|
||||
name := cmd.Name()
|
||||
|
||||
// Rule 1: never override reserved built-in commands.
|
||||
if reservedCommands[name] {
|
||||
slog.Warn("plugin: command name conflicts with built-in command, skipping",
|
||||
"command", name)
|
||||
continue
|
||||
}
|
||||
|
||||
// Rule 2: plugin vs plugin — first plugin wins.
|
||||
if pluginSeen[name] {
|
||||
slog.Warn("plugin: duplicate command from another plugin, skipping",
|
||||
"command", name)
|
||||
continue
|
||||
}
|
||||
pluginSeen[name] = true
|
||||
|
||||
// Rule 3: plugin vs Market — plugin wins, remove the old one.
|
||||
if existing[name] {
|
||||
for _, old := range root.Commands() {
|
||||
if old.Name() == name {
|
||||
root.RemoveCommand(old)
|
||||
slog.Debug("plugin: overriding Market command",
|
||||
"command", name)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
root.AddCommand(cmd)
|
||||
}
|
||||
}
|
||||
|
||||
// deduplicateCommands removes duplicate top-level commands, keeping the last
|
||||
// registered one. This ensures overlay commands take precedence over
|
||||
// open-source defaults when both register the same product name.
|
||||
func deduplicateCommands(root *cobra.Command) {
|
||||
seen := make(map[string]*cobra.Command)
|
||||
var dups []*cobra.Command
|
||||
for _, cmd := range root.Commands() {
|
||||
name := cmd.Name()
|
||||
if prev, ok := seen[name]; ok {
|
||||
dups = append(dups, prev)
|
||||
}
|
||||
seen[name] = cmd
|
||||
}
|
||||
for _, dup := range dups {
|
||||
root.RemoveCommand(dup)
|
||||
}
|
||||
}
|
||||
|
||||
func cacheStoreFromEnv() *cache.Store {
|
||||
cacheDir := strings.TrimSpace(os.Getenv(cli.CacheDirEnv))
|
||||
return cache.NewStore(cacheDir)
|
||||
@@ -824,11 +1014,445 @@ func CloseFileLogger() {
|
||||
}
|
||||
}
|
||||
|
||||
// loadPlugins scans plugin directories, injects their MCP servers into
|
||||
// the dynamic server registry, and registers their pipeline hooks.
|
||||
// This runs before legacy command construction so that plugin servers
|
||||
// are available for EnvironmentLoader.Load().
|
||||
func loadPlugins(engine *pipeline.Engine, runner executor.Runner) []*cobra.Command {
|
||||
pluginLoader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// 0a. Inject plugin config values from settings.json as environment
|
||||
// variables so that expandPluginVars can resolve ${KEY} references
|
||||
// in plugin.json headers, endpoints, etc. User-set env vars take
|
||||
// precedence (InjectPluginConfigEnv skips already-set keys).
|
||||
pluginLoader.InjectPluginConfigEnv()
|
||||
|
||||
// 0a. Ensure default managed plugins are installed (first-run bootstrap).
|
||||
updater := plugin.NewUpdater(pluginLoader.PluginsDir, RawVersion())
|
||||
accessToken, tokenErr := loadSkillAccessToken()
|
||||
if tokenErr == nil && accessToken != "" {
|
||||
bootstrapCtx, bootstrapCancel := context.WithTimeout(context.Background(), 2*time.Minute)
|
||||
installed := updater.EnsureManaged(bootstrapCtx, accessToken, os.Stderr)
|
||||
bootstrapCancel()
|
||||
if len(installed) > 0 {
|
||||
slog.Debug("plugin: bootstrapped managed plugins", "names", installed)
|
||||
}
|
||||
|
||||
// 0b. Check for managed plugin updates (non-blocking, best-effort).
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
updated := updater.CheckAndUpdate(ctx, accessToken, os.Stderr)
|
||||
cancel()
|
||||
if len(updated) > 0 {
|
||||
slog.Debug("plugin: updated managed plugins", "names", updated)
|
||||
}
|
||||
}
|
||||
|
||||
// 1. Load official plugins (always enabled)
|
||||
managedPlugins := pluginLoader.LoadManaged()
|
||||
|
||||
// 2. Load user plugins (per settings.json)
|
||||
userPlugins := pluginLoader.LoadUser()
|
||||
|
||||
// 3. Load dev plugins (registered via `dws plugin dev`)
|
||||
devPlugins := pluginLoader.LoadDev()
|
||||
|
||||
allPlugins := append(managedPlugins, userPlugins...)
|
||||
allPlugins = append(allPlugins, devPlugins...)
|
||||
|
||||
// 3. Discover tools from streamable-http servers and build CLI commands.
|
||||
// Third-party servers with auth headers are discovered in parallel
|
||||
// to avoid sequential 10s timeouts when multiple remote servers exist.
|
||||
var pluginCmds []*cobra.Command
|
||||
tc := transport.NewClient(nil)
|
||||
|
||||
// Collect all server descriptors and register auth first (fast, no I/O).
|
||||
type pluginServer struct {
|
||||
plugin *plugin.Plugin
|
||||
srv market.ServerDescriptor
|
||||
}
|
||||
var httpServers []pluginServer
|
||||
|
||||
for _, p := range allPlugins {
|
||||
for _, srv := range p.ToServerDescriptors() {
|
||||
AppendDynamicServer(srv)
|
||||
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
}
|
||||
|
||||
if srv.HasCLIMeta {
|
||||
httpServers = append(httpServers, pluginServer{plugin: p, srv: srv})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Discover tools from HTTP servers in parallel when there are multiple
|
||||
// servers with auth headers (third-party services with higher latency).
|
||||
if len(httpServers) > 1 {
|
||||
type discoveryResult struct {
|
||||
commands []*cobra.Command
|
||||
}
|
||||
results := make([]discoveryResult, len(httpServers))
|
||||
var wg sync.WaitGroup
|
||||
for i, ps := range httpServers {
|
||||
wg.Add(1)
|
||||
go func(idx int, ps pluginServer) {
|
||||
defer wg.Done()
|
||||
results[idx].commands = registerHTTPServer(ps.plugin, ps.srv, tc, runner)
|
||||
}(i, ps)
|
||||
}
|
||||
wg.Wait()
|
||||
for _, r := range results {
|
||||
pluginCmds = append(pluginCmds, r.commands...)
|
||||
}
|
||||
} else {
|
||||
for _, ps := range httpServers {
|
||||
cmds := registerHTTPServer(ps.plugin, ps.srv, tc, runner)
|
||||
pluginCmds = append(pluginCmds, cmds...)
|
||||
}
|
||||
}
|
||||
|
||||
// 4. Start stdio MCP servers, discover tools, and build CLI commands
|
||||
for _, p := range allPlugins {
|
||||
for _, sc := range p.StdioClients() {
|
||||
// Use background context so the subprocess lives for the CLI
|
||||
// process lifetime (not killed by a short timeout).
|
||||
if err := sc.Client.Start(context.Background()); err != nil {
|
||||
slog.Warn("plugin: failed to start stdio server",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
continue
|
||||
}
|
||||
cmds := registerStdioServer(p, sc, runner)
|
||||
pluginCmds = append(pluginCmds, cmds...)
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Register plugin hooks into pipeline engine
|
||||
if engine != nil {
|
||||
for _, p := range allPlugins {
|
||||
hooksCfg, err := p.LoadHooks()
|
||||
if err != nil {
|
||||
slog.Warn("plugin: failed to load hooks",
|
||||
"plugin", p.Manifest.Name, "error", err)
|
||||
continue
|
||||
}
|
||||
if hooksCfg == nil {
|
||||
continue
|
||||
}
|
||||
for _, entry := range hooksCfg.Hooks {
|
||||
engine.Register(plugin.NewHookAdapter(p.Manifest.Name, entry))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Sync plugin skills to agent directories
|
||||
plugin.SyncSkills(allPlugins)
|
||||
|
||||
if len(allPlugins) > 0 {
|
||||
slog.Debug("plugins loaded",
|
||||
"managed", len(managedPlugins),
|
||||
"user", len(userPlugins),
|
||||
"dev", len(devPlugins),
|
||||
)
|
||||
}
|
||||
|
||||
return pluginCmds
|
||||
}
|
||||
|
||||
// registerHTTPServer discovers tools from a streamable-http MCP server and
|
||||
// builds CLI commands. Used for plugin-owned HTTP servers that provide CLI metadata.
|
||||
//
|
||||
// When the server descriptor carries AuthHeaders (from plugin.json "headers"),
|
||||
// a dedicated transport.Client is created with the plugin's Bearer token and
|
||||
// trusted domains so that third-party MCP servers requiring independent
|
||||
// authentication (e.g. Alibaba Cloud Bailian) can be discovered at startup.
|
||||
func registerHTTPServer(p *plugin.Plugin, srv market.ServerDescriptor, tc *transport.Client, runner executor.Runner) []*cobra.Command {
|
||||
// Use a longer timeout for servers with custom auth headers (third-party
|
||||
// services may have higher latency than local/DingTalk endpoints).
|
||||
timeout := 2 * time.Second
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
timeout = 10 * time.Second
|
||||
}
|
||||
ctx, cancel := context.WithTimeout(context.Background(), timeout)
|
||||
defer cancel()
|
||||
|
||||
// If the plugin provides custom auth headers, create a dedicated client
|
||||
// so the Bearer token is sent to the third-party endpoint.
|
||||
discoveryClient := tc
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
discoveryClient = buildPluginAuthClient(tc, srv)
|
||||
}
|
||||
|
||||
if _, err := discoveryClient.Initialize(ctx, srv.Endpoint); err != nil {
|
||||
slog.Debug("plugin: http server offline, skipping tool discovery",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
toolsResult, err := discoveryClient.ListTools(ctx, srv.Endpoint)
|
||||
if err != nil {
|
||||
slog.Debug("plugin: http ListTools failed",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(toolsResult.Tools) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
detailsByID := make(map[string][]market.DetailTool)
|
||||
var detailTools []market.DetailTool
|
||||
for _, tool := range toolsResult.Tools {
|
||||
schemaJSON := ""
|
||||
if tool.InputSchema != nil {
|
||||
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
|
||||
schemaJSON = string(data)
|
||||
}
|
||||
}
|
||||
detailTools = append(detailTools, market.DetailTool{
|
||||
ToolName: tool.Name,
|
||||
ToolTitle: tool.Title,
|
||||
ToolDesc: tool.Description,
|
||||
IsSensitive: tool.Sensitive,
|
||||
ToolRequest: schemaJSON,
|
||||
})
|
||||
}
|
||||
detailsByID[strings.TrimSpace(srv.CLI.ID)] = detailTools
|
||||
|
||||
// If the server has no ToolOverrides (e.g. third-party MCP servers that
|
||||
// only declare cli.id and cli.command), auto-generate one override per
|
||||
// discovered tool so BuildDynamicCommands can create leaf commands.
|
||||
if len(srv.CLI.ToolOverrides) == 0 && len(toolsResult.Tools) > 0 {
|
||||
srv.CLI.ToolOverrides = make(map[string]market.CLIToolOverride, len(toolsResult.Tools))
|
||||
for _, tool := range toolsResult.Tools {
|
||||
srv.CLI.ToolOverrides[tool.Name] = market.CLIToolOverride{
|
||||
CLIName: deriveToolCLIName(tool.Name),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
cmds := compat.BuildDynamicCommands(
|
||||
[]market.ServerDescriptor{srv}, runner, detailsByID)
|
||||
|
||||
slog.Debug("plugin: http server registered",
|
||||
"plugin", p.Manifest.Name, "server", srv.Key,
|
||||
"tools", len(toolsResult.Tools), "commands", len(cmds))
|
||||
|
||||
return cmds
|
||||
}
|
||||
|
||||
// deriveToolCLIName converts an MCP tool name (e.g. "web_search" or
|
||||
// "maps.search_poi") into a kebab-case CLI command name ("search" or
|
||||
// "search-poi"). It strips common prefixes and replaces underscores/dots
|
||||
// with hyphens.
|
||||
func deriveToolCLIName(toolName string) string {
|
||||
// Use the last segment after "." as the base name.
|
||||
if idx := strings.LastIndex(toolName, "."); idx >= 0 {
|
||||
toolName = toolName[idx+1:]
|
||||
}
|
||||
// Replace underscores with hyphens for kebab-case.
|
||||
return strings.ReplaceAll(toolName, "_", "-")
|
||||
}
|
||||
|
||||
// buildPluginAuthClient creates a transport.Client copy with the plugin's
|
||||
// Bearer token and trusted domains injected. This allows third-party MCP
|
||||
// servers that require independent authentication to be discovered at startup.
|
||||
func buildPluginAuthClient(base *transport.Client, srv market.ServerDescriptor) *transport.Client {
|
||||
authToken := ""
|
||||
extraHeaders := make(map[string]string)
|
||||
for key, value := range srv.AuthHeaders {
|
||||
if strings.EqualFold(key, "Authorization") {
|
||||
authToken = strings.TrimPrefix(value, "Bearer ")
|
||||
authToken = strings.TrimSpace(authToken)
|
||||
} else {
|
||||
extraHeaders[key] = value
|
||||
}
|
||||
}
|
||||
if authToken == "" {
|
||||
return base
|
||||
}
|
||||
client := base.WithAuth(authToken, extraHeaders)
|
||||
// Trust the endpoint's hostname so the token is actually sent.
|
||||
if parsed, err := url.Parse(srv.Endpoint); err == nil {
|
||||
host := parsed.Hostname()
|
||||
client.TrustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
// registerPluginAuthFromHeaders extracts authentication credentials from
|
||||
// a server descriptor's AuthHeaders and registers them in the global
|
||||
// PluginAuth registry. The runner uses this registry at execution time
|
||||
// to inject the correct Bearer token for third-party MCP servers.
|
||||
func registerPluginAuthFromHeaders(srv market.ServerDescriptor) {
|
||||
authToken := ""
|
||||
extraHeaders := make(map[string]string)
|
||||
for key, value := range srv.AuthHeaders {
|
||||
if strings.EqualFold(key, "Authorization") {
|
||||
authToken = strings.TrimPrefix(value, "Bearer ")
|
||||
authToken = strings.TrimSpace(authToken)
|
||||
} else {
|
||||
extraHeaders[key] = value
|
||||
}
|
||||
}
|
||||
if authToken == "" {
|
||||
return
|
||||
}
|
||||
var trustedDomains []string
|
||||
if parsed, err := url.Parse(srv.Endpoint); err == nil {
|
||||
host := parsed.Hostname()
|
||||
trustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
productID := strings.TrimSpace(srv.CLI.ID)
|
||||
if productID == "" {
|
||||
productID = srv.Key
|
||||
}
|
||||
RegisterPluginAuth(productID, &PluginAuth{
|
||||
Token: authToken,
|
||||
ExtraHeaders: extraHeaders,
|
||||
TrustedDomains: trustedDomains,
|
||||
})
|
||||
}
|
||||
|
||||
// registerStdioServer initializes a stdio MCP server, discovers its tools
|
||||
// via ListTools, builds CLI commands, and registers the StdioClient for
|
||||
// runtime dispatch. Returns generated cobra commands.
|
||||
func registerStdioServer(p *plugin.Plugin, sc plugin.StdioServerClient, runner executor.Runner) []*cobra.Command {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
if _, err := sc.Client.Initialize(ctx); err != nil {
|
||||
slog.Warn("plugin: stdio initialize failed",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
toolsResult, err := sc.Client.ListTools(ctx)
|
||||
if err != nil {
|
||||
slog.Warn("plugin: stdio ListTools failed",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(toolsResult.Tools) == 0 {
|
||||
slog.Debug("plugin: stdio server has no tools",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build CLIOverlay: use manifest CLI metadata if present, else auto-generate.
|
||||
serverID := sc.Key
|
||||
overlay := market.CLIOverlay{
|
||||
ID: serverID,
|
||||
Command: serverID,
|
||||
}
|
||||
if srv, ok := p.Manifest.MCPServers[sc.Key]; ok && len(srv.CLI) > 0 {
|
||||
cliData := srv.CLI
|
||||
// If cli is a JSON string, treat it as a relative file path to an overlay file.
|
||||
if len(cliData) > 0 && cliData[0] == '"' {
|
||||
var cliPath string
|
||||
if err := json.Unmarshal(cliData, &cliPath); err == nil && cliPath != "" {
|
||||
absPath := filepath.Join(p.Root, cliPath)
|
||||
if fileData, readErr := os.ReadFile(absPath); readErr == nil {
|
||||
cliData = fileData
|
||||
} else {
|
||||
slog.Warn("plugin: failed to read CLI overlay file",
|
||||
"plugin", p.Manifest.Name, "path", absPath, "error", readErr)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := json.Unmarshal(cliData, &overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLI overlay for stdio server",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
}
|
||||
if overlay.ID == "" {
|
||||
overlay.ID = serverID
|
||||
}
|
||||
if overlay.Command == "" {
|
||||
overlay.Command = serverID
|
||||
}
|
||||
}
|
||||
|
||||
// Auto-generate ToolOverrides from discovered tools when not provided.
|
||||
if len(overlay.ToolOverrides) == 0 {
|
||||
overlay.ToolOverrides = make(map[string]market.CLIToolOverride)
|
||||
if len(overlay.Prefixes) == 0 {
|
||||
overlay.Prefixes = []string{serverID}
|
||||
}
|
||||
for _, tool := range toolsResult.Tools {
|
||||
overlay.ToolOverrides[tool.Name] = market.CLIToolOverride{
|
||||
IsSensitive: tool.Sensitive,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Construct virtual endpoint and server descriptor.
|
||||
endpoint := StdioEndpoint(p.Manifest.Name, sc.Key)
|
||||
|
||||
source := "plugin"
|
||||
if p.IsManaged {
|
||||
source = "plugin-managed"
|
||||
}
|
||||
|
||||
descriptor := market.ServerDescriptor{
|
||||
Key: sc.Key,
|
||||
DisplayName: p.Manifest.Name + "/" + sc.Key,
|
||||
Description: p.Manifest.Description,
|
||||
Endpoint: endpoint,
|
||||
Source: source,
|
||||
CLI: overlay,
|
||||
HasCLIMeta: true,
|
||||
}
|
||||
|
||||
AppendDynamicServer(descriptor)
|
||||
RegisterStdioClient(serverID, sc.Client)
|
||||
|
||||
// Convert tool descriptors to DetailTool entries for flag generation.
|
||||
detailsByID := make(map[string][]market.DetailTool)
|
||||
var detailTools []market.DetailTool
|
||||
for _, tool := range toolsResult.Tools {
|
||||
schemaJSON := ""
|
||||
if tool.InputSchema != nil {
|
||||
if data, marshalErr := json.Marshal(tool.InputSchema); marshalErr == nil {
|
||||
schemaJSON = string(data)
|
||||
}
|
||||
}
|
||||
detailTools = append(detailTools, market.DetailTool{
|
||||
ToolName: tool.Name,
|
||||
ToolTitle: tool.Title,
|
||||
ToolDesc: tool.Description,
|
||||
IsSensitive: tool.Sensitive,
|
||||
ToolRequest: schemaJSON,
|
||||
})
|
||||
}
|
||||
detailsByID[serverID] = detailTools
|
||||
|
||||
cmds := compat.BuildDynamicCommands(
|
||||
[]market.ServerDescriptor{descriptor}, runner, detailsByID)
|
||||
|
||||
slog.Debug("plugin: stdio server registered",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key,
|
||||
"tools", len(toolsResult.Tools), "commands", len(cmds))
|
||||
|
||||
return cmds
|
||||
}
|
||||
|
||||
// newPipelineEngine creates and configures the pipeline engine with
|
||||
// the standard set of handlers for model input correction.
|
||||
// handlers for all five pipeline phases. The phases execute in order:
|
||||
// Register → PreParse → PostParse → PreRequest → PostResponse.
|
||||
//
|
||||
// Phases are invoked at their respective integration points:
|
||||
// - Register: during command tree construction (newMCPCommand)
|
||||
// - PreParse: before Cobra parses raw argv (RunPreParse)
|
||||
// - PostParse: after Cobra parsing, before validation (canonical RunE)
|
||||
// - PreRequest: after validation, before JSON-RPC dispatch (canonical RunE)
|
||||
// - PostResponse: after transport returns, before stdout (canonical RunE)
|
||||
func newPipelineEngine() *pipeline.Engine {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
// Register handler runs during command tree building.
|
||||
handlers.RegisterHandler{},
|
||||
|
||||
// PreParse handlers run in order: alias → sticky → paramname.
|
||||
// Alias normalises case first (--userId → --user-id), then
|
||||
// sticky splits glued values (--limit100 → --limit 100), then
|
||||
@@ -839,6 +1463,12 @@ func newPipelineEngine() *pipeline.Engine {
|
||||
|
||||
// PostParse handlers normalise structured values.
|
||||
handlers.ParamValueHandler{},
|
||||
|
||||
// PreRequest handler inspects the validated payload before dispatch.
|
||||
handlers.PreRequestHandler{},
|
||||
|
||||
// PostResponse handler processes the response before output.
|
||||
handlers.PostResponseHandler{},
|
||||
)
|
||||
return engine
|
||||
}
|
||||
|
||||
@@ -29,6 +29,14 @@ import (
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
// patLikeError simulates an edition-specific PAT error that implements both
|
||||
// ExitCoder (exit code 4) and RawStderrError (raw JSON to stderr).
|
||||
type patLikeError struct{ raw string }
|
||||
|
||||
func (e *patLikeError) Error() string { return e.raw }
|
||||
func (e *patLikeError) ExitCode() int { return 4 }
|
||||
func (e *patLikeError) RawStderr() string { return e.raw }
|
||||
|
||||
func TestPrintExecutionErrorDefaultsToJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -172,8 +180,8 @@ func TestVersionCommandDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("Execute(version) error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), "\"version\"") {
|
||||
t.Fatalf("version output missing version key:\n%s", out.String())
|
||||
if !strings.Contains(out.String(), "Version:") {
|
||||
t.Fatalf("version output missing Version line:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -213,8 +221,8 @@ func TestVersionCommandUsesCachedRegistryWithoutBlockingAgedDiscovery(t *testing
|
||||
if elapsed := time.Since(start); elapsed >= 200*time.Millisecond {
|
||||
t.Fatalf("Execute(version) took %v, want cached startup under 200ms", elapsed)
|
||||
}
|
||||
if !strings.Contains(out.String(), "\"version\"") {
|
||||
t.Fatalf("version output missing version key:\n%s", out.String())
|
||||
if !strings.Contains(out.String(), "Version:") {
|
||||
t.Fatalf("version output missing Version line:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -255,6 +263,11 @@ func TestRootHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
if !strings.Contains(out.String(), "Discovered MCP Services:") {
|
||||
t.Fatalf("root help output missing MCP summary:\n%s", out.String())
|
||||
}
|
||||
for _, want := range []string{"Utility Commands:", "skill", "auth", "version"} {
|
||||
if !strings.Contains(out.String(), want) {
|
||||
t.Fatalf("root help output missing %q:\n%s", want, out.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRootShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
@@ -342,3 +355,56 @@ func TestNestedShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
t.Fatalf("nested short help output missing command title:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_writes_raw_JSON_to_stderr(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_LOW_RISK_NO_PERMISSION","data":{}}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty for RawStderrError", stdout.String())
|
||||
}
|
||||
got := strings.TrimSpace(stderr.String())
|
||||
if got != rawJSON {
|
||||
t.Fatalf("stderr = %q, want raw JSON %q", got, rawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_exit_code_is_4(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := &patLikeError{raw: `{"code":"PAT_MEDIUM_RISK_NO_PERMISSION"}`}
|
||||
exitCode := apperrors.ExitCode(err)
|
||||
if exitCode != 4 {
|
||||
t.Fatalf("apperrors.ExitCode(patLikeError) = %d, want 4", exitCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_takes_precedence_over_JSON_mode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_HIGH_RISK_NO_PERMISSION"}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
_ = root.PersistentFlags().Set("format", "json")
|
||||
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty — RawStderrError should bypass JSON mode", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stderr.String(), "PAT_HIGH_RISK_NO_PERMISSION") {
|
||||
t.Fatalf("stderr = %q, want raw PAT JSON", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -25,6 +26,7 @@ func configureRootHelp(root *cobra.Command) {
|
||||
|
||||
func renderRootHelp(root *cobra.Command) {
|
||||
services := visibleMCPRootCommands(root)
|
||||
utilities := visibleUtilityRootCommands(root)
|
||||
w := root.OutOrStdout()
|
||||
|
||||
if len(services) == 0 {
|
||||
@@ -44,8 +46,21 @@ func renderRootHelp(root *cobra.Command) {
|
||||
|
||||
_, _ = fmt.Fprintln(w, "Usage:")
|
||||
_, _ = fmt.Fprintln(w, " dws <service> [command] [flags]")
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, " dws <command> [flags]")
|
||||
}
|
||||
_, _ = fmt.Fprintln(w)
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service.`)
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, "Utility Commands:")
|
||||
_, _ = fmt.Fprintln(w)
|
||||
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
|
||||
for _, utility := range utilities {
|
||||
_, _ = fmt.Fprintf(tw, " %s\t%s\n", utility.Name(), strings.TrimSpace(utility.Short))
|
||||
}
|
||||
_ = tw.Flush()
|
||||
_, _ = fmt.Fprintln(w)
|
||||
}
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service or "dws <command> --help" for utility commands.`)
|
||||
}
|
||||
|
||||
func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
@@ -53,7 +68,16 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
return nil
|
||||
}
|
||||
|
||||
allowed := DirectRuntimeProductIDs()
|
||||
var allowed map[string]bool
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
products := fn()
|
||||
allowed = make(map[string]bool, len(products))
|
||||
for _, p := range products {
|
||||
allowed[p] = true
|
||||
}
|
||||
} else {
|
||||
allowed = DirectRuntimeProductIDs()
|
||||
}
|
||||
if len(allowed) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -70,3 +94,29 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
func visibleUtilityRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
if root == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
productCommands := DirectRuntimeProductIDs()
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
productCommands = make(map[string]bool, len(fn()))
|
||||
for _, product := range fn() {
|
||||
productCommands[product] = true
|
||||
}
|
||||
}
|
||||
|
||||
commands := make([]*cobra.Command, 0)
|
||||
for _, cmd := range root.Commands() {
|
||||
if cmd == nil || cmd.Hidden {
|
||||
continue
|
||||
}
|
||||
if productCommands[cmd.Name()] {
|
||||
continue
|
||||
}
|
||||
commands = append(commands, cmd)
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
+317
-25
@@ -15,23 +15,69 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/safety"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "启用 MCP 响应内容安全扫描",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN_ENFORCE",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "内容安全扫描发现问题时阻断响应",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_RUNTIME_CONTENT_SCAN_REPORT",
|
||||
Category: configmeta.CategoryRuntime,
|
||||
Description: "在 JSON 输出中包含安全扫描报告",
|
||||
Example: "true",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_AGENT",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-agent 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_TRACE_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-trace-id 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_SESSION_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-session-id 头",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DINGTALK_MESSAGE_ID",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "MCP 请求 x-dingtalk-message-id 头",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
runtimeContentScanEnv = "DWS_RUNTIME_CONTENT_SCAN"
|
||||
runtimeContentScanEnforceEnv = "DWS_RUNTIME_CONTENT_SCAN_ENFORCE"
|
||||
@@ -87,15 +133,25 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
return r.executeInvocation(ctx, endpoint, invocation)
|
||||
}
|
||||
|
||||
// Prefetch the Keychain token in the background. Keychain access costs
|
||||
// ~70ms on macOS; starting it here lets the load overlap with endpoint
|
||||
// resolution and catalog loading below.
|
||||
go getCachedRuntimeToken(ctx)
|
||||
|
||||
if shouldUseDirectRuntime(invocation) {
|
||||
if endpoint, ok := directRuntimeEndpoint(invocation.CanonicalProduct, invocation.Tool); ok {
|
||||
return r.executeInvocation(ctx, endpoint, invocation)
|
||||
}
|
||||
}
|
||||
|
||||
catalogStart := time.Now()
|
||||
catalog, err := r.loader.Load(ctx)
|
||||
RecordTiming(ctx, "catalog_load", time.Since(catalogStart))
|
||||
if err != nil {
|
||||
return executor.Result{}, err
|
||||
var degraded *cli.CatalogDegraded
|
||||
if !errors.As(err, °raded) {
|
||||
return executor.Result{}, err
|
||||
}
|
||||
}
|
||||
|
||||
product, ok := catalog.FindProduct(invocation.CanonicalProduct)
|
||||
@@ -116,8 +172,59 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
return r.executeInvocation(ctx, endpoint, invocation)
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
|
||||
tc := r.transport.WithAuth(r.resolveAuthToken(ctx), resolveIdentityHeaders())
|
||||
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (result executor.Result, retErr error) {
|
||||
// Route stdio:// endpoints to the local StdioClient — no HTTP, no auth.
|
||||
if IsStdioEndpoint(endpoint) {
|
||||
return r.executeStdioInvocation(ctx, invocation)
|
||||
}
|
||||
|
||||
invokeStart := time.Now()
|
||||
execID := generateExecutionID()
|
||||
r.transport.ExecutionId = execID
|
||||
|
||||
// Lazy bind FileLogger: it may be nil at construction time because
|
||||
// configureLogLevel runs later in PersistentPreRunE.
|
||||
if r.transport.FileLogger == nil {
|
||||
r.transport.FileLogger = FileLoggerInstance()
|
||||
}
|
||||
|
||||
fl := r.transport.FileLogger
|
||||
|
||||
defer func() {
|
||||
var errCat, errReason string
|
||||
if retErr != nil {
|
||||
var typed *apperrors.Error
|
||||
if errors.As(retErr, &typed) {
|
||||
errCat = string(typed.Category)
|
||||
errReason = typed.Reason
|
||||
} else {
|
||||
errCat = "unknown"
|
||||
errReason = retErr.Error()
|
||||
}
|
||||
}
|
||||
logging.LogCommandEnd(fl, execID,
|
||||
invocation.CanonicalProduct, invocation.Tool,
|
||||
retErr == nil, time.Since(invokeStart), errCat, errReason)
|
||||
}()
|
||||
|
||||
// Check if this product has plugin-level auth credentials registered.
|
||||
// If so, use the plugin's token instead of the default DingTalk OAuth token.
|
||||
// This allows third-party MCP servers (e.g. Bailian) to use their own API keys.
|
||||
pluginAuth, hasPluginAuth := LookupPluginAuth(invocation.CanonicalProduct)
|
||||
|
||||
authToken := ""
|
||||
if hasPluginAuth {
|
||||
authToken = pluginAuth.Token
|
||||
} else {
|
||||
authToken = r.resolveAuthToken(ctx)
|
||||
}
|
||||
|
||||
var timeoutSec int
|
||||
if r.globalFlags != nil {
|
||||
timeoutSec = r.globalFlags.Timeout
|
||||
}
|
||||
logging.LogCommandStart(fl, execID,
|
||||
invocation.CanonicalProduct, invocation.Tool, endpoint, version, authToken != "", timeoutSec)
|
||||
|
||||
if invocation.DryRun {
|
||||
return executor.Result{
|
||||
@@ -148,19 +255,66 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}, nil
|
||||
}
|
||||
|
||||
callResult, err := tc.CallTool(ctx, endpoint, invocation.Tool, invocation.Params)
|
||||
// Fail-fast: reject unauthenticated requests before making network calls.
|
||||
// This provides a clear error message instead of cryptic HTTP 400 from MCP.
|
||||
if strings.TrimSpace(authToken) == "" {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"未登录,请先执行 dws auth login",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("运行 'dws auth login' 完成登录后重试"),
|
||||
apperrors.WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
var tc *transport.Client
|
||||
if hasPluginAuth {
|
||||
// Use plugin-level auth: inject the plugin's token and trust its domains.
|
||||
tc = r.transport.WithAuth(authToken, pluginAuth.ExtraHeaders)
|
||||
tc.TrustedDomains = pluginAuth.TrustedDomains
|
||||
} else {
|
||||
// Default path: use DingTalk OAuth token with identity headers.
|
||||
tc = r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
|
||||
var cancel context.CancelFunc
|
||||
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
|
||||
defer cancel()
|
||||
}
|
||||
|
||||
callStart := time.Now()
|
||||
callResult, err := tc.CallTool(callCtx, endpoint, invocation.Tool, invocation.Params)
|
||||
RecordTiming(ctx, "mcp_call", time.Since(callStart))
|
||||
if err != nil {
|
||||
if isAuthError(err) {
|
||||
if fn := edition.Get().OnAuthError; fn != nil {
|
||||
if overrideErr := fn(defaultConfigDir(), err); overrideErr != nil {
|
||||
captureRuntimeFailure(invocation, err, overrideErr)
|
||||
return executor.Result{}, overrideErr
|
||||
}
|
||||
}
|
||||
}
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return executor.Result{}, err
|
||||
}
|
||||
|
||||
if fn := edition.Get().ClassifyToolResult; fn != nil {
|
||||
if editionErr := fn(callResult.Content); editionErr != nil {
|
||||
return executor.Result{}, editionErr
|
||||
}
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
|
||||
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
|
||||
mcpErr := apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("mcp_tool_error"),
|
||||
apperrors.WithServerKey(invocation.CanonicalProduct),
|
||||
apperrors.WithHint("MCP tool returned a business error; check tool parameters and refer to skill documentation."),
|
||||
apperrors.WithServerDiag(diag),
|
||||
)
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return executor.Result{}, mcpErr
|
||||
@@ -172,11 +326,14 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}
|
||||
|
||||
if bizErr := detectBusinessError(callResult.Content); bizErr != "" {
|
||||
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
|
||||
logBusinessError(r.transport.FileLogger, "business_error", invocation, callResult.Content, diag)
|
||||
return executor.Result{}, apperrors.NewAPI(bizErr,
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("business_error"),
|
||||
apperrors.WithServerKey(invocation.CanonicalProduct),
|
||||
apperrors.WithHint("The API returned a business-level error. Check required parameters and values."),
|
||||
apperrors.WithServerDiag(diag),
|
||||
)
|
||||
}
|
||||
|
||||
@@ -191,37 +348,128 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
return executor.Result{Invocation: invocation, Response: response}, nil
|
||||
}
|
||||
|
||||
// executeStdioInvocation dispatches a tool call through a local StdioClient
|
||||
// subprocess instead of the HTTP transport. This is used for plugin stdio
|
||||
// servers whose endpoints use the stdio:// scheme.
|
||||
func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
if invocation.DryRun {
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
Response: map[string]any{
|
||||
"dry_run": true,
|
||||
"transport": "stdio",
|
||||
"request": executor.ToolCallRequest(invocation.Tool, invocation.Params),
|
||||
"note": "execution skipped by --dry-run",
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
client, ok := LookupStdioClient(invocation.CanonicalProduct)
|
||||
if !ok {
|
||||
return executor.Result{}, apperrors.NewInternal(
|
||||
fmt.Sprintf("stdio client not found for %q", invocation.CanonicalProduct))
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
if r.globalFlags != nil && r.globalFlags.Timeout > 0 {
|
||||
var cancel context.CancelFunc
|
||||
callCtx, cancel = context.WithTimeout(ctx, time.Duration(r.globalFlags.Timeout)*time.Second)
|
||||
defer cancel()
|
||||
}
|
||||
|
||||
callResult, err := client.CallTool(callCtx, invocation.Tool, invocation.Params)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
fmt.Sprintf("stdio call failed: %v", err),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("stdio_error"),
|
||||
)
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("mcp_tool_error"),
|
||||
apperrors.WithServerKey(invocation.CanonicalProduct),
|
||||
)
|
||||
}
|
||||
|
||||
invocation.Implemented = true
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
Response: map[string]any{
|
||||
"transport": "stdio",
|
||||
"content": callResult.Content,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
|
||||
explicitToken := ""
|
||||
if r != nil && r.globalFlags != nil {
|
||||
explicitToken = r.globalFlags.Token
|
||||
}
|
||||
return resolveRuntimeAuthToken(ctx, explicitToken)
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
}
|
||||
if tp := edition.Get().TokenProvider; tp != nil {
|
||||
token, _ := tp(ctx, func() (string, error) {
|
||||
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
|
||||
})
|
||||
return token
|
||||
}
|
||||
return getCachedRuntimeToken(ctx)
|
||||
}
|
||||
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
}
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token)
|
||||
}
|
||||
// If the error is a decryption failure (corrupted data), surface
|
||||
// it immediately instead of falling back to empty token.
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return ""
|
||||
}
|
||||
manager := authpkg.NewManager(configDir, nil)
|
||||
configureLegacyAuthManagerCompatibility(manager)
|
||||
if token, _, err := manager.GetToken(); err == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token)
|
||||
}
|
||||
return ""
|
||||
// Use cached token to avoid repeated Keychain access (~70ms per call)
|
||||
return getCachedRuntimeToken(ctx)
|
||||
}
|
||||
|
||||
// Cached token state for process lifetime
|
||||
var (
|
||||
cachedRuntimeToken string
|
||||
cachedRuntimeTokenOnce sync.Once
|
||||
)
|
||||
|
||||
// getCachedRuntimeToken returns a cached access token, loading it only once per process.
|
||||
// This avoids repeated Keychain access which takes ~70ms each time.
|
||||
func getCachedRuntimeToken(ctx context.Context) string {
|
||||
cachedRuntimeTokenOnce.Do(func() {
|
||||
loadStart := time.Now()
|
||||
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return
|
||||
}
|
||||
if token != "" {
|
||||
cachedRuntimeToken = token
|
||||
}
|
||||
})
|
||||
return cachedRuntimeToken
|
||||
}
|
||||
|
||||
// generateExecutionID returns a random 16-char hex string used to correlate
|
||||
// all log entries (command_start, jsonrpc_request, command_end, etc.) belonging
|
||||
// to a single command invocation.
|
||||
func generateExecutionID() string {
|
||||
b := make([]byte, 8)
|
||||
_, _ = rand.Read(b)
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
|
||||
// This should be called after login/logout operations.
|
||||
func ResetRuntimeTokenCache() {
|
||||
cachedRuntimeTokenOnce = sync.Once{}
|
||||
cachedRuntimeToken = ""
|
||||
}
|
||||
|
||||
func newRuntimeContentScanner() safety.Scanner {
|
||||
@@ -255,6 +503,14 @@ func runtimeFlagEnabled(raw string, defaultValue bool) bool {
|
||||
}
|
||||
}
|
||||
|
||||
func isAuthError(err error) bool {
|
||||
var appErr *apperrors.Error
|
||||
if errors.As(err, &appErr) {
|
||||
return appErr.Category == apperrors.CategoryAuth
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func productEndpointOverride(productID string) (string, bool) {
|
||||
key := "DINGTALK_" + strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(productID), "-", "_")) + "_MCP_URL"
|
||||
value := strings.TrimSpace(os.Getenv(key))
|
||||
@@ -285,6 +541,9 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers[k] = v
|
||||
}
|
||||
}
|
||||
if fn := edition.Get().MergeHeaders; fn != nil {
|
||||
headers = fn(headers)
|
||||
}
|
||||
return headers
|
||||
}
|
||||
|
||||
@@ -328,3 +587,36 @@ func extractMCPErrorMessage(result transport.ToolCallResult) string {
|
||||
}
|
||||
return "MCP tool returned an error response"
|
||||
}
|
||||
|
||||
// logBusinessError logs MCP tool errors and business errors to the file logger
|
||||
// so they can be diagnosed offline. These errors arrive as HTTP 200 responses
|
||||
// and would otherwise not be captured by transport-level logging.
|
||||
func logBusinessError(logger *slog.Logger, reason string, inv executor.Invocation, content map[string]any, diag apperrors.ServerDiagnostics) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []any{
|
||||
"product", inv.CanonicalProduct,
|
||||
"tool", inv.Tool,
|
||||
"reason", reason,
|
||||
}
|
||||
if diag.TraceID != "" {
|
||||
attrs = append(attrs, "trace_id", diag.TraceID)
|
||||
}
|
||||
if diag.ServerErrorCode != "" {
|
||||
attrs = append(attrs, "server_error_code", diag.ServerErrorCode)
|
||||
}
|
||||
if diag.TechnicalDetail != "" {
|
||||
attrs = append(attrs, "technical_detail", diag.TechnicalDetail)
|
||||
}
|
||||
if msg, ok := content["error"].(string); ok {
|
||||
attrs = append(attrs, "error", msg)
|
||||
}
|
||||
if msg, ok := content["errorMsg"].(string); ok {
|
||||
attrs = append(attrs, "errorMsg", msg)
|
||||
}
|
||||
if msg, ok := content["message"].(string); ok {
|
||||
attrs = append(attrs, "message", msg)
|
||||
}
|
||||
logger.Warn("business_error", attrs...)
|
||||
}
|
||||
|
||||
+176
-6
@@ -17,6 +17,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -25,12 +26,61 @@ import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
func setupRuntimeCommandTest(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
discoverySrv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_ = json.NewEncoder(w).Encode(contactDiscoveryResponse())
|
||||
}))
|
||||
t.Cleanup(func() { discoverySrv.Close() })
|
||||
SetDiscoveryBaseURL(discoverySrv.URL)
|
||||
t.Cleanup(func() { SetDiscoveryBaseURL("") })
|
||||
}
|
||||
|
||||
func contactDiscoveryResponse() map[string]any {
|
||||
return map[string]any{
|
||||
"metadata": map[string]any{"count": 1, "nextCursor": ""},
|
||||
"servers": []any{
|
||||
map[string]any{
|
||||
"server": map[string]any{
|
||||
"name": "Contact",
|
||||
"description": "通讯录",
|
||||
"remotes": []any{
|
||||
map[string]any{
|
||||
"type": "streamable-http",
|
||||
"url": "https://mcp.dingtalk.com/contact/v1",
|
||||
},
|
||||
},
|
||||
},
|
||||
"_meta": map[string]any{
|
||||
"com.dingtalk.mcp.registry/metadata": map[string]any{
|
||||
"status": "active", "isLatest": true,
|
||||
},
|
||||
"com.dingtalk.mcp.registry/cli": map[string]any{
|
||||
"id": "contact",
|
||||
"command": "contact",
|
||||
"groups": map[string]any{
|
||||
"user": map[string]any{
|
||||
"description": "用户管理",
|
||||
},
|
||||
},
|
||||
"toolOverrides": map[string]any{
|
||||
"get_current_user_profile": map[string]any{
|
||||
"cliName": "get-self",
|
||||
"group": "user",
|
||||
"flags": map[string]any{},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
|
||||
@@ -45,7 +95,7 @@ func TestRuntimeRunnerIncludesContentScanReportWhenEnabled(t *testing.T) {
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
@@ -90,7 +140,7 @@ func TestRuntimeRunnerBlocksUnsafeContentWhenEnforced(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetOut(&bytes.Buffer{})
|
||||
cmd.SetErr(&bytes.Buffer{})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
@@ -121,7 +171,7 @@ func TestCanonicalCommandUsesRuntimeRunnerWhenEnabled(t *testing.T) {
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--yes"})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--yes", "--token", "test-token"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
@@ -255,6 +305,37 @@ func TestRuntimeRunnerInjectsAuthTokenFromFlag(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestRuntimeRunnerRejectsUnauthenticatedRequest verifies that requests without
|
||||
// a valid token are rejected with a clear error before making any network call.
|
||||
func TestRuntimeRunnerRejectsUnauthenticatedRequest(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
server := mockmcp.DefaultServer()
|
||||
defer server.Close()
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.RemoteURL("/server/doc"), false))
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var stdout, stderr bytes.Buffer
|
||||
cmd.SetOut(&stdout)
|
||||
cmd.SetErr(&stderr)
|
||||
// No --token flag, should be rejected
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want authentication error")
|
||||
}
|
||||
|
||||
// Verify we get a clear auth error, not a cryptic HTTP 400
|
||||
errMsg := err.Error()
|
||||
if !strings.Contains(errMsg, "未登录") {
|
||||
t.Fatalf("Execute() error = %v, want error containing '未登录'", err)
|
||||
}
|
||||
if !strings.Contains(errMsg, "auth login") {
|
||||
t.Fatalf("Execute() error = %v, want error containing 'auth login'", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerFallsBackForUnavailableProduct(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
server := mockmcp.DefaultServer()
|
||||
@@ -415,7 +496,7 @@ func TestCanonicalSensitiveToolAcceptsInteractiveConfirmation(t *testing.T) {
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
cmd.SetIn(strings.NewReader("yes\n"))
|
||||
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "create_document", "--json", `{"title":"Quarterly"}`, "--token", "test-token"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
@@ -465,7 +546,7 @@ func TestRuntimeRunnerUsesProductEndpointOverride(t *testing.T) {
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
@@ -565,6 +646,95 @@ func contentScanServer() *mockmcp.Server {
|
||||
return mockmcp.MustNewServer(fixture)
|
||||
}
|
||||
|
||||
func TestClassifyToolResultHookPreemptsBusinessError(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
t.Setenv("DWS_TRUSTED_DOMAINS", "*")
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
var req map[string]any
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, "bad request", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
method, _ := req["method"].(string)
|
||||
switch method {
|
||||
case "initialize":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"protocolVersion": "2025-03-26",
|
||||
"capabilities": map[string]any{"tools": map[string]any{"listChanged": false}},
|
||||
"serverInfo": map[string]any{"name": "doc", "version": "1.0.0"},
|
||||
},
|
||||
})
|
||||
case "notifications/initialized":
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
case "tools/list":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"tools": []map[string]any{{
|
||||
"name": "search_documents",
|
||||
"title": "Search",
|
||||
"description": "Search documents",
|
||||
"inputSchema": map[string]any{"type": "object"},
|
||||
}},
|
||||
},
|
||||
})
|
||||
case "tools/call":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": req["id"],
|
||||
"result": map[string]any{
|
||||
"content": map[string]any{
|
||||
"success": false,
|
||||
"code": "PAT_LOW_RISK_NO_PERMISSION",
|
||||
"data": map[string]any{"requiredScopes": []any{}},
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
hookCalled := false
|
||||
sentinelMsg := "hook-intercepted-PAT"
|
||||
edition.Override(&edition.Hooks{
|
||||
ClassifyToolResult: func(content map[string]any) error {
|
||||
if code, ok := content["code"].(string); ok && strings.Contains(code, "PAT") {
|
||||
hookCalled = true
|
||||
return fmt.Errorf("%s", sentinelMsg)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
})
|
||||
t.Cleanup(func() { edition.Override(&edition.Hooks{}) })
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.URL, false))
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetOut(&bytes.Buffer{})
|
||||
cmd.SetErr(&bytes.Buffer{})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want hook sentinel error")
|
||||
}
|
||||
if !hookCalled {
|
||||
t.Fatal("ClassifyToolResult hook was not called")
|
||||
}
|
||||
if !strings.Contains(err.Error(), sentinelMsg) {
|
||||
t.Fatalf("error = %q, want hook sentinel %q (not generic business error)", err.Error(), sentinelMsg)
|
||||
}
|
||||
if strings.Contains(err.Error(), "business_error") || strings.Contains(err.Error(), "mcp_tool_error") {
|
||||
t.Fatalf("error = %q, should NOT contain generic framework error category", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
t.Setenv("DWS_ALLOW_HTTP_ENDPOINTS", "1")
|
||||
@@ -628,7 +798,7 @@ func TestRuntimeRunnerReturnsErrorWhenMCPIsErrorTrue(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetOut(&bytes.Buffer{})
|
||||
cmd.SetErr(&bytes.Buffer{})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`})
|
||||
cmd.SetArgs([]string{"mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
|
||||
@@ -0,0 +1,613 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_SKILL_API_HOST",
|
||||
Category: configmeta.CategoryNetwork,
|
||||
Description: "覆盖 Skill API 地址",
|
||||
DefaultValue: "https://mcp.dingtalk.com",
|
||||
Example: "https://custom-mcp.example.com",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// legacySkillAPIHost is the legacy skill market host used by the old cli.
|
||||
legacySkillAPIHost = "https://mcp.dingtalk.com"
|
||||
// skillDownloadEndpoint is the API endpoint for downloading skills.
|
||||
skillDownloadEndpoint = "https://aihub.dingtalk.com/cli/download"
|
||||
// skillDownloadTimeout is the timeout for skill download operations.
|
||||
skillDownloadTimeout = 5 * time.Minute
|
||||
)
|
||||
|
||||
// downloadSkillResponse represents the API response for skill download.
|
||||
type downloadSkillResponse struct {
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result *downloadSkillResult `json:"result,omitempty"`
|
||||
}
|
||||
|
||||
// downloadSkillResult contains the download URL and file name.
|
||||
type downloadSkillResult struct {
|
||||
DownloadURL string `json:"downloadUrl"`
|
||||
FileName string `json:"fileName"`
|
||||
}
|
||||
|
||||
// findSkillsResponse represents the legacy skill search API response.
|
||||
type findSkillsResponse struct {
|
||||
Success bool `json:"success"`
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
Result []CliSkillDTO `json:"result,omitempty"`
|
||||
}
|
||||
|
||||
// CliSkillDTO mirrors the old cli response payload for `skill find`.
|
||||
type CliSkillDTO struct {
|
||||
SkillID string `json:"skillId"`
|
||||
Name string `json:"name"`
|
||||
Desc string `json:"desc"`
|
||||
Icon string `json:"icon"`
|
||||
}
|
||||
|
||||
// agentSkillPaths maps target names to their relative skill installation paths.
|
||||
// These paths are relative to the user's home directory.
|
||||
var agentSkillPaths = map[string]string{
|
||||
"qoder": ".qoder/skills",
|
||||
"claude": ".claude/skills",
|
||||
"cursor": ".cursor/skills",
|
||||
"codex": ".codex/skills",
|
||||
"opencode": filepath.Join(".config", "opencode", "skills"),
|
||||
}
|
||||
|
||||
// supportedTargets returns a comma-separated list of supported targets.
|
||||
func supportedTargets() string {
|
||||
targets := make([]string, 0, len(agentSkillPaths)+1)
|
||||
for target := range agentSkillPaths {
|
||||
targets = append(targets, target)
|
||||
}
|
||||
targets = append(targets, ".")
|
||||
return strings.Join(targets, ", ")
|
||||
}
|
||||
|
||||
func buildSkillCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "skill",
|
||||
Short: "技能管理",
|
||||
Long: "管理钉钉技能市场的技能。支持搜索、下载与安装到指定 Agent 目录。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(
|
||||
newSkillAddCommand(),
|
||||
newSkillGetCommand(),
|
||||
newSkillFindCommand(),
|
||||
newSkillSearchHintCommand(),
|
||||
)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillGetCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "get",
|
||||
Short: "获取技能压缩文件",
|
||||
Long: "从服务端下载技能包到本地临时目录。命令执行成功后会输出临时目录路径,供调用方使用。",
|
||||
Example: " dws skill get --skill-id <skillId>",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillGet,
|
||||
}
|
||||
cmd.Flags().String("skill-id", "", "技能 ID(必填)")
|
||||
_ = cmd.MarkFlagRequired("skill-id")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillFindCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "find",
|
||||
Short: "从钉钉技能市场搜索技能",
|
||||
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
|
||||
Example: " dws skill find --context 关键词",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillFind,
|
||||
}
|
||||
cmd.Flags().String("context", "", "搜索关键词(必填)")
|
||||
_ = cmd.MarkFlagRequired("context")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillSearchHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "search",
|
||||
Short: "兼容旧用法,提示使用 skill find",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill find --context <关键词>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newSkillAddCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "add <skillId> <target>",
|
||||
Short: "下载并安装技能到指定目录",
|
||||
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
|
||||
|
||||
参数:
|
||||
skillId 技能 ID(必填),可从钉钉技能市场获取
|
||||
target 安装目标(必填),支持: %s
|
||||
|
||||
安装路径:
|
||||
qoder -> ~/.qoder/skills/
|
||||
claude -> ~/.claude/skills/
|
||||
cursor -> ~/.cursor/skills/
|
||||
codex -> ~/.codex/skills/
|
||||
opencode -> ~/.config/opencode/skills/
|
||||
. -> 当前目录
|
||||
|
||||
示例:
|
||||
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillAdd,
|
||||
}
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/install?skillId=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(skillID)))
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "⬇️ 下载技能包...")
|
||||
|
||||
tmpDir, err := downloadSkillToTmpDir(cmd.Context(), apiURL, accessToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), tmpDir)
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
keyword, _ := cmd.Flags().GetString("context")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
|
||||
req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
var result findSkillsResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to parse search response: %v", err))
|
||||
}
|
||||
if !result.Success {
|
||||
errMsg := strings.TrimSpace(result.ErrorMsg)
|
||||
if errMsg == "" {
|
||||
errMsg = strings.TrimSpace(result.ErrorCode)
|
||||
}
|
||||
if errMsg == "" {
|
||||
errMsg = "unknown error"
|
||||
}
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %s", errMsg))
|
||||
}
|
||||
|
||||
if len(result.Result) == 0 {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "未找到匹配的技能")
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, skill := range result.Result {
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "SkillID: %s\n", skill.SkillID)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Name: %s\n", skill.Name)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Desc: %s\n", skill.Desc)
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "---")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
skillID := strings.TrimSpace(args[0])
|
||||
target := strings.TrimSpace(args[1])
|
||||
|
||||
if skillID == "" {
|
||||
return apperrors.NewValidation("skillId is required")
|
||||
}
|
||||
|
||||
// Resolve target path
|
||||
destPath, err := resolveSkillTargetPath(target)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
|
||||
}
|
||||
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(cmd.Context(), skillDownloadTimeout)
|
||||
defer cancel()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
|
||||
// Step 1: Get download URL from API
|
||||
fmt.Fprintf(w, "正在获取技能信息...\n")
|
||||
downloadResp, err := fetchSkillDownloadInfo(ctx, accessToken, skillID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if !downloadResp.Success {
|
||||
errMsg := downloadResp.ErrorMsg
|
||||
if errMsg == "" {
|
||||
errMsg = downloadResp.ErrorCode
|
||||
}
|
||||
if errMsg == "" {
|
||||
errMsg = "unknown error"
|
||||
}
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to get skill download info: %s", errMsg),
|
||||
apperrors.WithReason(downloadResp.ErrorCode))
|
||||
}
|
||||
|
||||
if downloadResp.Result == nil || downloadResp.Result.DownloadURL == "" {
|
||||
return apperrors.NewAPI("skill download URL not found in response")
|
||||
}
|
||||
|
||||
// Step 2: Download the skill zip file
|
||||
fmt.Fprintf(w, "正在下载技能...\n")
|
||||
tempZipPath, err := downloadSkillFile(ctx, downloadResp.Result.DownloadURL, downloadResp.Result.FileName)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer cleanupTempFile(tempZipPath)
|
||||
|
||||
// Step 3: Extract zip to destination
|
||||
fmt.Fprintf(w, "正在解压到 %s...\n", destPath)
|
||||
if err := extractSkillZip(tempZipPath, destPath); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "\n[OK] 技能安装成功!\n")
|
||||
fmt.Fprintf(w, "安装路径: %s\n", destPath)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSkillAccessToken() (string, error) {
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return "", skillAuthError()
|
||||
}
|
||||
return tokenData.AccessToken, nil
|
||||
}
|
||||
|
||||
func skillAuthError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return apperrors.NewAuth("认证信息已失效",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("请先完成钉钉账号登录后重试"))
|
||||
}
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
}
|
||||
|
||||
func skillAPIHost() string {
|
||||
if override := strings.TrimSpace(os.Getenv("DWS_SKILL_API_HOST")); override != "" {
|
||||
return strings.TrimRight(override, "/")
|
||||
}
|
||||
return legacySkillAPIHost
|
||||
}
|
||||
|
||||
// resolveSkillTargetPath resolves the target argument to an absolute path.
|
||||
func resolveSkillTargetPath(target string) (string, error) {
|
||||
target = strings.TrimSpace(target)
|
||||
if target == "" {
|
||||
return "", fmt.Errorf("target is required")
|
||||
}
|
||||
|
||||
// Special case: current directory
|
||||
if target == "." {
|
||||
return os.Getwd()
|
||||
}
|
||||
|
||||
// Look up predefined agent paths
|
||||
relPath, ok := agentSkillPaths[strings.ToLower(target)]
|
||||
if !ok {
|
||||
return "", fmt.Errorf("unsupported target")
|
||||
}
|
||||
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get home directory: %w", err)
|
||||
}
|
||||
|
||||
return filepath.Join(homeDir, relPath), nil
|
||||
}
|
||||
|
||||
// fetchSkillDownloadInfo calls the download API to get the skill download URL.
|
||||
func fetchSkillDownloadInfo(ctx context.Context, accessToken, skillID string) (*downloadSkillResponse, error) {
|
||||
url := fmt.Sprintf("%s?skillId=%s", skillDownloadEndpoint, skillID)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return nil, apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, apperrors.NewAPI(fmt.Sprintf("failed to call download API: %v", err),
|
||||
apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
return nil, skillAuthError()
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, apperrors.NewAPI(fmt.Sprintf("download API returned HTTP %d", resp.StatusCode),
|
||||
apperrors.WithRetryable(resp.StatusCode >= 500))
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024)) // 10MB limit
|
||||
if err != nil {
|
||||
return nil, apperrors.NewAPI(fmt.Sprintf("failed to read response: %v", err))
|
||||
}
|
||||
|
||||
var result downloadSkillResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return nil, apperrors.NewAPI(fmt.Sprintf("failed to parse response: %v", err))
|
||||
}
|
||||
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func downloadSkillToTmpDir(ctx context.Context, apiURL, accessToken string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: skillDownloadTimeout}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill package: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
tmpDir, err := os.MkdirTemp("", "dws-skill-*")
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp dir: %v", err))
|
||||
}
|
||||
|
||||
filename := filenameFromDisposition(resp.Header.Get("Content-Disposition"))
|
||||
destPath := filepath.Join(tmpDir, filename)
|
||||
file, err := os.Create(destPath)
|
||||
if err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
if _, err := io.Copy(file, resp.Body); err != nil {
|
||||
os.RemoveAll(tmpDir)
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
|
||||
}
|
||||
return tmpDir, nil
|
||||
}
|
||||
|
||||
func filenameFromDisposition(cd string) string {
|
||||
if cd != "" {
|
||||
if _, params, err := mime.ParseMediaType(cd); err == nil {
|
||||
if name := strings.TrimSpace(params["filename"]); name != "" {
|
||||
return name
|
||||
}
|
||||
}
|
||||
}
|
||||
return "skill.zip"
|
||||
}
|
||||
|
||||
func parseLegacySkillAPIError(resp *http.Response) error {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusUnauthorized:
|
||||
return skillAuthError()
|
||||
case http.StatusBadRequest:
|
||||
return apperrors.NewValidation("request parameters are invalid")
|
||||
case http.StatusNotFound:
|
||||
return apperrors.NewValidation("skill does not exist or corresponding file was not found")
|
||||
default:
|
||||
return apperrors.NewAPI(fmt.Sprintf("skill API returned HTTP %d", resp.StatusCode),
|
||||
apperrors.WithRetryable(resp.StatusCode >= 500))
|
||||
}
|
||||
}
|
||||
|
||||
// downloadSkillFile downloads the skill zip file to a temporary location.
|
||||
func downloadSkillFile(ctx context.Context, downloadURL, fileName string) (string, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, downloadURL, nil)
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create download request: %v", err))
|
||||
}
|
||||
|
||||
client := &http.Client{Timeout: skillDownloadTimeout}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to download skill: %v", err),
|
||||
apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("download returned HTTP %d", resp.StatusCode),
|
||||
apperrors.WithRetryable(resp.StatusCode >= 500))
|
||||
}
|
||||
|
||||
// Create temp file
|
||||
if fileName == "" {
|
||||
fileName = "skill.zip"
|
||||
}
|
||||
tempFile, err := os.CreateTemp("", "dws-skill-*.zip")
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to create temp file: %v", err))
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
|
||||
// Copy response body to temp file
|
||||
_, err = io.Copy(tempFile, resp.Body)
|
||||
closeErr := tempFile.Close()
|
||||
if err != nil {
|
||||
os.Remove(tempPath)
|
||||
return "", apperrors.NewAPI(fmt.Sprintf("failed to save downloaded file: %v", err))
|
||||
}
|
||||
if closeErr != nil {
|
||||
os.Remove(tempPath)
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to close temp file: %v", closeErr))
|
||||
}
|
||||
|
||||
return tempPath, nil
|
||||
}
|
||||
|
||||
// extractSkillZip extracts a zip file to the destination directory.
|
||||
func extractSkillZip(zipPath, destDir string) error {
|
||||
// Ensure destination directory exists
|
||||
if err := os.MkdirAll(destDir, 0755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create destination directory: %v", err))
|
||||
}
|
||||
|
||||
reader, err := zip.OpenReader(zipPath)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to open zip file: %v", err))
|
||||
}
|
||||
defer reader.Close()
|
||||
|
||||
for _, file := range reader.File {
|
||||
if err := extractZipFile(file, destDir); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// extractZipFile extracts a single file from the zip archive.
|
||||
func extractZipFile(file *zip.File, destDir string) error {
|
||||
// Sanitize file path to prevent zip slip attacks
|
||||
filePath := filepath.Join(destDir, file.Name)
|
||||
if !strings.HasPrefix(filepath.Clean(filePath), filepath.Clean(destDir)+string(os.PathSeparator)) {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid file path in zip: %s", file.Name))
|
||||
}
|
||||
|
||||
if file.FileInfo().IsDir() {
|
||||
// Use 0755 to ensure we have write permission for creating files inside
|
||||
return os.MkdirAll(filePath, 0755)
|
||||
}
|
||||
|
||||
// Ensure parent directory exists with write permission
|
||||
if err := os.MkdirAll(filepath.Dir(filePath), 0755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
|
||||
}
|
||||
|
||||
// Extract file
|
||||
srcFile, err := file.Open()
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to open file in zip: %v", err))
|
||||
}
|
||||
defer srcFile.Close()
|
||||
|
||||
// Use file mode from zip but ensure at least 0644 for files
|
||||
fileMode := file.Mode()
|
||||
if fileMode&0600 == 0 {
|
||||
fileMode = 0644
|
||||
}
|
||||
destFile, err := os.OpenFile(filePath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create file: %v", err))
|
||||
}
|
||||
defer destFile.Close()
|
||||
|
||||
if _, err := io.Copy(destFile, srcFile); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to extract file: %v", err))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// cleanupTempFile removes a temporary file, ignoring errors.
|
||||
func cleanupTempFile(path string) {
|
||||
if path != "" {
|
||||
os.Remove(path)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,784 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
)
|
||||
|
||||
func TestResolveSkillTargetPath(t *testing.T) {
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get home directory: %v", err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
target string
|
||||
wantSuffix string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "qoder target",
|
||||
target: "qoder",
|
||||
wantSuffix: filepath.Join(".qoder", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "claude target",
|
||||
target: "claude",
|
||||
wantSuffix: filepath.Join(".claude", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "cursor target",
|
||||
target: "cursor",
|
||||
wantSuffix: filepath.Join(".cursor", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "codex target",
|
||||
target: "codex",
|
||||
wantSuffix: filepath.Join(".codex", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "opencode target",
|
||||
target: "opencode",
|
||||
wantSuffix: filepath.Join(".config", "opencode", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "case insensitive - QODER",
|
||||
target: "QODER",
|
||||
wantSuffix: filepath.Join(".qoder", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "case insensitive - Claude",
|
||||
target: "Claude",
|
||||
wantSuffix: filepath.Join(".claude", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid target",
|
||||
target: "invalid",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "empty target",
|
||||
target: "",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "whitespace only",
|
||||
target: " ",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := resolveSkillTargetPath(tt.target)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("resolveSkillTargetPath() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if !tt.wantErr {
|
||||
expected := filepath.Join(homeDir, tt.wantSuffix)
|
||||
if got != expected {
|
||||
t.Errorf("resolveSkillTargetPath() = %v, want %v", got, expected)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSkillTargetPathCurrentDir(t *testing.T) {
|
||||
// Test "." target returns current working directory
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get working directory: %v", err)
|
||||
}
|
||||
|
||||
got, err := resolveSkillTargetPath(".")
|
||||
if err != nil {
|
||||
t.Errorf("resolveSkillTargetPath(\".\") error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if got != cwd {
|
||||
t.Errorf("resolveSkillTargetPath(\".\") = %v, want %v", got, cwd)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDownloadSkillResponse(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
jsonInput string
|
||||
wantSuccess bool
|
||||
wantURL string
|
||||
wantFile string
|
||||
wantErrCode string
|
||||
wantErrMsg string
|
||||
}{
|
||||
{
|
||||
name: "successful response",
|
||||
jsonInput: `{
|
||||
"success": true,
|
||||
"result": {
|
||||
"downloadUrl": "https://example.com/skill.zip",
|
||||
"fileName": "my-skill.zip"
|
||||
}
|
||||
}`,
|
||||
wantSuccess: true,
|
||||
wantURL: "https://example.com/skill.zip",
|
||||
wantFile: "my-skill.zip",
|
||||
},
|
||||
{
|
||||
name: "error response",
|
||||
jsonInput: `{
|
||||
"success": false,
|
||||
"errorCode": "SKILL_NOT_FOUND",
|
||||
"errorMsg": "The skill does not exist"
|
||||
}`,
|
||||
wantSuccess: false,
|
||||
wantErrCode: "SKILL_NOT_FOUND",
|
||||
wantErrMsg: "The skill does not exist",
|
||||
},
|
||||
{
|
||||
name: "success with empty result",
|
||||
jsonInput: `{
|
||||
"success": true,
|
||||
"result": null
|
||||
}`,
|
||||
wantSuccess: true,
|
||||
wantURL: "",
|
||||
wantFile: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var resp downloadSkillResponse
|
||||
if err := json.Unmarshal([]byte(tt.jsonInput), &resp); err != nil {
|
||||
t.Fatalf("failed to unmarshal JSON: %v", err)
|
||||
}
|
||||
|
||||
if resp.Success != tt.wantSuccess {
|
||||
t.Errorf("Success = %v, want %v", resp.Success, tt.wantSuccess)
|
||||
}
|
||||
|
||||
if tt.wantSuccess && resp.Result != nil {
|
||||
if resp.Result.DownloadURL != tt.wantURL {
|
||||
t.Errorf("DownloadURL = %v, want %v", resp.Result.DownloadURL, tt.wantURL)
|
||||
}
|
||||
if resp.Result.FileName != tt.wantFile {
|
||||
t.Errorf("FileName = %v, want %v", resp.Result.FileName, tt.wantFile)
|
||||
}
|
||||
}
|
||||
|
||||
if !tt.wantSuccess {
|
||||
if resp.ErrorCode != tt.wantErrCode {
|
||||
t.Errorf("ErrorCode = %v, want %v", resp.ErrorCode, tt.wantErrCode)
|
||||
}
|
||||
if resp.ErrorMsg != tt.wantErrMsg {
|
||||
t.Errorf("ErrorMsg = %v, want %v", resp.ErrorMsg, tt.wantErrMsg)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZip(t *testing.T) {
|
||||
// Create a temporary zip file with test content
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "test.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create zip file with test content
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
|
||||
// Add a file to the zip
|
||||
fileContent := []byte("test content")
|
||||
writer, err := zipWriter.Create("test-file.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write(fileContent); err != nil {
|
||||
t.Fatalf("failed to write file content: %v", err)
|
||||
}
|
||||
|
||||
// Add a subdirectory with a file
|
||||
writer, err = zipWriter.Create("subdir/nested-file.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create nested file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write([]byte("nested content")); err != nil {
|
||||
t.Fatalf("failed to write nested file content: %v", err)
|
||||
}
|
||||
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract the zip
|
||||
if err := extractSkillZip(zipPath, destDir); err != nil {
|
||||
t.Fatalf("extractSkillZip() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify extracted files
|
||||
extractedFile := filepath.Join(destDir, "test-file.txt")
|
||||
content, err := os.ReadFile(extractedFile)
|
||||
if err != nil {
|
||||
t.Errorf("failed to read extracted file: %v", err)
|
||||
}
|
||||
if string(content) != "test content" {
|
||||
t.Errorf("extracted content = %v, want %v", string(content), "test content")
|
||||
}
|
||||
|
||||
// Verify nested file
|
||||
nestedFile := filepath.Join(destDir, "subdir", "nested-file.txt")
|
||||
content, err = os.ReadFile(nestedFile)
|
||||
if err != nil {
|
||||
t.Errorf("failed to read nested file: %v", err)
|
||||
}
|
||||
if string(content) != "nested content" {
|
||||
t.Errorf("nested content = %v, want %v", string(content), "nested content")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZipPreventZipSlip(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "malicious.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create a zip file with a path traversal attempt
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
|
||||
// Try to create a file with path traversal
|
||||
writer, err := zipWriter.Create("../../../etc/passwd")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create malicious file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write([]byte("malicious content")); err != nil {
|
||||
t.Fatalf("failed to write malicious content: %v", err)
|
||||
}
|
||||
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract should fail due to zip slip protection
|
||||
err = extractSkillZip(zipPath, destDir)
|
||||
if err == nil {
|
||||
t.Error("extractSkillZip() should have failed for zip slip attack")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "invalid file path") {
|
||||
t.Errorf("error should mention invalid file path, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
wantErr bool
|
||||
errMsg string
|
||||
}{
|
||||
{
|
||||
name: "missing arguments",
|
||||
args: []string{"skill", "add"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "missing target",
|
||||
args: []string{"skill", "add", "skill-123"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "too many arguments",
|
||||
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs(tt.args)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Execute() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if tt.wantErr && !strings.Contains(err.Error(), tt.errMsg) {
|
||||
t.Errorf("error = %v, should contain %v", err, tt.errMsg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
// Setup: Create config directory with valid token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
// Save a valid token
|
||||
err := authpkg.SaveTokenData(configDir, &authpkg.TokenData{
|
||||
AccessToken: "test-token",
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
})
|
||||
if err != nil {
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err = cmd.Execute()
|
||||
if err == nil {
|
||||
t.Error("Execute() should have failed for invalid target")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "invalid target") {
|
||||
t.Errorf("error should mention invalid target, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
// Setup: Create config directory without token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
// Ensure the config directory exists but has no token
|
||||
if err := os.MkdirAll(configDir, 0755); err != nil {
|
||||
t.Fatalf("failed to create config dir: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Error("Execute() should have failed without auth")
|
||||
}
|
||||
// Check for authentication-related error (English or Chinese)
|
||||
errStr := err.Error()
|
||||
if !strings.Contains(errStr, "not logged in") && !strings.Contains(errStr, "token") && !strings.Contains(errStr, "未登录") && !strings.Contains(errStr, "auth") {
|
||||
t.Errorf("error should mention authentication, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchSkillDownloadInfoUnauthorized(t *testing.T) {
|
||||
// Create mock server that returns 401
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// We can't easily test the actual fetchSkillDownloadInfo function
|
||||
// because it uses a hardcoded URL. This test verifies HTTP 401 handling pattern.
|
||||
client := &http.Client{Timeout: 5 * time.Second}
|
||||
resp, err := client.Get(server.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Errorf("expected 401, got %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSupportedTargets(t *testing.T) {
|
||||
targets := supportedTargets()
|
||||
|
||||
// Should contain all predefined targets
|
||||
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode", "."}
|
||||
for _, expected := range expectedTargets {
|
||||
if !strings.Contains(targets, expected) {
|
||||
t.Errorf("supportedTargets() should contain %s, got: %s", expected, targets)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentSkillPathsCrossPlatform(t *testing.T) {
|
||||
// Verify that paths use platform-appropriate separators
|
||||
for target, path := range agentSkillPaths {
|
||||
if runtime.GOOS == "windows" {
|
||||
if strings.Contains(path, "/") && !strings.Contains(path, "\\") {
|
||||
// On Windows, filepath.Join should use backslashes
|
||||
// But raw map values may use forward slashes
|
||||
t.Logf("Note: %s path '%s' uses forward slashes (will be converted by filepath.Join)", target, path)
|
||||
}
|
||||
}
|
||||
|
||||
// Test that resolveSkillTargetPath produces valid paths
|
||||
resolved, err := resolveSkillTargetPath(target)
|
||||
if err != nil {
|
||||
t.Errorf("resolveSkillTargetPath(%s) failed: %v", target, err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Path should be absolute
|
||||
if !filepath.IsAbs(resolved) {
|
||||
t.Errorf("resolveSkillTargetPath(%s) returned non-absolute path: %s", target, resolved)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupTempFile(t *testing.T) {
|
||||
// Create a temporary file
|
||||
tempFile, err := os.CreateTemp("", "test-cleanup-*.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create temp file: %v", err)
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
tempFile.Close()
|
||||
|
||||
// Verify file exists
|
||||
if _, err := os.Stat(tempPath); os.IsNotExist(err) {
|
||||
t.Fatalf("temp file should exist before cleanup")
|
||||
}
|
||||
|
||||
// Clean up
|
||||
cleanupTempFile(tempPath)
|
||||
|
||||
// Verify file is deleted
|
||||
if _, err := os.Stat(tempPath); !os.IsNotExist(err) {
|
||||
t.Errorf("temp file should be deleted after cleanup")
|
||||
}
|
||||
|
||||
// Cleanup should not panic on empty path
|
||||
cleanupTempFile("")
|
||||
|
||||
// Cleanup should not panic on non-existent file
|
||||
cleanupTempFile("/nonexistent/path/file.txt")
|
||||
}
|
||||
|
||||
func TestDownloadSkillResponseJSON(t *testing.T) {
|
||||
// Test JSON marshaling/unmarshaling round-trip
|
||||
original := downloadSkillResponse{
|
||||
Success: true,
|
||||
Result: &downloadSkillResult{
|
||||
DownloadURL: "https://example.com/skill.zip",
|
||||
FileName: "skill.zip",
|
||||
},
|
||||
}
|
||||
|
||||
data, err := json.Marshal(original)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var parsed downloadSkillResponse
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatalf("failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if parsed.Success != original.Success {
|
||||
t.Errorf("Success mismatch: got %v, want %v", parsed.Success, original.Success)
|
||||
}
|
||||
if parsed.Result.DownloadURL != original.Result.DownloadURL {
|
||||
t.Errorf("DownloadURL mismatch: got %v, want %v", parsed.Result.DownloadURL, original.Result.DownloadURL)
|
||||
}
|
||||
if parsed.Result.FileName != original.Result.FileName {
|
||||
t.Errorf("FileName mismatch: got %v, want %v", parsed.Result.FileName, original.Result.FileName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
output := out.String()
|
||||
// Check for the Long description which is shown in help
|
||||
if !strings.Contains(output, "技能") {
|
||||
t.Errorf("help should mention '技能', got: %s", output)
|
||||
}
|
||||
for _, subcmd := range []string{"add", "find", "get"} {
|
||||
if !strings.Contains(output, subcmd) {
|
||||
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
output := out.String()
|
||||
// Should mention supported targets
|
||||
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode"}
|
||||
for _, target := range expectedTargets {
|
||||
if !strings.Contains(output, target) {
|
||||
t.Errorf("help should mention target '%s', got: %s", target, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillGetCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "get"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillFindCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "find"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillSearchHintCommand(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "search"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), "dws skill find --context") {
|
||||
t.Fatalf("output = %q, want legacy hint", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSkillFileSuccess(t *testing.T) {
|
||||
// Create a mock server that returns a zip file
|
||||
expectedContent := []byte("fake zip content")
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/zip")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
w.Write(expectedContent)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// Download the file
|
||||
ctx := context.Background()
|
||||
tempPath, err := downloadSkillFile(ctx, server.URL, "test.zip")
|
||||
if err != nil {
|
||||
t.Fatalf("downloadSkillFile() error = %v", err)
|
||||
}
|
||||
defer os.Remove(tempPath)
|
||||
|
||||
// Verify the downloaded content
|
||||
content, err := os.ReadFile(tempPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to read downloaded file: %v", err)
|
||||
}
|
||||
|
||||
if !bytes.Equal(content, expectedContent) {
|
||||
t.Errorf("downloaded content mismatch: got %v, want %v", content, expectedContent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSkillFileServerError(t *testing.T) {
|
||||
// Create a mock server that returns 500
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
ctx := context.Background()
|
||||
_, err := downloadSkillFile(ctx, server.URL, "test.zip")
|
||||
if err == nil {
|
||||
t.Error("downloadSkillFile() should fail on server error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZipEmptyZip(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "empty.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create an empty zip file
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract should succeed even for empty zip
|
||||
if err := extractSkillZip(zipPath, destDir); err != nil {
|
||||
t.Errorf("extractSkillZip() should not fail for empty zip: %v", err)
|
||||
}
|
||||
|
||||
// Destination directory should be created
|
||||
if _, err := os.Stat(destDir); os.IsNotExist(err) {
|
||||
t.Errorf("destination directory should be created")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZipWithDirectories(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "test.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create zip with directory entries
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
|
||||
// Add a directory entry with proper permissions
|
||||
header := &zip.FileHeader{
|
||||
Name: "mydir/",
|
||||
Method: zip.Deflate,
|
||||
}
|
||||
header.SetMode(0755 | os.ModeDir)
|
||||
_, err = zipWriter.CreateHeader(header)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create directory in zip: %v", err)
|
||||
}
|
||||
|
||||
// Add a file in the directory
|
||||
fileHeader := &zip.FileHeader{
|
||||
Name: "mydir/file.txt",
|
||||
Method: zip.Deflate,
|
||||
}
|
||||
fileHeader.SetMode(0644)
|
||||
writer, err := zipWriter.CreateHeader(fileHeader)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write([]byte("content")); err != nil {
|
||||
t.Fatalf("failed to write content: %v", err)
|
||||
}
|
||||
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract
|
||||
if err := extractSkillZip(zipPath, destDir); err != nil {
|
||||
t.Fatalf("extractSkillZip() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify directory was created
|
||||
dirPath := filepath.Join(destDir, "mydir")
|
||||
info, err := os.Stat(dirPath)
|
||||
if err != nil {
|
||||
t.Errorf("directory should exist: %v", err)
|
||||
} else if !info.IsDir() {
|
||||
t.Errorf("mydir should be a directory")
|
||||
}
|
||||
|
||||
// Verify file exists
|
||||
filePath := filepath.Join(destDir, "mydir", "file.txt")
|
||||
content, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
t.Errorf("file should exist: %v", err)
|
||||
} else if string(content) != "content" {
|
||||
t.Errorf("file content mismatch: got %s, want 'content'", string(content))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
const stdioEndpointScheme = "stdio://"
|
||||
|
||||
var (
|
||||
stdioMu sync.RWMutex
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
)
|
||||
|
||||
// RegisterStdioClient stores a StdioClient keyed by its canonical product ID
|
||||
// (the CLI.ID used in the server descriptor). The runner looks up this client
|
||||
// when a stdio:// endpoint is resolved at execution time.
|
||||
func RegisterStdioClient(productID string, client *transport.StdioClient) {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
stdioClients[productID] = client
|
||||
}
|
||||
|
||||
// LookupStdioClient returns the StdioClient registered for the given product ID.
|
||||
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
|
||||
stdioMu.RLock()
|
||||
defer stdioMu.RUnlock()
|
||||
c, ok := stdioClients[productID]
|
||||
return c, ok
|
||||
}
|
||||
|
||||
// StdioEndpoint returns a virtual endpoint URL for a stdio-based MCP server.
|
||||
// Format: stdio://{pluginName}/{serverKey}
|
||||
func StdioEndpoint(pluginName, serverKey string) string {
|
||||
return stdioEndpointScheme + pluginName + "/" + serverKey
|
||||
}
|
||||
|
||||
// IsStdioEndpoint returns true if the endpoint uses the stdio:// scheme.
|
||||
func IsStdioEndpoint(endpoint string) bool {
|
||||
return strings.HasPrefix(endpoint, stdioEndpointScheme)
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,381 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PERF_DEBUG",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "启用性能计时输出到 stderr",
|
||||
Example: "1",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PERF_REPORT",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "JSON 性能报告输出路径 (auto=~/.dws/perf/latest.json)",
|
||||
Example: "auto",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// PerfDebugEnv is the environment variable to enable performance timing output.
|
||||
PerfDebugEnv = "DWS_PERF_DEBUG"
|
||||
|
||||
// PerfReportEnv is the environment variable to enable JSON perf report output.
|
||||
// Set to "auto" to write to ~/.dws/perf/latest.json, or a custom file path.
|
||||
PerfReportEnv = "DWS_PERF_REPORT"
|
||||
|
||||
perfReportDir = "perf"
|
||||
perfReportFile = "latest.json"
|
||||
)
|
||||
|
||||
// timingContextKey is the context key for TimingCollector.
|
||||
type timingContextKey struct{}
|
||||
|
||||
// TimingEntry represents a single timing measurement.
|
||||
type TimingEntry struct {
|
||||
Name string
|
||||
Duration time.Duration
|
||||
Timestamp time.Time
|
||||
Seq int // insertion order
|
||||
}
|
||||
|
||||
// TimingCollector collects timing measurements for a single command execution.
|
||||
// It is safe for concurrent use.
|
||||
type TimingCollector struct {
|
||||
mu sync.Mutex
|
||||
start time.Time
|
||||
entries []TimingEntry
|
||||
seq int
|
||||
}
|
||||
|
||||
// NewTimingCollector creates a new collector with the start time set to now.
|
||||
func NewTimingCollector() *TimingCollector {
|
||||
return &TimingCollector{
|
||||
start: time.Now(),
|
||||
entries: make([]TimingEntry, 0, 16),
|
||||
}
|
||||
}
|
||||
|
||||
// Record adds a timing entry with the given name and duration.
|
||||
func (tc *TimingCollector) Record(name string, d time.Duration) {
|
||||
if tc == nil {
|
||||
return
|
||||
}
|
||||
tc.mu.Lock()
|
||||
defer tc.mu.Unlock()
|
||||
tc.entries = append(tc.entries, TimingEntry{
|
||||
Name: name,
|
||||
Duration: d,
|
||||
Timestamp: time.Now(),
|
||||
Seq: tc.seq,
|
||||
})
|
||||
tc.seq++
|
||||
}
|
||||
|
||||
// StartTimer returns a function that, when called, records the elapsed time
|
||||
// since StartTimer was called. This is convenient for defer usage:
|
||||
//
|
||||
// defer tc.StartTimer("operation")()
|
||||
func (tc *TimingCollector) StartTimer(name string) func() {
|
||||
if tc == nil {
|
||||
return func() {}
|
||||
}
|
||||
start := time.Now()
|
||||
return func() {
|
||||
tc.Record(name, time.Since(start))
|
||||
}
|
||||
}
|
||||
|
||||
// Total returns the total elapsed time since the collector was created.
|
||||
func (tc *TimingCollector) Total() time.Duration {
|
||||
if tc == nil {
|
||||
return 0
|
||||
}
|
||||
return time.Since(tc.start)
|
||||
}
|
||||
|
||||
// Entries returns a copy of all recorded entries in insertion order.
|
||||
func (tc *TimingCollector) Entries() []TimingEntry {
|
||||
if tc == nil {
|
||||
return nil
|
||||
}
|
||||
tc.mu.Lock()
|
||||
defer tc.mu.Unlock()
|
||||
result := make([]TimingEntry, len(tc.entries))
|
||||
copy(result, tc.entries)
|
||||
sort.Slice(result, func(i, j int) bool {
|
||||
return result[i].Seq < result[j].Seq
|
||||
})
|
||||
return result
|
||||
}
|
||||
|
||||
// formatDuration returns a human-friendly duration string.
|
||||
// Sub-µs → "0µs", sub-ms → microsecond precision (e.g. "142µs"), else → ms.
|
||||
func formatDuration(d time.Duration) string {
|
||||
switch {
|
||||
case d < time.Microsecond:
|
||||
return "0µs"
|
||||
case d < time.Millisecond:
|
||||
return d.Truncate(time.Microsecond).String()
|
||||
default:
|
||||
return d.Truncate(time.Millisecond).String()
|
||||
}
|
||||
}
|
||||
|
||||
// Print writes a summary of all timing entries to the given writer.
|
||||
func (tc *TimingCollector) Print(w io.Writer) {
|
||||
if tc == nil || w == nil {
|
||||
return
|
||||
}
|
||||
entries := tc.Entries()
|
||||
total := tc.Total()
|
||||
if len(entries) == 0 {
|
||||
fmt.Fprintf(w, "\n[Perf] Total: %v (no detailed entries)\n", formatDuration(total))
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintln(w, "[Perf] Execution breakdown:")
|
||||
for _, e := range entries {
|
||||
fmt.Fprintf(w, " %-30s %v\n", e.Name, formatDuration(e.Duration))
|
||||
}
|
||||
fmt.Fprintf(w, " %-30s %v\n", "──────────────────────────────", "──────────")
|
||||
fmt.Fprintf(w, " %-30s %v\n", "Total", formatDuration(total))
|
||||
}
|
||||
|
||||
// PrintIfEnabled prints timing info to stderr if DWS_PERF_DEBUG is set.
|
||||
func (tc *TimingCollector) PrintIfEnabled() {
|
||||
if tc == nil {
|
||||
return
|
||||
}
|
||||
if os.Getenv(PerfDebugEnv) == "" {
|
||||
return
|
||||
}
|
||||
tc.Print(os.Stderr)
|
||||
}
|
||||
|
||||
// WithTimingCollector returns a new context with the TimingCollector attached.
|
||||
func WithTimingCollector(ctx context.Context, tc *TimingCollector) context.Context {
|
||||
return context.WithValue(ctx, timingContextKey{}, tc)
|
||||
}
|
||||
|
||||
// TimingCollectorFromContext extracts the TimingCollector from context, or nil.
|
||||
func TimingCollectorFromContext(ctx context.Context) *TimingCollector {
|
||||
if ctx == nil {
|
||||
return nil
|
||||
}
|
||||
tc, _ := ctx.Value(timingContextKey{}).(*TimingCollector)
|
||||
return tc
|
||||
}
|
||||
|
||||
// RecordTiming is a convenience function to record timing to the collector in context.
|
||||
func RecordTiming(ctx context.Context, name string, d time.Duration) {
|
||||
if tc := TimingCollectorFromContext(ctx); tc != nil {
|
||||
tc.Record(name, d)
|
||||
}
|
||||
}
|
||||
|
||||
// StartTiming is a convenience function that returns a stop function for defer usage.
|
||||
// Example:
|
||||
//
|
||||
// defer StartTiming(ctx, "operation")()
|
||||
func StartTiming(ctx context.Context, name string) func() {
|
||||
tc := TimingCollectorFromContext(ctx)
|
||||
if tc == nil {
|
||||
return func() {}
|
||||
}
|
||||
return tc.StartTimer(name)
|
||||
}
|
||||
|
||||
// IsPerfDebugEnabled returns true if performance debug output is enabled.
|
||||
func IsPerfDebugEnabled() bool {
|
||||
return os.Getenv(PerfDebugEnv) != ""
|
||||
}
|
||||
|
||||
// ── Structured Performance Report ──────────────────────────────────────
|
||||
|
||||
// PerfPhase is a single phase in the performance report.
|
||||
type PerfPhase struct {
|
||||
Name string `json:"name"`
|
||||
DurationMs int64 `json:"duration_ms"`
|
||||
Seq int `json:"seq"`
|
||||
}
|
||||
|
||||
// PerfReport is the JSON-serialisable performance report.
|
||||
type PerfReport struct {
|
||||
Kind string `json:"kind"`
|
||||
Version string `json:"version"`
|
||||
CLIVersion string `json:"cli_version"`
|
||||
Command string `json:"command"`
|
||||
Timestamp time.Time `json:"timestamp"`
|
||||
TotalMs int64 `json:"total_ms"`
|
||||
Phases []PerfPhase `json:"phases"`
|
||||
Slowest string `json:"slowest"`
|
||||
OverheadMs int64 `json:"overhead_ms"`
|
||||
}
|
||||
|
||||
// BuildReport constructs a PerfReport from the collected timing entries.
|
||||
func (tc *TimingCollector) BuildReport(cliVersion, command string) PerfReport {
|
||||
entries := tc.Entries()
|
||||
total := tc.Total()
|
||||
totalMs := total.Milliseconds()
|
||||
|
||||
phases := make([]PerfPhase, len(entries))
|
||||
var sumMs int64
|
||||
var slowestName string
|
||||
var slowestMs int64
|
||||
|
||||
for i, e := range entries {
|
||||
ms := e.Duration.Milliseconds()
|
||||
phases[i] = PerfPhase{
|
||||
Name: e.Name,
|
||||
DurationMs: ms,
|
||||
Seq: e.Seq,
|
||||
}
|
||||
sumMs += ms
|
||||
if ms > slowestMs {
|
||||
slowestMs = ms
|
||||
slowestName = e.Name
|
||||
}
|
||||
}
|
||||
|
||||
overhead := totalMs - sumMs
|
||||
if overhead < 0 {
|
||||
overhead = 0
|
||||
}
|
||||
|
||||
return PerfReport{
|
||||
Kind: "perf_report",
|
||||
Version: "1",
|
||||
CLIVersion: cliVersion,
|
||||
Command: command,
|
||||
Timestamp: time.Now(),
|
||||
TotalMs: totalMs,
|
||||
Phases: phases,
|
||||
Slowest: slowestName,
|
||||
OverheadMs: overhead,
|
||||
}
|
||||
}
|
||||
|
||||
// WriteReportIfEnabled checks DWS_PERF_REPORT and writes a JSON report if set.
|
||||
func (tc *TimingCollector) WriteReportIfEnabled(cliVersion, command string) {
|
||||
if tc == nil {
|
||||
return
|
||||
}
|
||||
dest := os.Getenv(PerfReportEnv)
|
||||
if dest == "" {
|
||||
return
|
||||
}
|
||||
|
||||
report := tc.BuildReport(cliVersion, command)
|
||||
data, err := json.MarshalIndent(report, "", " ")
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
path := resolvePerfReportPath(dest)
|
||||
if path == "" {
|
||||
return
|
||||
}
|
||||
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
tmp := path + ".tmp"
|
||||
if err := os.WriteFile(tmp, data, 0o600); err != nil {
|
||||
_ = os.Remove(tmp)
|
||||
return
|
||||
}
|
||||
_ = os.Rename(tmp, path)
|
||||
}
|
||||
|
||||
// LoadLatestReport reads the default perf report file (~/.dws/perf/latest.json).
|
||||
func LoadLatestReport() (*PerfReport, error) {
|
||||
path := defaultPerfReportPath()
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var report PerfReport
|
||||
if err := json.Unmarshal(data, &report); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &report, nil
|
||||
}
|
||||
|
||||
// resolvePerfReportPath resolves the DWS_PERF_REPORT value to an absolute path.
|
||||
func resolvePerfReportPath(dest string) string {
|
||||
if dest == "auto" {
|
||||
return defaultPerfReportPath()
|
||||
}
|
||||
return dest
|
||||
}
|
||||
|
||||
func defaultPerfReportPath() string {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return filepath.Join(home, ".dws", perfReportDir, perfReportFile)
|
||||
}
|
||||
|
||||
// sensitiveFlags are flag names whose values should be masked in commands.
|
||||
var sensitiveFlags = map[string]bool{
|
||||
"--token": true,
|
||||
"--client-secret": true,
|
||||
"--client-id": true,
|
||||
}
|
||||
|
||||
// SanitizeCommand redacts sensitive flag values from a command arg slice.
|
||||
func SanitizeCommand(args []string) string {
|
||||
sanitized := make([]string, 0, len(args))
|
||||
skipNext := false
|
||||
for _, arg := range args {
|
||||
if skipNext {
|
||||
sanitized = append(sanitized, "***")
|
||||
skipNext = false
|
||||
continue
|
||||
}
|
||||
if idx := strings.IndexByte(arg, '='); idx > 0 {
|
||||
key := arg[:idx]
|
||||
if sensitiveFlags[key] {
|
||||
sanitized = append(sanitized, key+"=***")
|
||||
continue
|
||||
}
|
||||
}
|
||||
if sensitiveFlags[arg] {
|
||||
skipNext = true
|
||||
}
|
||||
sanitized = append(sanitized, arg)
|
||||
}
|
||||
return strings.Join(sanitized, " ")
|
||||
}
|
||||
@@ -0,0 +1,453 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestTimingCollector_Basic(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
if tc == nil {
|
||||
t.Fatal("NewTimingCollector returned nil")
|
||||
}
|
||||
|
||||
// Record some timings
|
||||
tc.Record("op1", 10*time.Millisecond)
|
||||
tc.Record("op2", 20*time.Millisecond)
|
||||
|
||||
entries := tc.Entries()
|
||||
if len(entries) != 2 {
|
||||
t.Errorf("expected 2 entries, got %d", len(entries))
|
||||
}
|
||||
|
||||
// Check ordering
|
||||
if entries[0].Name != "op1" {
|
||||
t.Errorf("expected first entry to be 'op1', got %q", entries[0].Name)
|
||||
}
|
||||
if entries[1].Name != "op2" {
|
||||
t.Errorf("expected second entry to be 'op2', got %q", entries[1].Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTimingCollector_StartTimer(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
|
||||
stop := tc.StartTimer("timed_op")
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
stop()
|
||||
|
||||
entries := tc.Entries()
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("expected 1 entry, got %d", len(entries))
|
||||
}
|
||||
if entries[0].Name != "timed_op" {
|
||||
t.Errorf("expected entry name 'timed_op', got %q", entries[0].Name)
|
||||
}
|
||||
if entries[0].Duration < 5*time.Millisecond {
|
||||
t.Errorf("expected duration >= 5ms, got %v", entries[0].Duration)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTimingCollector_NilSafe(t *testing.T) {
|
||||
var tc *TimingCollector
|
||||
|
||||
// Should not panic on nil collector
|
||||
tc.Record("op", 10*time.Millisecond)
|
||||
stop := tc.StartTimer("op")
|
||||
stop()
|
||||
_ = tc.Total()
|
||||
_ = tc.Entries()
|
||||
tc.Print(nil)
|
||||
tc.PrintIfEnabled()
|
||||
}
|
||||
|
||||
func TestTimingCollector_Print(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("auth_token", 44*time.Millisecond)
|
||||
tc.Record("mcp_call", 150*time.Millisecond)
|
||||
|
||||
var buf bytes.Buffer
|
||||
tc.Print(&buf)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "[Perf]") {
|
||||
t.Error("output should contain [Perf] header")
|
||||
}
|
||||
if !strings.Contains(output, "auth_token") {
|
||||
t.Error("output should contain 'auth_token'")
|
||||
}
|
||||
if !strings.Contains(output, "mcp_call") {
|
||||
t.Error("output should contain 'mcp_call'")
|
||||
}
|
||||
if !strings.Contains(output, "Total") {
|
||||
t.Error("output should contain 'Total'")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTimingCollector_PrintIfEnabled(t *testing.T) {
|
||||
// Set environment variable
|
||||
os.Setenv(PerfDebugEnv, "1")
|
||||
defer os.Unsetenv(PerfDebugEnv)
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("test_op", 10*time.Millisecond)
|
||||
|
||||
// This should not panic and should print to stderr
|
||||
tc.PrintIfEnabled()
|
||||
}
|
||||
|
||||
func TestTimingCollector_ContextIntegration(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
ctx := WithTimingCollector(context.Background(), tc)
|
||||
|
||||
// Retrieve from context
|
||||
retrieved := TimingCollectorFromContext(ctx)
|
||||
if retrieved != tc {
|
||||
t.Error("TimingCollectorFromContext should return the same collector")
|
||||
}
|
||||
|
||||
// Use convenience functions
|
||||
RecordTiming(ctx, "ctx_op", 30*time.Millisecond)
|
||||
stop := StartTiming(ctx, "ctx_timed")
|
||||
time.Sleep(2 * time.Millisecond)
|
||||
stop()
|
||||
|
||||
entries := tc.Entries()
|
||||
if len(entries) != 2 {
|
||||
t.Errorf("expected 2 entries, got %d", len(entries))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTimingCollectorFromContext_NilContext(t *testing.T) {
|
||||
//lint:ignore SA1012 Testing explicit nil-context guard in TimingCollectorFromContext.
|
||||
tc := TimingCollectorFromContext(nil)
|
||||
if tc != nil {
|
||||
t.Error("TimingCollectorFromContext(nil) should return nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTimingCollectorFromContext_NoCollector(t *testing.T) {
|
||||
tc := TimingCollectorFromContext(context.Background())
|
||||
if tc != nil {
|
||||
t.Error("TimingCollectorFromContext with no collector should return nil")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStartTiming_NoCollector(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
stop := StartTiming(ctx, "no_collector")
|
||||
// Should not panic
|
||||
stop()
|
||||
}
|
||||
|
||||
func TestIsPerfDebugEnabled(t *testing.T) {
|
||||
// Clear the env var first
|
||||
os.Unsetenv(PerfDebugEnv)
|
||||
|
||||
if IsPerfDebugEnabled() {
|
||||
t.Error("IsPerfDebugEnabled should return false when env var is not set")
|
||||
}
|
||||
|
||||
os.Setenv(PerfDebugEnv, "1")
|
||||
defer os.Unsetenv(PerfDebugEnv)
|
||||
|
||||
if !IsPerfDebugEnabled() {
|
||||
t.Error("IsPerfDebugEnabled should return true when env var is set")
|
||||
}
|
||||
}
|
||||
|
||||
// ── PerfReport tests ────────────────────────────────────────────────────
|
||||
|
||||
func TestBuildReport(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 45*time.Millisecond)
|
||||
tc.Record("auth_keychain", 72*time.Millisecond)
|
||||
tc.Record("mcp_call", 620*time.Millisecond)
|
||||
|
||||
report := tc.BuildReport("v1.0.8", "dws aitable list-records")
|
||||
|
||||
if report.Kind != "perf_report" {
|
||||
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
|
||||
}
|
||||
if report.Version != "1" {
|
||||
t.Errorf("expected version '1', got %q", report.Version)
|
||||
}
|
||||
if report.CLIVersion != "v1.0.8" {
|
||||
t.Errorf("expected cli_version 'v1.0.8', got %q", report.CLIVersion)
|
||||
}
|
||||
if report.Command != "dws aitable list-records" {
|
||||
t.Errorf("expected command 'dws aitable list-records', got %q", report.Command)
|
||||
}
|
||||
if len(report.Phases) != 3 {
|
||||
t.Fatalf("expected 3 phases, got %d", len(report.Phases))
|
||||
}
|
||||
if report.Phases[0].Name != "cmd_init" || report.Phases[0].DurationMs != 45 {
|
||||
t.Errorf("unexpected first phase: %+v", report.Phases[0])
|
||||
}
|
||||
if report.Slowest != "mcp_call" {
|
||||
t.Errorf("expected slowest 'mcp_call', got %q", report.Slowest)
|
||||
}
|
||||
if report.TotalMs < 0 {
|
||||
t.Errorf("total_ms should be >= 0, got %d", report.TotalMs)
|
||||
}
|
||||
if report.OverheadMs < 0 {
|
||||
t.Errorf("overhead_ms should be >= 0, got %d", report.OverheadMs)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReportEmpty(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
report := tc.BuildReport("dev", "dws version")
|
||||
|
||||
if len(report.Phases) != 0 {
|
||||
t.Errorf("expected 0 phases, got %d", len(report.Phases))
|
||||
}
|
||||
if report.Slowest != "" {
|
||||
t.Errorf("expected empty slowest, got %q", report.Slowest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildReportJSON(t *testing.T) {
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 10*time.Millisecond)
|
||||
|
||||
report := tc.BuildReport("v1.0.0", "dws version")
|
||||
data, err := json.Marshal(report)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal failed: %v", err)
|
||||
}
|
||||
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatalf("json.Unmarshal failed: %v", err)
|
||||
}
|
||||
|
||||
requiredKeys := []string{"kind", "version", "cli_version", "command", "timestamp", "total_ms", "phases", "slowest", "overhead_ms"}
|
||||
for _, key := range requiredKeys {
|
||||
if _, ok := parsed[key]; !ok {
|
||||
t.Errorf("missing key %q in JSON output", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
reportPath := filepath.Join(dir, "report.json")
|
||||
|
||||
t.Setenv(PerfReportEnv, reportPath)
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 50*time.Millisecond)
|
||||
tc.Record("mcp_call", 200*time.Millisecond)
|
||||
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
|
||||
data, err := os.ReadFile(reportPath)
|
||||
if err != nil {
|
||||
t.Fatalf("report file not written: %v", err)
|
||||
}
|
||||
|
||||
var report PerfReport
|
||||
if err := json.Unmarshal(data, &report); err != nil {
|
||||
t.Fatalf("invalid JSON in report: %v", err)
|
||||
}
|
||||
if report.Kind != "perf_report" {
|
||||
t.Errorf("expected kind 'perf_report', got %q", report.Kind)
|
||||
}
|
||||
if len(report.Phases) != 2 {
|
||||
t.Errorf("expected 2 phases, got %d", len(report.Phases))
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_Auto(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
expected := filepath.Join(tmpHome, ".dws", "perf", "latest.json")
|
||||
|
||||
// Temporarily override HOME for defaultPerfReportPath
|
||||
t.Setenv("HOME", tmpHome)
|
||||
t.Setenv(PerfReportEnv, "auto")
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("cmd_init", 10*time.Millisecond)
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
|
||||
if _, err := os.Stat(expected); err != nil {
|
||||
t.Fatalf("expected report at %s: %v", expected, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_Disabled(t *testing.T) {
|
||||
t.Setenv(PerfReportEnv, "")
|
||||
|
||||
tc := NewTimingCollector()
|
||||
tc.Record("op", 10*time.Millisecond)
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
// No file should be written; no error expected
|
||||
}
|
||||
|
||||
func TestWriteReportIfEnabled_NilCollector(t *testing.T) {
|
||||
t.Setenv(PerfReportEnv, "/tmp/should-not-exist.json")
|
||||
var tc *TimingCollector
|
||||
tc.WriteReportIfEnabled("v1.0.0", "dws version")
|
||||
}
|
||||
|
||||
func TestLoadLatestReport(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
t.Setenv("HOME", tmpHome)
|
||||
|
||||
perfDir := filepath.Join(tmpHome, ".dws", "perf")
|
||||
if err := os.MkdirAll(perfDir, 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
report := PerfReport{
|
||||
Kind: "perf_report",
|
||||
Version: "1",
|
||||
CLIVersion: "v1.0.0",
|
||||
Command: "dws version",
|
||||
TotalMs: 100,
|
||||
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}},
|
||||
Slowest: "cmd_init",
|
||||
OverheadMs: 50,
|
||||
}
|
||||
data, _ := json.MarshalIndent(report, "", " ")
|
||||
if err := os.WriteFile(filepath.Join(perfDir, "latest.json"), data, 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
loaded, err := LoadLatestReport()
|
||||
if err != nil {
|
||||
t.Fatalf("LoadLatestReport failed: %v", err)
|
||||
}
|
||||
if loaded.CLIVersion != "v1.0.0" {
|
||||
t.Errorf("expected cli_version 'v1.0.0', got %q", loaded.CLIVersion)
|
||||
}
|
||||
if len(loaded.Phases) != 1 {
|
||||
t.Errorf("expected 1 phase, got %d", len(loaded.Phases))
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadLatestReport_NotFound(t *testing.T) {
|
||||
tmpHome := t.TempDir()
|
||||
t.Setenv("HOME", tmpHome)
|
||||
|
||||
_, err := LoadLatestReport()
|
||||
if err == nil {
|
||||
t.Error("expected error when report file does not exist")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeCommand(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "no sensitive flags",
|
||||
args: []string{"dws", "aitable", "list-records"},
|
||||
want: "dws aitable list-records",
|
||||
},
|
||||
{
|
||||
name: "token with space-separated value",
|
||||
args: []string{"dws", "--token", "secret123", "version"},
|
||||
want: "dws --token *** version",
|
||||
},
|
||||
{
|
||||
name: "token with equals sign",
|
||||
args: []string{"dws", "--token=secret123", "version"},
|
||||
want: "dws --token=*** version",
|
||||
},
|
||||
{
|
||||
name: "client-secret space-separated",
|
||||
args: []string{"dws", "--client-secret", "mysecret", "--client-id", "myid", "auth"},
|
||||
want: "dws --client-secret *** --client-id *** auth",
|
||||
},
|
||||
{
|
||||
name: "client-id with equals",
|
||||
args: []string{"dws", "--client-id=abc123"},
|
||||
want: "dws --client-id=***",
|
||||
},
|
||||
{
|
||||
name: "empty args",
|
||||
args: []string{},
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := SanitizeCommand(tt.args)
|
||||
if got != tt.want {
|
||||
t.Errorf("SanitizeCommand(%v) = %q, want %q", tt.args, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePerfReportPath_Auto(t *testing.T) {
|
||||
p := resolvePerfReportPath("auto")
|
||||
if p == "" {
|
||||
t.Skip("HOME not available")
|
||||
}
|
||||
if !strings.HasSuffix(p, filepath.Join("perf", "latest.json")) {
|
||||
t.Errorf("expected path ending in perf/latest.json, got %q", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolvePerfReportPath_Custom(t *testing.T) {
|
||||
p := resolvePerfReportPath("/tmp/my-report.json")
|
||||
if p != "/tmp/my-report.json" {
|
||||
t.Errorf("expected '/tmp/my-report.json', got %q", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPerfReportSummary(t *testing.T) {
|
||||
report := &PerfReport{
|
||||
Command: "dws version",
|
||||
Timestamp: time.Now(),
|
||||
TotalMs: 300,
|
||||
Phases: []PerfPhase{{Name: "cmd_init", DurationMs: 50, Seq: 0}, {Name: "mcp_call", DurationMs: 200, Seq: 1}},
|
||||
Slowest: "mcp_call",
|
||||
OverheadMs: 50,
|
||||
}
|
||||
|
||||
var buf bytes.Buffer
|
||||
printPerfReportSummary(&buf, report)
|
||||
out := buf.String()
|
||||
|
||||
if !strings.Contains(out, "cmd_init") {
|
||||
t.Error("output should contain 'cmd_init'")
|
||||
}
|
||||
if !strings.Contains(out, "mcp_call") {
|
||||
t.Error("output should contain 'mcp_call'")
|
||||
}
|
||||
if !strings.Contains(out, "← 最慢") {
|
||||
t.Error("output should contain '← 最慢' marker")
|
||||
}
|
||||
if !strings.Contains(out, "总耗时") {
|
||||
t.Error("output should contain '总耗时'")
|
||||
}
|
||||
if !strings.Contains(out, "框架开销") {
|
||||
t.Error("output should contain '框架开销'")
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
+25
-1
@@ -13,7 +13,22 @@
|
||||
|
||||
package app
|
||||
|
||||
var version = "v1.0.2"
|
||||
var version = "dev"
|
||||
|
||||
// SetVersion overrides the version, build time and git commit strings.
|
||||
// Called by pkg/cli.SetVersion for overlay modules that inject their own
|
||||
// version info via ldflags.
|
||||
func SetVersion(v, bt, gc string) {
|
||||
if v != "" {
|
||||
version = v
|
||||
}
|
||||
if bt != "" {
|
||||
buildTime = bt
|
||||
}
|
||||
if gc != "" {
|
||||
gitCommit = gc
|
||||
}
|
||||
}
|
||||
|
||||
// Version returns the current CLI version string, including build metadata
|
||||
// when injected via ldflags (buildTime, gitCommit).
|
||||
@@ -23,3 +38,12 @@ func Version() string {
|
||||
}
|
||||
return version
|
||||
}
|
||||
|
||||
// RawVersion returns the bare version string without build metadata.
|
||||
func RawVersion() string { return version }
|
||||
|
||||
// BuildTime returns the build timestamp injected via ldflags.
|
||||
func BuildTime() string { return buildTime }
|
||||
|
||||
// GitCommit returns the git commit hash injected via ldflags.
|
||||
func GitCommit() string { return gitCommit }
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/helpers"
|
||||
)
|
||||
|
||||
const (
|
||||
// appConfigFile is the filename for storing app credentials.
|
||||
appConfigFile = "app.json"
|
||||
)
|
||||
|
||||
// AppConfig represents the application credentials configuration.
|
||||
// This is stored in ~/.dws/app.json with the client secret securely stored in keychain.
|
||||
type AppConfig struct {
|
||||
ClientID string `json:"clientId"`
|
||||
ClientSecret SecretInput `json:"clientSecret"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
UpdatedAt time.Time `json:"updatedAt,omitempty"`
|
||||
}
|
||||
|
||||
// Cached app config for performance (avoid repeated file reads).
|
||||
var (
|
||||
cachedAppConfig *AppConfig
|
||||
cachedAppConfigOnce sync.Once
|
||||
cachedAppConfigMu sync.RWMutex
|
||||
)
|
||||
|
||||
// Cached resolved credentials (avoid repeated keychain access).
|
||||
var (
|
||||
cachedResolvedID string
|
||||
cachedResolvedSecret string
|
||||
cachedResolvedValid bool
|
||||
cachedResolvedMu sync.RWMutex
|
||||
)
|
||||
|
||||
// GetAppConfigPath returns the path to the app config file.
|
||||
func GetAppConfigPath(configDir string) string {
|
||||
return filepath.Join(configDir, appConfigFile)
|
||||
}
|
||||
|
||||
// LoadAppConfig loads the app configuration from disk.
|
||||
// Returns nil, nil if the config file does not exist.
|
||||
func LoadAppConfig(configDir string) (*AppConfig, error) {
|
||||
path := GetAppConfigPath(configDir)
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, nil
|
||||
}
|
||||
return nil, fmt.Errorf("reading app config: %w", err)
|
||||
}
|
||||
|
||||
var config AppConfig
|
||||
if err := json.Unmarshal(data, &config); err != nil {
|
||||
return nil, fmt.Errorf("parsing app config: %w", err)
|
||||
}
|
||||
return &config, nil
|
||||
}
|
||||
|
||||
// SaveAppConfig saves the app configuration to disk.
|
||||
// If the client secret is a plain string, it will be stored in keychain
|
||||
// and the config file will contain a reference to it.
|
||||
func SaveAppConfig(configDir string, config *AppConfig) error {
|
||||
// Store plain secret in keychain, convert to reference
|
||||
if config.ClientSecret.IsPlain() && config.ClientID != "" {
|
||||
storedRef, err := StoreSecret(config.ClientID, config.ClientSecret)
|
||||
if err != nil {
|
||||
return fmt.Errorf("storing client secret: %w", err)
|
||||
}
|
||||
config.ClientSecret = storedRef
|
||||
}
|
||||
|
||||
// Update timestamps
|
||||
if config.CreatedAt.IsZero() {
|
||||
config.CreatedAt = time.Now()
|
||||
}
|
||||
config.UpdatedAt = time.Now()
|
||||
|
||||
data, err := json.MarshalIndent(config, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling app config: %w", err)
|
||||
}
|
||||
|
||||
path := GetAppConfigPath(configDir)
|
||||
if err := helpers.AtomicWriteJSON(path, append(data, '\n')); err != nil {
|
||||
return fmt.Errorf("writing app config: %w", err)
|
||||
}
|
||||
|
||||
// Update cache
|
||||
cachedAppConfigMu.Lock()
|
||||
cachedAppConfig = config
|
||||
cachedAppConfigMu.Unlock()
|
||||
|
||||
// Invalidate resolved credentials cache so next access re-resolves
|
||||
cachedResolvedMu.Lock()
|
||||
cachedResolvedValid = false
|
||||
cachedResolvedID = ""
|
||||
cachedResolvedSecret = ""
|
||||
cachedResolvedMu.Unlock()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteAppConfig removes the app configuration and associated keychain secrets.
|
||||
func DeleteAppConfig(configDir string) error {
|
||||
// Load existing config to clean up keychain
|
||||
existing, _ := LoadAppConfig(configDir)
|
||||
if existing != nil {
|
||||
RemoveSecretStore(existing.ClientSecret)
|
||||
}
|
||||
|
||||
// Remove config file
|
||||
path := GetAppConfigPath(configDir)
|
||||
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
|
||||
return fmt.Errorf("removing app config: %w", err)
|
||||
}
|
||||
|
||||
// Clear cache
|
||||
cachedAppConfigMu.Lock()
|
||||
cachedAppConfig = nil
|
||||
cachedAppConfigMu.Unlock()
|
||||
|
||||
// Clear resolved credentials cache
|
||||
cachedResolvedMu.Lock()
|
||||
cachedResolvedValid = false
|
||||
cachedResolvedID = ""
|
||||
cachedResolvedSecret = ""
|
||||
cachedResolvedMu.Unlock()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetCachedAppConfig returns the cached app configuration.
|
||||
// It loads from disk on first call and caches the result.
|
||||
// Returns nil if no configuration exists or loading fails.
|
||||
func GetCachedAppConfig(configDir string) *AppConfig {
|
||||
cachedAppConfigOnce.Do(func() {
|
||||
cfg, err := LoadAppConfig(configDir)
|
||||
if err == nil && cfg != nil {
|
||||
cachedAppConfigMu.Lock()
|
||||
cachedAppConfig = cfg
|
||||
cachedAppConfigMu.Unlock()
|
||||
}
|
||||
})
|
||||
|
||||
cachedAppConfigMu.RLock()
|
||||
defer cachedAppConfigMu.RUnlock()
|
||||
return cachedAppConfig
|
||||
}
|
||||
|
||||
// ReloadAppConfig forces a reload of the app configuration from disk.
|
||||
// This should be called after SaveAppConfig to ensure the cache is updated.
|
||||
func ReloadAppConfig(configDir string) (*AppConfig, error) {
|
||||
cfg, err := LoadAppConfig(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
cachedAppConfigMu.Lock()
|
||||
cachedAppConfig = cfg
|
||||
cachedAppConfigMu.Unlock()
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// HasAppConfig returns true if an app configuration file exists.
|
||||
func HasAppConfig(configDir string) bool {
|
||||
path := GetAppConfigPath(configDir)
|
||||
_, err := os.Stat(path)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// ResolveAppCredentials resolves the client ID and secret from the app config.
|
||||
// Results are cached to avoid repeated keychain access.
|
||||
// Returns empty strings if the config doesn't exist or resolution fails.
|
||||
func ResolveAppCredentials(configDir string) (clientID, clientSecret string) {
|
||||
// Fast path: check cache first
|
||||
cachedResolvedMu.RLock()
|
||||
if cachedResolvedValid {
|
||||
id, secret := cachedResolvedID, cachedResolvedSecret
|
||||
cachedResolvedMu.RUnlock()
|
||||
return id, secret
|
||||
}
|
||||
cachedResolvedMu.RUnlock()
|
||||
|
||||
// Slow path: load and cache
|
||||
cachedResolvedMu.Lock()
|
||||
defer cachedResolvedMu.Unlock()
|
||||
// Double-check after acquiring write lock
|
||||
if cachedResolvedValid {
|
||||
return cachedResolvedID, cachedResolvedSecret
|
||||
}
|
||||
|
||||
cfg := GetCachedAppConfig(configDir)
|
||||
if cfg != nil {
|
||||
cachedResolvedID = cfg.ClientID
|
||||
if secret, err := ResolveSecret(cfg.ClientSecret); err == nil {
|
||||
cachedResolvedSecret = secret
|
||||
}
|
||||
}
|
||||
cachedResolvedValid = true
|
||||
return cachedResolvedID, cachedResolvedSecret
|
||||
}
|
||||
@@ -26,6 +26,7 @@ func TestClientID_RuntimeOverride(t *testing.T) {
|
||||
|
||||
func TestClientID_EnvFallback(t *testing.T) {
|
||||
SetClientID("")
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
|
||||
t.Setenv("DWS_CLIENT_ID", "env-id")
|
||||
if got := ClientID(); got != "env-id" {
|
||||
t.Fatalf("expected env-id, got %s", got)
|
||||
@@ -34,9 +35,17 @@ func TestClientID_EnvFallback(t *testing.T) {
|
||||
|
||||
func TestClientID_Default(t *testing.T) {
|
||||
SetClientID("")
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
|
||||
t.Setenv("DWS_CLIENT_ID", "")
|
||||
if got := ClientID(); got != DefaultClientID {
|
||||
t.Fatalf("expected default, got %s", got)
|
||||
// When DefaultClientID is a placeholder (starts with "<"), ClientID() returns empty string
|
||||
if strings.HasPrefix(DefaultClientID, "<") {
|
||||
if got := ClientID(); got != "" {
|
||||
t.Fatalf("expected empty string for placeholder, got %s", got)
|
||||
}
|
||||
} else {
|
||||
if got := ClientID(); got != DefaultClientID {
|
||||
t.Fatalf("expected default, got %s", got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -51,6 +60,7 @@ func TestClientSecret_RuntimeOverride(t *testing.T) {
|
||||
|
||||
func TestClientSecret_EnvFallback(t *testing.T) {
|
||||
SetClientSecret("")
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
|
||||
t.Setenv("DWS_CLIENT_SECRET", "env-secret")
|
||||
if got := ClientSecret(); got != "env-secret" {
|
||||
t.Fatalf("expected env-secret, got %s", got)
|
||||
@@ -59,6 +69,7 @@ func TestClientSecret_EnvFallback(t *testing.T) {
|
||||
|
||||
func TestClientSecret_Default(t *testing.T) {
|
||||
SetClientSecret("")
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir()) // Use temp dir to avoid reading persisted config
|
||||
t.Setenv("DWS_CLIENT_SECRET", "")
|
||||
if got := ClientSecret(); got != DefaultClientSecret {
|
||||
t.Fatalf("expected default, got %s", got)
|
||||
|
||||
@@ -0,0 +1,840 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// setupMCPConfigDir creates a temp config dir with mcp_url pointing to the
|
||||
// given test server and sets DWS_CONFIG_DIR via t.Setenv.
|
||||
// NOTE: tests calling this must NOT use t.Parallel().
|
||||
func setupMCPConfigDir(t *testing.T, srvURL string) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
os.WriteFile(filepath.Join(dir, "mcp_url"), []byte(srvURL), 0o600)
|
||||
t.Setenv("DWS_CONFIG_DIR", dir)
|
||||
return dir
|
||||
}
|
||||
|
||||
// resetClientIDFromMCP clears the MCP-sourced flag (test helper).
|
||||
func resetClientIDFromMCP() {
|
||||
clientMu.Lock()
|
||||
defer clientMu.Unlock()
|
||||
clientIDFromMCP = false
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 1. CheckCLIAuthEnabled: interface error → fail-closed with retry
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error from CheckCLIAuthEnabled when server returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Server 500 → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
configDir := setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 2 * time.Second},
|
||||
}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ Connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `{this is not valid json}`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Timeout_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(5 * time.Second)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 200 * time.Millisecond},
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(ctx, "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error on request timeout, got nil")
|
||||
}
|
||||
t.Logf("✅ Timeout → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 2. CheckCLIAuthEnabled: transient error then recovery → succeeds
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !status.Success || !status.Result.CLIAuthEnabled {
|
||||
t.Fatalf("expected CLIAuthEnabled=true, got %+v", status)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ Transient error then success: attempts=%d, enabled=%v", calls.Load(), status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 3. CheckCLIAuthEnabled: normal responses (pass-through)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=true")
|
||||
}
|
||||
t.Logf("✅ Normal enabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=false")
|
||||
}
|
||||
t.Logf("✅ Normal disabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 4. OAuth callback: simulates the fail-closed logic at the /callback level
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestOAuthCallback_CLIAuthError_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error = fmt.Errorf("simulated network error")
|
||||
var authStatus *CLIAuthStatus
|
||||
_ = authStatus
|
||||
|
||||
// This is the exact expression used in oauth_provider.go callback:
|
||||
// cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
cliAuthEnabled := statusErr == nil // false
|
||||
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when statusErr != nil")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=%v → cliAuthEnabled=%v → shows notEnabledHTML (fail-closed)", statusErr, cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthEnabled_ShowsSuccessPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
if !cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be true when API returns enabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=true → cliAuthEnabled=%v → shows successHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthDisabledByServer_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when server says disabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=false → cliAuthEnabled=%v → shows notEnabledHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 5. Device Flow: loginOnce with broken cliAuthEnabled endpoint
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
|
||||
default:
|
||||
http.Error(w, "not found: "+r.URL.Path, http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth check fails, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "检查 CLI 授权状态失败") && !strings.Contains(err.Error(), "Failed to check CLI auth status") {
|
||||
t.Fatalf("unexpected error message: %s", err)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth check error → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, SuperAdminPath):
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "张三"}},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth is disabled, got nil")
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth disabled by server → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
token, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected loginOnce to succeed, got error: %v", err)
|
||||
}
|
||||
if token.AccessToken != "test-access-token" {
|
||||
t.Fatalf("unexpected token: %s", token.AccessToken)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth enabled → login succeeded, token=%s", token.AccessToken)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 6. FetchClientIDFromMCP: /cli/clientId error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestFetchClientIDFromMCP_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/clientId returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId 500 → fail-closed with retry: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `not json at all`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_BusinessError_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: false,
|
||||
ErrorCode: "forbidden",
|
||||
ErrorMsg: "access denied",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when server returns success=false, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "access denied") {
|
||||
t.Fatalf("error should contain server error message, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId business error → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "recovered-client-id",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if id != "recovered-client-id" {
|
||||
t.Fatalf("expected client ID 'recovered-client-id', got %q", id)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId transient then success: attempts=%d, id=%s", calls.Load(), id)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != ClientIDPath {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "my-client-id-123",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if id != "my-client-id-123" {
|
||||
t.Fatalf("expected 'my-client-id-123', got %q", id)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId normal success: id=%s", id)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 7. GetSuperAdmins: /cli/superAdmin error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetSuperAdmins_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/superAdmin returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "a1", Name: "张三"}, {StaffID: "a2", Name: "李四"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !result.Success || len(result.Result) != 2 {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin transient then success: attempts=%d, admins=%v", calls.Load(), result.Result)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "王五"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(result.Result) != 1 || result.Result[0].Name != "王五" {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin normal success: admins=%v", result.Result)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 8. SendCliAuthApply: /cli/sendCliAuthApply error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendCliAuthApply_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/sendCliAuthApply returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 1 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failure, got error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 2 {
|
||||
t.Fatalf("expected 2 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply transient then success: attempts=%d", calls.Load())
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
if !strings.Contains(r.URL.RawQuery, "adminStaffId=admin123") {
|
||||
t.Errorf("missing or wrong adminStaffId param: %s", r.URL.RawQuery)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "good-token", "admin123")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply normal success: result=%+v", result)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_BusinessError(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{
|
||||
Success: false,
|
||||
ErrorCode: "invalid_admin",
|
||||
ErrorMsg: "admin not found",
|
||||
Result: false,
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "nonexistent")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected transport error: %v", err)
|
||||
}
|
||||
if result.Success {
|
||||
t.Fatal("expected success=false for business error")
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply business error: errorCode=%s, errorMsg=%s", result.ErrorCode, result.ErrorMsg)
|
||||
}
|
||||
@@ -26,8 +26,8 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/fatih/color"
|
||||
)
|
||||
|
||||
@@ -96,6 +96,23 @@ type serviceResult struct {
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) Login(ctx context.Context) (*TokenData, error) {
|
||||
// Ensure we have a valid client ID (fetch from MCP if not available)
|
||||
if p.clientID == "" {
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("client ID not configured, fetching from MCP server")
|
||||
}
|
||||
mcpClientID, mcpErr := FetchClientIDFromMCP(ctx)
|
||||
if mcpErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("获取 Client ID 失败"), mcpErr)
|
||||
}
|
||||
p.clientID = mcpClientID
|
||||
// Mark that clientID is from MCP
|
||||
SetClientIDFromMCP(mcpClientID)
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("fetched client ID from MCP server", "clientID", mcpClientID)
|
||||
}
|
||||
}
|
||||
|
||||
const maxAttempts = 3
|
||||
for attempt := 1; attempt <= maxAttempts; attempt++ {
|
||||
tokenData, err := p.loginOnce(ctx, attempt)
|
||||
@@ -149,9 +166,54 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
|
||||
}
|
||||
|
||||
// Check if CLI auth is enabled for this organization (fail-closed: block on error)
|
||||
dfPrintStep(p.output(), 4, i18n.T("检查组织 CLI 授权状态..."), 0)
|
||||
authStatus, authErr := oauthProvider.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
|
||||
if authErr != nil {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 无法检查 CLI 数据访问权限状态")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请检查网络连接后重试。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("检查 CLI 授权状态失败"), authErr)
|
||||
} else if authStatus.Success && !authStatus.Result.CLIAuthEnabled {
|
||||
// CLI auth is disabled - show detailed error with admin info
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
|
||||
// Try to get super admin list
|
||||
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
|
||||
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
|
||||
// Show up to 3 admins
|
||||
maxAdmins := 3
|
||||
if len(admins.Result) < maxAdmins {
|
||||
maxAdmins = len(admins.Result)
|
||||
}
|
||||
var adminNames []string
|
||||
for i := 0; i < maxAdmins; i++ {
|
||||
adminNames = append(adminNames, admins.Result[i].Name)
|
||||
}
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
|
||||
}
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := SaveTokenData(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
oauthProvider.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -46,6 +46,10 @@ func writeServiceResult(w http.ResponseWriter, success bool, result any, errCode
|
||||
func TestRequestDeviceCodeSuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Set a test client ID
|
||||
SetClientID("test-client-id")
|
||||
t.Cleanup(func() { SetClientID("") })
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
t.Fatalf("method = %s, want POST", r.Method)
|
||||
|
||||
+170
-5
@@ -15,9 +15,31 @@ package auth
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_ID",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppKey (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// AuthorizeURL is the DingTalk OAuth authorization page.
|
||||
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
|
||||
@@ -58,16 +80,107 @@ const (
|
||||
|
||||
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
|
||||
LogoutContinueURL = "https://login.dingtalk.com"
|
||||
|
||||
// MCP API endpoints for CLI authorization management.
|
||||
DefaultMCPBaseURL = "https://mcp.dingtalk.com"
|
||||
CLIAuthEnabledPath = "/cli/cliAuthEnabled"
|
||||
SuperAdminPath = "/cli/superAdmin"
|
||||
SendCliAuthApplyPath = "/cli/sendCliAuthApply"
|
||||
ClientIDPath = "/cli/clientId"
|
||||
|
||||
// MCP OAuth endpoints (used when clientId is fetched from MCP).
|
||||
MCPOAuthTokenPath = "/oauth2/getToken"
|
||||
MCPRefreshTokenPath = "/oauth2/refreshToken"
|
||||
MCPRevokeTokenPath = "/oauth2/revokeToken"
|
||||
)
|
||||
|
||||
// GetMCPBaseURL returns the MCP base URL with priority:
|
||||
// 1. ~/.dws/mcp_url file content (for pre-release environment)
|
||||
// 2. Default value (https://mcp.dingtalk.com)
|
||||
func GetMCPBaseURL() string {
|
||||
mcpURLPath := filepath.Join(getDefaultConfigDir(), "mcp_url")
|
||||
if data, err := os.ReadFile(mcpURLPath); err == nil {
|
||||
if url := strings.TrimSpace(string(data)); url != "" {
|
||||
return url
|
||||
}
|
||||
}
|
||||
return DefaultMCPBaseURL
|
||||
}
|
||||
|
||||
// Runtime overrides set via CLI flags (--client-id, --client-secret).
|
||||
// These take highest priority over environment variables and defaults.
|
||||
var (
|
||||
clientMu sync.RWMutex
|
||||
runtimeClientID string
|
||||
runtimeClientSecret string
|
||||
// clientIDFromMCP indicates whether the clientID was fetched from MCP server.
|
||||
// When true, MCP OAuth endpoints should be used instead of direct DingTalk API.
|
||||
clientIDFromMCP bool
|
||||
)
|
||||
|
||||
// SetClientIDFromMCP sets the clientID fetched from MCP server and marks it as MCP-sourced.
|
||||
func SetClientIDFromMCP(id string) {
|
||||
clientMu.Lock()
|
||||
defer clientMu.Unlock()
|
||||
runtimeClientID = id
|
||||
clientIDFromMCP = true
|
||||
}
|
||||
|
||||
// IsClientIDFromMCP returns true if the current clientID was fetched from MCP server.
|
||||
func IsClientIDFromMCP() bool {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return clientIDFromMCP || edition.Get().AuthClientFromMCP
|
||||
}
|
||||
|
||||
// GetUserAccessTokenURL returns the appropriate token exchange URL.
|
||||
// Uses MCP endpoint when clientID is from MCP, otherwise uses direct DingTalk API.
|
||||
func GetUserAccessTokenURL() string {
|
||||
if IsClientIDFromMCP() {
|
||||
return GetMCPBaseURL() + MCPOAuthTokenPath
|
||||
}
|
||||
return UserAccessTokenURL
|
||||
}
|
||||
|
||||
// GetRefreshTokenURL returns the appropriate token refresh URL.
|
||||
// Uses MCP endpoint when clientID is from MCP, otherwise uses direct DingTalk API.
|
||||
func GetRefreshTokenURL() string {
|
||||
if IsClientIDFromMCP() {
|
||||
return GetMCPBaseURL() + MCPRefreshTokenPath
|
||||
}
|
||||
return UserAccessTokenURL // DingTalk uses same endpoint for refresh
|
||||
}
|
||||
|
||||
// GetRevokeTokenURL returns the token revocation URL (MCP only).
|
||||
// Returns empty string if not using MCP mode.
|
||||
func GetRevokeTokenURL() string {
|
||||
if IsClientIDFromMCP() {
|
||||
return GetMCPBaseURL() + MCPRevokeTokenPath
|
||||
}
|
||||
return "" // Direct mode doesn't have revoke endpoint
|
||||
}
|
||||
|
||||
// resolveCredentialSource determines the source of the current credentials.
|
||||
// Returns one of: "flag", "env", "app", "default".
|
||||
// This is used to track where credentials came from for token refresh.
|
||||
func resolveCredentialSource() string {
|
||||
clientMu.RLock()
|
||||
hasRuntimeOverride := runtimeClientID != "" || runtimeClientSecret != ""
|
||||
clientMu.RUnlock()
|
||||
|
||||
if hasRuntimeOverride {
|
||||
return "flag"
|
||||
}
|
||||
// Check if loaded from app config
|
||||
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
|
||||
return "app"
|
||||
}
|
||||
if os.Getenv("DWS_CLIENT_ID") != "" || os.Getenv("DWS_CLIENT_SECRET") != "" {
|
||||
return "env"
|
||||
}
|
||||
return "default"
|
||||
}
|
||||
|
||||
// SetClientID allows runtime override of the client ID (e.g., from CLI flags).
|
||||
func SetClientID(id string) {
|
||||
clientMu.Lock()
|
||||
@@ -84,8 +197,11 @@ func SetClientSecret(secret string) {
|
||||
|
||||
// ClientID returns the OAuth client ID with priority:
|
||||
// 1. Runtime override (CLI flag --client-id)
|
||||
// 2. Environment variable (DWS_CLIENT_ID)
|
||||
// 3. Default hardcoded value
|
||||
// 2. Persisted app config (from previous login)
|
||||
// 3. Environment variable (DWS_CLIENT_ID)
|
||||
// 4. Default hardcoded value (if not a placeholder)
|
||||
// Returns empty string if no valid client ID is available.
|
||||
// Note: MCP server fetch (priority 4 in the full flow) is handled in OAuthProvider.Login()
|
||||
func ClientID() string {
|
||||
clientMu.RLock()
|
||||
override := runtimeClientID
|
||||
@@ -93,16 +209,28 @@ func ClientID() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
if id := edition.Get().AuthClientID; id != "" {
|
||||
return id
|
||||
}
|
||||
// Try loading from persisted app config
|
||||
if id, _ := ResolveAppCredentials(getDefaultConfigDir()); id != "" {
|
||||
return id
|
||||
}
|
||||
if v := os.Getenv("DWS_CLIENT_ID"); v != "" {
|
||||
return v
|
||||
}
|
||||
return DefaultClientID
|
||||
// Only return default if it's not a placeholder
|
||||
if !strings.HasPrefix(DefaultClientID, "<") {
|
||||
return DefaultClientID
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// ClientSecret returns the OAuth client secret with priority:
|
||||
// 1. Runtime override (CLI flag --client-secret)
|
||||
// 2. Environment variable (DWS_CLIENT_SECRET)
|
||||
// 3. Default hardcoded value
|
||||
// 2. Persisted app config (from previous login, stored in keychain)
|
||||
// 3. Environment variable (DWS_CLIENT_SECRET)
|
||||
// 4. Default hardcoded value
|
||||
func ClientSecret() string {
|
||||
clientMu.RLock()
|
||||
override := runtimeClientSecret
|
||||
@@ -110,8 +238,45 @@ func ClientSecret() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
// Try loading from persisted app config (secret is in keychain)
|
||||
if _, secret := ResolveAppCredentials(getDefaultConfigDir()); secret != "" {
|
||||
return secret
|
||||
}
|
||||
if v := os.Getenv("DWS_CLIENT_SECRET"); v != "" {
|
||||
return v
|
||||
}
|
||||
return DefaultClientSecret
|
||||
}
|
||||
|
||||
// HasValidClientSecret returns true if a valid client secret is available.
|
||||
// A valid secret is one that is not a placeholder (e.g., <YOUR_CLIENT_SECRET>).
|
||||
func HasValidClientSecret() bool {
|
||||
secret := ClientSecret()
|
||||
return secret != "" && !strings.HasPrefix(secret, "<")
|
||||
}
|
||||
|
||||
// getRuntimeCredentials returns the runtime-override credentials if set.
|
||||
// Returns empty strings if no runtime overrides were provided.
|
||||
func getRuntimeCredentials() (clientID, clientSecret string) {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return runtimeClientID, runtimeClientSecret
|
||||
}
|
||||
|
||||
// getEnvClientID returns the environment variable client ID if set.
|
||||
func getEnvClientID() string {
|
||||
return os.Getenv("DWS_CLIENT_ID")
|
||||
}
|
||||
|
||||
// getDefaultConfigDir returns the default configuration directory.
|
||||
// Priority: DWS_CONFIG_DIR env var > ~/.dws
|
||||
func getDefaultConfigDir() string {
|
||||
if envDir := os.Getenv("DWS_CONFIG_DIR"); envDir != "" {
|
||||
return envDir
|
||||
}
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return ".dws"
|
||||
}
|
||||
return filepath.Join(homeDir, ".dws")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -105,3 +105,45 @@ func EnsureMigration(configDir string, logger *slog.Logger) {
|
||||
func IsMigrationDone() bool {
|
||||
return migrationDone
|
||||
}
|
||||
|
||||
// Client credential storage functions.
|
||||
// These store the clientSecret associated with a specific clientId,
|
||||
// allowing token refresh to work even if environment variables change.
|
||||
|
||||
const clientSecretPrefix = "client-secret:"
|
||||
|
||||
// SaveClientSecret stores the client secret for a specific client ID.
|
||||
// This is called during login to snapshot the credentials used.
|
||||
func SaveClientSecret(clientID, clientSecret string) error {
|
||||
if clientID == "" || clientSecret == "" {
|
||||
return nil // Nothing to save
|
||||
}
|
||||
account := clientSecretPrefix + clientID
|
||||
if err := keychain.Set(keychain.Service, account, clientSecret); err != nil {
|
||||
return fmt.Errorf("save client secret: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadClientSecret retrieves the stored client secret for a specific client ID.
|
||||
// Returns empty string if not found.
|
||||
func LoadClientSecret(clientID string) string {
|
||||
if clientID == "" {
|
||||
return ""
|
||||
}
|
||||
account := clientSecretPrefix + clientID
|
||||
secret, err := keychain.Get(keychain.Service, account)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return secret
|
||||
}
|
||||
|
||||
// DeleteClientSecret removes the stored client secret for a specific client ID.
|
||||
func DeleteClientSecret(clientID string) error {
|
||||
if clientID == "" {
|
||||
return nil
|
||||
}
|
||||
account := clientSecretPrefix + clientID
|
||||
return keychain.Remove(keychain.Service, account)
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
|
||||
+1022
-12
File diff suppressed because it is too large
Load Diff
+326
-12
@@ -15,6 +15,7 @@ package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
@@ -22,6 +23,7 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
@@ -92,6 +94,23 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
}
|
||||
|
||||
// Fall through: full browser OAuth flow.
|
||||
// Ensure we have a valid client ID (fetch from MCP if not available)
|
||||
if p.clientID == "" {
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("client ID not configured, fetching from MCP server")
|
||||
}
|
||||
mcpClientID, mcpErr := FetchClientIDFromMCP(ctx)
|
||||
if mcpErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("获取 Client ID 失败"), mcpErr)
|
||||
}
|
||||
p.clientID = mcpClientID
|
||||
// Mark that clientID is from MCP, so we use MCP OAuth endpoints
|
||||
SetClientIDFromMCP(mcpClientID)
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("fetched client ID from MCP server", "clientID", mcpClientID)
|
||||
}
|
||||
}
|
||||
|
||||
// Find a free port for the callback server.
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
@@ -100,15 +119,79 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
port := listener.Addr().(*net.TCPAddr).Port
|
||||
redirectURI := fmt.Sprintf("http://127.0.0.1:%d%s", port, CallbackPath)
|
||||
|
||||
codeCh := make(chan string, 1)
|
||||
// Channel to pass callback result (token data or error with CLI auth status)
|
||||
type callbackResult struct {
|
||||
token *TokenData
|
||||
err error
|
||||
cliAuthDisabled bool
|
||||
}
|
||||
resultCh := make(chan callbackResult, 1)
|
||||
errCh := make(chan error, 1)
|
||||
|
||||
// Shared state for API handlers (protected by mutex)
|
||||
var (
|
||||
callbackToken *TokenData
|
||||
callbackProcessedCode string // The auth code that has been successfully processed
|
||||
callbackAuthDisabled bool
|
||||
callbackApplySent bool // Whether apply request was sent
|
||||
callbackSelectedAdminId string // Selected admin ID for apply
|
||||
callbackCodeInProgress string // Code currently being processed (to prevent concurrent exchange)
|
||||
callbackTokenMu sync.Mutex
|
||||
)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc(CallbackPath, func(w http.ResponseWriter, r *http.Request) {
|
||||
// Get code first to check if this is a new authorization or page refresh
|
||||
code := r.URL.Query().Get("authCode")
|
||||
if code == "" {
|
||||
code = r.URL.Query().Get("code")
|
||||
}
|
||||
|
||||
// Check state and handle page refresh or concurrent requests
|
||||
callbackTokenMu.Lock()
|
||||
processedCode := callbackProcessedCode
|
||||
processedAuthDisabled := callbackAuthDisabled
|
||||
codeInProgress := callbackCodeInProgress
|
||||
hasToken := callbackToken != nil
|
||||
|
||||
// Case 1: This code was already successfully processed - show cached page
|
||||
if code != "" && code == processedCode {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if processedAuthDisabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Case 2: This code is being processed by another request - show wait page
|
||||
if code != "" && code == codeInProgress {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, `<html><head><meta http-equiv="refresh" content="1"></head><body><p>正在处理授权,请稍候...</p></body></html>`)
|
||||
return
|
||||
}
|
||||
|
||||
// Case 3: No code but we have a processed token - show cached page
|
||||
if code == "" && hasToken {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if processedAuthDisabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Case 4: New code - mark as in-progress and process
|
||||
if code != "" {
|
||||
callbackCodeInProgress = code
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
if code == "" {
|
||||
select {
|
||||
case errCh <- errors.New(i18n.T("回调中未收到授权码")):
|
||||
@@ -118,14 +201,149 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
_, _ = fmt.Fprint(w, i18n.T("授权失败:未收到授权码"))
|
||||
return
|
||||
}
|
||||
select {
|
||||
case codeCh <- code:
|
||||
|
||||
// Exchange code for token
|
||||
tokenData, exchangeErr := p.exchangeCode(ctx, code)
|
||||
if exchangeErr != nil {
|
||||
// Clear in-progress state on error
|
||||
callbackTokenMu.Lock()
|
||||
if callbackCodeInProgress == code {
|
||||
callbackCodeInProgress = ""
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
default:
|
||||
// Select already exited (timeout/cancel); discard late callback.
|
||||
w.WriteHeader(http.StatusGone)
|
||||
_, _ = fmt.Fprintf(w, "<html><body><h1>授权失败</h1><p>%s</p></body></html>", exchangeErr.Error())
|
||||
select {
|
||||
case resultCh <- callbackResult{err: exchangeErr}:
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Mark as processed immediately after successful exchange
|
||||
callbackTokenMu.Lock()
|
||||
previouslyProcessed := callbackProcessedCode != ""
|
||||
callbackToken = tokenData
|
||||
callbackProcessedCode = code // Remember this code was successfully processed
|
||||
callbackCodeInProgress = "" // Clear in-progress state
|
||||
// Reset apply state for new authorization (user switched org)
|
||||
if previouslyProcessed {
|
||||
callbackApplySent = false
|
||||
callbackSelectedAdminId = ""
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Check CLI auth enabled status (fail-closed: treat errors as disabled)
|
||||
authStatus, statusErr := p.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
|
||||
// Update CLI auth disabled state
|
||||
callbackTokenMu.Lock()
|
||||
callbackAuthDisabled = !cliAuthEnabled
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Display appropriate HTML based on CLI auth status
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if !cliAuthEnabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
// Ensure response is flushed to client
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
// Notify main goroutine with full result
|
||||
select {
|
||||
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled}:
|
||||
default:
|
||||
}
|
||||
})
|
||||
|
||||
// API endpoint: get super admins
|
||||
mux.HandleFunc("/api/superAdmin", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := GetSuperAdmins(ctx, token.AccessToken)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// API endpoint: send CLI auth apply
|
||||
mux.HandleFunc("/api/sendApply", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
adminStaffID := r.URL.Query().Get("adminStaffId")
|
||||
if adminStaffID == "" {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"缺少 adminStaffId 参数"}`))
|
||||
return
|
||||
}
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := SendCliAuthApply(ctx, token.AccessToken, adminStaffID)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
// Mark apply as sent and save selected admin on success
|
||||
if result.Success && result.Result {
|
||||
callbackTokenMu.Lock()
|
||||
callbackApplySent = true
|
||||
callbackSelectedAdminId = adminStaffID
|
||||
callbackTokenMu.Unlock()
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// API endpoint: get current status (clientId, applySent, selectedAdminId)
|
||||
mux.HandleFunc("/api/status", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
applySent := callbackApplySent
|
||||
selectedAdminId := callbackSelectedAdminId
|
||||
callbackTokenMu.Unlock()
|
||||
_, _ = fmt.Fprintf(w, `{"clientId":"%s","applySent":%t,"selectedAdminId":"%s"}`, p.clientID, applySent, selectedAdminId)
|
||||
})
|
||||
|
||||
// API endpoint: check CLI auth enabled status
|
||||
mux.HandleFunc("/api/cliAuthEnabled", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := p.CheckCLIAuthEnabled(ctx, token.AccessToken)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// Success page endpoint
|
||||
mux.HandleFunc("/success", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
})
|
||||
|
||||
server := &http.Server{Handler: mux}
|
||||
@@ -161,9 +379,9 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
timeout := time.NewTimer(5 * time.Minute)
|
||||
defer timeout.Stop()
|
||||
|
||||
var authCode string
|
||||
var result callbackResult
|
||||
select {
|
||||
case authCode = <-codeCh:
|
||||
case result = <-resultCh:
|
||||
case err := <-errCh:
|
||||
return nil, err
|
||||
case <-timeout.C:
|
||||
@@ -172,13 +390,82 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
tokenData, err := p.exchangeCode(ctx, authCode)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
|
||||
// Handle callback errors
|
||||
if result.err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), result.err)
|
||||
}
|
||||
|
||||
// Handle CLI auth disabled - keep server running for user to apply
|
||||
if result.cliAuthDisabled {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
|
||||
|
||||
// Poll for CLI auth status while waiting
|
||||
applyTimeout := time.NewTimer(10 * time.Minute)
|
||||
defer applyTimeout.Stop()
|
||||
pollTicker := time.NewTicker(5 * time.Second)
|
||||
defer pollTicker.Stop()
|
||||
|
||||
elapsedSeconds := 0
|
||||
for {
|
||||
select {
|
||||
case <-applyTimeout.C:
|
||||
return nil, errors.New(i18n.T("操作超时,请重新登录"))
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-pollTicker.C:
|
||||
elapsedSeconds += 5
|
||||
|
||||
// Get latest token and state (user may have switched org)
|
||||
callbackTokenMu.Lock()
|
||||
currentToken := callbackToken
|
||||
currentAuthDisabled := callbackAuthDisabled
|
||||
applySent := callbackApplySent
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Check if user switched to an org with CLI auth enabled
|
||||
if currentToken != nil && !currentAuthDisabled {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
|
||||
time.Sleep(2 * time.Second)
|
||||
result.token = currentToken
|
||||
result.cliAuthDisabled = false
|
||||
goto continueLogin
|
||||
}
|
||||
|
||||
// Check if CLI auth is now enabled (admin approved)
|
||||
if currentToken != nil {
|
||||
authStatus, err := p.CheckCLIAuthEnabled(ctx, currentToken.AccessToken)
|
||||
if err == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
|
||||
time.Sleep(2 * time.Second)
|
||||
result.token = currentToken
|
||||
result.cliAuthDisabled = false
|
||||
goto continueLogin
|
||||
}
|
||||
}
|
||||
|
||||
// Show polling status based on apply state
|
||||
if applySent {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待管理员审批中"), elapsedSeconds)
|
||||
} else {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待提交申请中"), elapsedSeconds)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
continueLogin:
|
||||
tokenData := result.token
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := SaveTokenData(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
p.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
@@ -290,3 +577,30 @@ func (p *OAuthProvider) Logout() error {
|
||||
func (p *OAuthProvider) Status() (*TokenData, error) {
|
||||
return LoadTokenData(p.configDir)
|
||||
}
|
||||
|
||||
// persistAppConfigIfNeeded saves app credentials if custom ones were used.
|
||||
// This ensures the client secret is available for future token refreshes.
|
||||
func (p *OAuthProvider) persistAppConfigIfNeeded() {
|
||||
// Check if custom credentials were provided via runtime flags
|
||||
clientID, clientSecret := getRuntimeCredentials()
|
||||
if clientID == "" || clientSecret == "" {
|
||||
return
|
||||
}
|
||||
|
||||
// Only persist if they differ from environment/default values
|
||||
envID := getEnvClientID()
|
||||
if clientID == envID || clientID == DefaultClientID {
|
||||
return
|
||||
}
|
||||
|
||||
// Save app config with secret stored in keychain
|
||||
config := &AppConfig{
|
||||
ClientID: clientID,
|
||||
ClientSecret: PlainSecret(clientSecret),
|
||||
}
|
||||
if err := SaveAppConfig(p.configDir, config); err != nil {
|
||||
if p.logger != nil {
|
||||
p.logger.Warn("failed to persist app credentials", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
const (
|
||||
// secretKeyPrefix is the keychain account prefix for app secrets.
|
||||
secretKeyPrefix = "appsecret:"
|
||||
)
|
||||
|
||||
// SecretRef references a secret stored externally.
|
||||
type SecretRef struct {
|
||||
Source string `json:"source"` // "keychain" | "file"
|
||||
ID string `json:"id"` // keychain key or file path
|
||||
}
|
||||
|
||||
// SecretInput represents a secret value: either a plain string or a SecretRef object.
|
||||
type SecretInput struct {
|
||||
Plain string // non-empty for plain string values
|
||||
Ref *SecretRef // non-nil for SecretRef values
|
||||
}
|
||||
|
||||
// PlainSecret creates a SecretInput from a plain string.
|
||||
func PlainSecret(s string) SecretInput {
|
||||
return SecretInput{Plain: s}
|
||||
}
|
||||
|
||||
// IsZero returns true if the SecretInput has no value.
|
||||
func (s SecretInput) IsZero() bool {
|
||||
return s.Plain == "" && s.Ref == nil
|
||||
}
|
||||
|
||||
// IsSecretRef returns true if this is a SecretRef object.
|
||||
func (s SecretInput) IsSecretRef() bool {
|
||||
return s.Ref != nil
|
||||
}
|
||||
|
||||
// IsPlain returns true if this is a plain text string (not a SecretRef).
|
||||
func (s SecretInput) IsPlain() bool {
|
||||
return s.Ref == nil && s.Plain != ""
|
||||
}
|
||||
|
||||
// MarshalJSON serializes SecretInput: plain string → JSON string, SecretRef → JSON object.
|
||||
func (s SecretInput) MarshalJSON() ([]byte, error) {
|
||||
if s.Ref != nil {
|
||||
return json.Marshal(s.Ref)
|
||||
}
|
||||
return json.Marshal(s.Plain)
|
||||
}
|
||||
|
||||
// UnmarshalJSON deserializes SecretInput from either a JSON string or a SecretRef object.
|
||||
func (s *SecretInput) UnmarshalJSON(data []byte) error {
|
||||
// Try string first
|
||||
var plain string
|
||||
if err := json.Unmarshal(data, &plain); err == nil {
|
||||
s.Plain = plain
|
||||
s.Ref = nil
|
||||
return nil
|
||||
}
|
||||
// Try SecretRef object
|
||||
var ref SecretRef
|
||||
if err := json.Unmarshal(data, &ref); err == nil && isValidSource(ref.Source) && ref.ID != "" {
|
||||
s.Ref = &ref
|
||||
s.Plain = ""
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("clientSecret must be a string or {source, id} object")
|
||||
}
|
||||
|
||||
// ValidSecretSources is the set of recognized SecretRef sources.
|
||||
var ValidSecretSources = map[string]bool{
|
||||
"file": true, "keychain": true,
|
||||
}
|
||||
|
||||
func isValidSource(source string) bool {
|
||||
return ValidSecretSources[source]
|
||||
}
|
||||
|
||||
// secretAccountKey generates the keychain account key for an app's secret.
|
||||
func secretAccountKey(clientID string) string {
|
||||
return secretKeyPrefix + clientID
|
||||
}
|
||||
|
||||
// ResolveSecret resolves a SecretInput to a plain string.
|
||||
// SecretRef objects are resolved by source (file / keychain).
|
||||
func ResolveSecret(input SecretInput) (string, error) {
|
||||
if input.Ref == nil {
|
||||
return input.Plain, nil
|
||||
}
|
||||
switch input.Ref.Source {
|
||||
case "file":
|
||||
data, err := os.ReadFile(input.Ref.ID)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read secret file %s: %w", input.Ref.ID, err)
|
||||
}
|
||||
return strings.TrimSpace(string(data)), nil
|
||||
case "keychain":
|
||||
val, err := keychain.Get(keychain.Service, input.Ref.ID)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get secret from keychain: %w", err)
|
||||
}
|
||||
return val, nil
|
||||
default:
|
||||
return "", fmt.Errorf("unknown secret source: %s", input.Ref.Source)
|
||||
}
|
||||
}
|
||||
|
||||
// StoreSecret stores a plain text secret in keychain and returns a SecretRef.
|
||||
// If the input is already a SecretRef, it is returned as-is.
|
||||
// Returns error if keychain is unavailable.
|
||||
func StoreSecret(clientID string, input SecretInput) (SecretInput, error) {
|
||||
if !input.IsPlain() {
|
||||
return input, nil // SecretRef → keep as-is
|
||||
}
|
||||
key := secretAccountKey(clientID)
|
||||
if err := keychain.Set(keychain.Service, key, input.Plain); err != nil {
|
||||
return SecretInput{}, fmt.Errorf("keychain unavailable: %w\nhint: use file reference in config to bypass keychain", err)
|
||||
}
|
||||
return SecretInput{Ref: &SecretRef{Source: "keychain", ID: key}}, nil
|
||||
}
|
||||
|
||||
// RemoveSecretStore cleans up keychain entries when an app is removed.
|
||||
// Errors are intentionally ignored — cleanup is best-effort.
|
||||
func RemoveSecretStore(input SecretInput) {
|
||||
if input.IsSecretRef() && input.Ref.Source == "keychain" {
|
||||
_ = keychain.Remove(keychain.Service, input.Ref.ID)
|
||||
}
|
||||
}
|
||||
@@ -21,8 +21,8 @@ import (
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/security"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
const secureDataFile = ".data"
|
||||
|
||||
+117
-18
@@ -14,11 +14,17 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// TokenData holds the OAuth token set persisted to disk.
|
||||
@@ -32,6 +38,7 @@ type TokenData struct {
|
||||
UserID string `json:"user_id,omitempty"`
|
||||
UserName string `json:"user_name,omitempty"`
|
||||
CorpName string `json:"corp_name,omitempty"`
|
||||
ClientID string `json:"client_id,omitempty"` // Associated app client ID for refresh
|
||||
UpdatedAt string `json:"updated_at,omitempty"`
|
||||
Source string `json:"source,omitempty"`
|
||||
}
|
||||
@@ -58,54 +65,104 @@ func (t *TokenData) HasPersistentCode() bool {
|
||||
return t != nil && t.PersistentCode != ""
|
||||
}
|
||||
|
||||
// SaveTokenData saves TokenData to the platform keychain.
|
||||
// Uses the new keychain-based storage with random master key for better security.
|
||||
const tokenJSONFile = "token.json"
|
||||
|
||||
// TokenMarker is a lightweight file the host application reads to detect
|
||||
// whether the CLI has a valid token without accessing the keychain.
|
||||
type TokenMarker struct {
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
}
|
||||
|
||||
// WriteTokenMarker writes a token.json marker containing only an updated_at
|
||||
// timestamp. The host application uses this file's presence and mtime to
|
||||
// decide whether it needs to trigger a new auth exchange.
|
||||
func WriteTokenMarker(configDir string) error {
|
||||
marker := TokenMarker{UpdatedAt: time.Now().Format(time.RFC3339)}
|
||||
data, _ := json.MarshalIndent(marker, "", " ")
|
||||
if err := os.MkdirAll(configDir, 0o700); err != nil {
|
||||
return err
|
||||
}
|
||||
tmp := filepath.Join(configDir, tokenJSONFile+".tmp")
|
||||
if err := os.WriteFile(tmp, data, 0o600); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.Rename(tmp, filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// DeleteTokenMarker removes the token.json marker file.
|
||||
func DeleteTokenMarker(configDir string) error {
|
||||
return os.Remove(filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to the default keychain-based storage.
|
||||
func SaveTokenData(configDir string, data *TokenData) error {
|
||||
if h := edition.Get(); h.SaveToken != nil {
|
||||
jsonData, err := json.MarshalIndent(data, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling token data for hook: %w", err)
|
||||
}
|
||||
return h.SaveToken(configDir, jsonData)
|
||||
}
|
||||
return SaveTokenDataKeychain(data)
|
||||
}
|
||||
|
||||
// LoadTokenData reads TokenData from the platform keychain.
|
||||
// On first call, it attempts to migrate legacy .data file if present.
|
||||
// LoadTokenData reads TokenData. When an edition hook (LoadToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain with legacy .data migration.
|
||||
func LoadTokenData(configDir string) (*TokenData, error) {
|
||||
// Try loading from new keychain first
|
||||
if h := edition.Get(); h.LoadToken != nil {
|
||||
jsonData, err := h.LoadToken(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var td TokenData
|
||||
if err := json.Unmarshal(jsonData, &td); err != nil {
|
||||
return nil, fmt.Errorf("parsing token data from hook: %w", err)
|
||||
}
|
||||
return &td, nil
|
||||
}
|
||||
|
||||
// Default: keychain with legacy .data migration
|
||||
if TokenDataExistsKeychain() {
|
||||
return LoadTokenDataKeychain()
|
||||
}
|
||||
|
||||
// Fallback: try legacy .data file and migrate
|
||||
data, err := LoadSecureTokenData(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Migrate to keychain for future use
|
||||
if err := SaveTokenDataKeychain(data); err == nil {
|
||||
// Successfully migrated, delete legacy file
|
||||
_ = DeleteSecureData(configDir)
|
||||
}
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes token data from both keychain and legacy storage.
|
||||
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain + legacy cleanup.
|
||||
func DeleteTokenData(configDir string) error {
|
||||
// Delete from keychain
|
||||
if h := edition.Get(); h.DeleteToken != nil {
|
||||
return h.DeleteToken(configDir)
|
||||
}
|
||||
keychainErr := DeleteTokenDataKeychain()
|
||||
|
||||
// Also clean up any legacy .data file
|
||||
legacyErr := DeleteSecureData(configDir)
|
||||
|
||||
// Return keychain error if any, otherwise legacy error
|
||||
if keychainErr != nil {
|
||||
return keychainErr
|
||||
}
|
||||
return legacyErr
|
||||
}
|
||||
|
||||
// RevokeTokenRemote calls the DingTalk logout endpoint to invalidate the access token.
|
||||
// RevokeTokenRemote calls the appropriate logout/revoke endpoint to invalidate the access token.
|
||||
// Uses MCP revoke endpoint when clientID is from MCP, otherwise uses DingTalk logout.
|
||||
// This should be called before deleting local token data.
|
||||
// The function is best-effort: errors are returned but callers may choose to ignore them.
|
||||
func RevokeTokenRemote(ctx context.Context) error {
|
||||
// Use MCP revoke endpoint when clientID is from MCP
|
||||
if IsClientIDFromMCP() {
|
||||
return revokeTokenViaMCP(ctx)
|
||||
}
|
||||
// Direct mode: use DingTalk logout endpoint
|
||||
logoutURL, err := url.Parse(LogoutURL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parsing logout URL: %w", err)
|
||||
@@ -142,3 +199,45 @@ func RevokeTokenRemote(ctx context.Context) error {
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// revokeTokenViaMCP revokes token via MCP endpoint.
|
||||
func revokeTokenViaMCP(ctx context.Context) error {
|
||||
revokeURL := GetRevokeTokenURL()
|
||||
if revokeURL == "" {
|
||||
return nil // No revoke endpoint available
|
||||
}
|
||||
|
||||
// Load current token to get accessToken
|
||||
tokenData, err := LoadTokenData(getDefaultConfigDir())
|
||||
if err != nil || tokenData == nil {
|
||||
return nil // No token to revoke
|
||||
}
|
||||
|
||||
body := map[string]string{
|
||||
"clientId": ClientID(),
|
||||
"accessToken": tokenData.AccessToken,
|
||||
}
|
||||
bodyBytes, err := json.Marshal(body)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshaling revoke request: %w", err)
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, revokeURL, bytes.NewReader(bodyBytes))
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating revoke request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return fmt.Errorf("calling revoke endpoint: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("revoke endpoint returned status %d", resp.StatusCode)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+91
-8
@@ -17,18 +17,90 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/ir"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CACHE_DIR",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "覆盖缓存目录",
|
||||
DefaultValue: "~/.dws/cache",
|
||||
Example: "/tmp/dws-cache",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CATALOG_FIXTURE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "使用本地 JSON 文件替代在线目录发现",
|
||||
Example: "/path/to/catalog.json",
|
||||
Hidden: true,
|
||||
})
|
||||
}
|
||||
|
||||
// CatalogDegradedReason identifies why catalog discovery returned empty.
|
||||
type CatalogDegradedReason string
|
||||
|
||||
const (
|
||||
DegradedUnauthenticated CatalogDegradedReason = "unauthenticated"
|
||||
DegradedMarketUnreachable CatalogDegradedReason = "market_unreachable"
|
||||
DegradedRuntimeAllFailed CatalogDegradedReason = "runtime_all_failed"
|
||||
)
|
||||
|
||||
// CatalogDegraded is returned by EnvironmentLoader.Load when discovery
|
||||
// fails for a diagnosable reason. Callers that need graceful degradation
|
||||
// (e.g. the runtime runner) can check errors.As and fall back to an
|
||||
// empty catalog; callers like the schema command can surface the hint.
|
||||
type CatalogDegraded struct {
|
||||
Reason CatalogDegradedReason
|
||||
Hint string
|
||||
ServerCount int // number of servers discovered (only set for runtime_all_failed)
|
||||
}
|
||||
|
||||
func (e *CatalogDegraded) Error() string { return string(e.Reason) + ": " + e.Hint }
|
||||
|
||||
func degradedHint(reason CatalogDegradedReason, serverCount int) string {
|
||||
embedded := edition.Get().IsEmbedded
|
||||
switch reason {
|
||||
case DegradedUnauthenticated:
|
||||
if embedded {
|
||||
return "未登录,请重新认证"
|
||||
}
|
||||
return "未登录,无法发现 MCP 服务。请先执行: dws auth login"
|
||||
case DegradedMarketUnreachable:
|
||||
if embedded {
|
||||
return "无法连接 MCP 市场,请检查网络"
|
||||
}
|
||||
return "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络"
|
||||
case DegradedRuntimeAllFailed:
|
||||
if embedded {
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试", serverCount)
|
||||
}
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试或执行: dws cache refresh", serverCount)
|
||||
default:
|
||||
return "MCP 服务发现失败"
|
||||
}
|
||||
}
|
||||
|
||||
func newCatalogDegraded(reason CatalogDegradedReason, serverCount int) *CatalogDegraded {
|
||||
return &CatalogDegraded{
|
||||
Reason: reason,
|
||||
Hint: degradedHint(reason, serverCount),
|
||||
ServerCount: serverCount,
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
CatalogFixtureEnv = "DWS_CATALOG_FIXTURE"
|
||||
CacheDirEnv = "DWS_CACHE_DIR"
|
||||
@@ -92,6 +164,10 @@ type EnvironmentLoader struct {
|
||||
// AuthTokenFunc returns an access token for MCP discovery requests
|
||||
// (initialize, tools/list). When nil, discovery runs without auth.
|
||||
AuthTokenFunc func(context.Context) string
|
||||
// LoggerFunc returns a structured logger for discovery diagnostics.
|
||||
// Called lazily because the file logger may not be initialized at
|
||||
// construction time (it's set up during PersistentPreRunE).
|
||||
LoggerFunc func() *slog.Logger
|
||||
}
|
||||
|
||||
type cachedCatalogState struct {
|
||||
@@ -123,17 +199,23 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
// Startup command construction should not block on synchronous discovery
|
||||
// just because the cache has aged past the short revalidation window.
|
||||
cached := l.loadFromCache(store)
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
|
||||
transportClient := transport.NewClient(nil)
|
||||
hasAuth := false
|
||||
if l.AuthTokenFunc != nil {
|
||||
if token := l.AuthTokenFunc(ctx); token != "" {
|
||||
transportClient = transportClient.WithAuth(token, nil)
|
||||
hasAuth = true
|
||||
}
|
||||
}
|
||||
|
||||
if !hasAuth {
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedUnauthenticated, 0)
|
||||
}
|
||||
|
||||
// Use a bounded context so discovery doesn't hang in test or CI environments.
|
||||
timeout := defaultDiscoveryTimeout
|
||||
if l.DiscoveryTimeout > 0 {
|
||||
@@ -147,14 +229,15 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
transportClient,
|
||||
store,
|
||||
)
|
||||
if l.LoggerFunc != nil {
|
||||
service.Logger = l.LoggerFunc()
|
||||
}
|
||||
response, err := service.MarketClient.FetchServers(discoverCtx, 200)
|
||||
if err != nil {
|
||||
// Graceful degradation: return empty catalog on discovery failure.
|
||||
// The runtime runner will fall back to EchoRunner for unknown products.
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedMarketUnreachable, 0)
|
||||
}
|
||||
|
||||
servers := market.NormalizeServers(response, "live_market")
|
||||
@@ -184,10 +267,10 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
|
||||
refreshed, failures := service.DiscoverAllRuntime(discoverCtx, toRefresh)
|
||||
if len(unchangedRuntime) == 0 && len(refreshed) == 0 && len(failures) > 0 {
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedRuntimeAllFailed, len(servers))
|
||||
}
|
||||
|
||||
refreshedByKey := make(map[string]discovery.RuntimeServer, len(refreshed))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -14,8 +14,10 @@
|
||||
package compat
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"sort"
|
||||
"strconv"
|
||||
"strings"
|
||||
@@ -24,11 +26,12 @@ import (
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/convert"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/convert"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
)
|
||||
|
||||
type ValueKind string
|
||||
@@ -137,13 +140,31 @@ func NewDirectCommand(route Route, runner executor.Runner) *cobra.Command {
|
||||
for key, value := range bindingParams {
|
||||
params[key] = value
|
||||
}
|
||||
|
||||
// Collect schema-derived flags (from buildFlagsFromDetailSchema)
|
||||
// that are not covered by explicit bindings.
|
||||
collectSchemaFlags(cmd, route.Bindings, params)
|
||||
|
||||
if route.Normalizer != nil {
|
||||
if err := route.Normalizer(cmd, params); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if blocked, _ := params["_blocked"].(bool); blocked {
|
||||
return nil
|
||||
// Interactive confirmation for destructive operations (consistent with Helper commands)
|
||||
fmt.Fprintln(cmd.ErrOrStderr(), "⚠️ This is a destructive operation.")
|
||||
fmt.Fprint(cmd.ErrOrStderr(), "Confirm? (yes/no): ")
|
||||
|
||||
reader := bufio.NewReader(os.Stdin)
|
||||
answer, _ := reader.ReadString('\n')
|
||||
answer = strings.TrimSpace(strings.ToLower(answer))
|
||||
|
||||
if answer != "yes" && answer != "y" {
|
||||
fmt.Fprintln(cmd.ErrOrStderr(), "Operation cancelled")
|
||||
return nil
|
||||
}
|
||||
// User confirmed, continue execution
|
||||
delete(params, "_blocked")
|
||||
}
|
||||
|
||||
invocation := executor.NewCompatibilityInvocation(
|
||||
@@ -231,6 +252,70 @@ func ApplyBindings(cmd *cobra.Command, bindings []FlagBinding) {
|
||||
_ = cmd.Flags().MarkHidden("params")
|
||||
}
|
||||
|
||||
// collectSchemaFlags picks up flags created by buildFlagsFromDetailSchema that
|
||||
// have no explicit FlagBinding. This bridges the gap for plugin-defined tools
|
||||
// whose parameters come from the MCP inputSchema rather than CLIToolOverride.Flags.
|
||||
func collectSchemaFlags(cmd *cobra.Command, bindings []FlagBinding, params map[string]any) {
|
||||
// Build a set of flag names already covered by bindings.
|
||||
bound := make(map[string]bool, len(bindings)*2)
|
||||
for _, b := range bindings {
|
||||
if n := strings.TrimSpace(b.FlagName); n != "" {
|
||||
bound[n] = true
|
||||
}
|
||||
if a := strings.TrimSpace(b.Alias); a != "" {
|
||||
bound[a] = true
|
||||
}
|
||||
}
|
||||
|
||||
// Reserved/internal flags that should never be forwarded as tool params.
|
||||
skip := map[string]bool{
|
||||
"json": true, "params": true, "help": true,
|
||||
"format": true, "fields": true, "jq": true,
|
||||
"debug": true, "verbose": true, "dry-run": true,
|
||||
"yes": true, "mock": true, "timeout": true,
|
||||
"client-id": true, "client-secret": true,
|
||||
}
|
||||
|
||||
cmd.Flags().Visit(func(f *pflag.Flag) {
|
||||
if bound[f.Name] || skip[f.Name] {
|
||||
return
|
||||
}
|
||||
// Convert flag name back to the original parameter name (kebab → snake/camel)
|
||||
// For simplicity, use the flag name as-is since MCP tools typically
|
||||
// use snake_case which maps to kebab-case flags.
|
||||
paramName := toOriginalParamName(f.Name)
|
||||
if _, exists := params[paramName]; exists {
|
||||
return // already set by --json/--params
|
||||
}
|
||||
|
||||
switch f.Value.Type() {
|
||||
case "int":
|
||||
if v, err := cmd.Flags().GetInt(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
case "bool":
|
||||
if v, err := cmd.Flags().GetBool(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
case "stringSlice":
|
||||
if v, err := cmd.Flags().GetStringSlice(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
default:
|
||||
if v, err := cmd.Flags().GetString(f.Name); err == nil {
|
||||
params[paramName] = v
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// toOriginalParamName converts a kebab-case flag name back to the original
|
||||
// MCP parameter name. Since toKebabCase converts both camelCase and snake_case
|
||||
// to kebab-case, we default to snake_case (the MCP convention).
|
||||
func toOriginalParamName(flagName string) string {
|
||||
return strings.ReplaceAll(flagName, "-", "_")
|
||||
}
|
||||
|
||||
func CollectBindings(cmd *cobra.Command, bindings []FlagBinding, existing map[string]any) (map[string]any, error) {
|
||||
if existing == nil {
|
||||
existing = map[string]any{}
|
||||
|
||||
@@ -146,3 +146,102 @@ func TestCollectBindingsParsesJSONFlagValue(t *testing.T) {
|
||||
t.Fatalf("config.options = %#v, want array of 1", config["options"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsPicksUpUnboundFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Simulate a plugin command with schema-generated flags but no bindings.
|
||||
cmd := &cobra.Command{Use: "greet"}
|
||||
cmd.Flags().String("name", "", "Name of person")
|
||||
cmd.Flags().String("language", "en", "Language")
|
||||
cmd.Flags().Int("count", 0, "Repeat count")
|
||||
cmd.Flags().Bool("loud", false, "Loud mode")
|
||||
cmd.Flags().String("json", "", "")
|
||||
cmd.Flags().String("params", "", "")
|
||||
|
||||
// User sets --name and --count but not --language
|
||||
_ = cmd.Flags().Set("name", "Alice")
|
||||
_ = cmd.Flags().Set("count", "3")
|
||||
_ = cmd.Flags().Set("loud", "true")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, nil, params)
|
||||
|
||||
if params["name"] != "Alice" {
|
||||
t.Errorf("name = %v, want Alice", params["name"])
|
||||
}
|
||||
if params["count"] != 3 {
|
||||
t.Errorf("count = %v, want 3", params["count"])
|
||||
}
|
||||
if params["loud"] != true {
|
||||
t.Errorf("loud = %v, want true", params["loud"])
|
||||
}
|
||||
// language was not set by user, should not appear
|
||||
if _, exists := params["language"]; exists {
|
||||
t.Errorf("language should not be in params (not set by user)")
|
||||
}
|
||||
// json/params are reserved, should not appear
|
||||
if _, exists := params["json"]; exists {
|
||||
t.Error("json should be skipped")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsSkipsBoundFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
bindings := []FlagBinding{
|
||||
{FlagName: "dept-id", Property: "deptId", Kind: ValueString},
|
||||
}
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
ApplyBindings(cmd, bindings)
|
||||
// Also add a schema-generated flag
|
||||
cmd.Flags().String("title", "", "Title")
|
||||
|
||||
_ = cmd.Flags().Set("dept-id", "D001")
|
||||
_ = cmd.Flags().Set("title", "Hello")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, bindings, params)
|
||||
|
||||
// dept-id is bound, should NOT be collected by collectSchemaFlags
|
||||
if _, exists := params["dept_id"]; exists {
|
||||
t.Error("dept-id should be skipped (already has binding)")
|
||||
}
|
||||
// title is unbound, should be collected
|
||||
if params["title"] != "Hello" {
|
||||
t.Errorf("title = %v, want Hello", params["title"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectSchemaFlagsSkipsGlobalFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("name", "", "Name")
|
||||
cmd.Flags().Bool("debug", false, "Debug")
|
||||
cmd.Flags().Bool("verbose", false, "Verbose")
|
||||
cmd.Flags().Bool("dry-run", false, "Dry run")
|
||||
cmd.Flags().String("format", "json", "Format")
|
||||
cmd.Flags().String("json", "", "")
|
||||
cmd.Flags().String("params", "", "")
|
||||
|
||||
_ = cmd.Flags().Set("name", "Bob")
|
||||
_ = cmd.Flags().Set("debug", "true")
|
||||
_ = cmd.Flags().Set("verbose", "true")
|
||||
_ = cmd.Flags().Set("dry-run", "true")
|
||||
_ = cmd.Flags().Set("format", "table")
|
||||
|
||||
params := make(map[string]any)
|
||||
collectSchemaFlags(cmd, nil, params)
|
||||
|
||||
if params["name"] != "Bob" {
|
||||
t.Errorf("name = %v, want Bob", params["name"])
|
||||
}
|
||||
// Global flags should be skipped
|
||||
for _, skip := range []string{"debug", "verbose", "dry_run", "format"} {
|
||||
if _, exists := params[skip]; exists {
|
||||
t.Errorf("%s should be skipped (global flag)", skip)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+116
-15
@@ -21,12 +21,30 @@ import (
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_TENANT",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "缓存分区的租户标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_AUTH_IDENTITY",
|
||||
Category: configmeta.CategorySecurity,
|
||||
Description: "缓存分区的认证身份标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
tenantEnv = "DWS_TENANT"
|
||||
authIdentityEnv = "DWS_AUTH_IDENTITY"
|
||||
@@ -152,29 +170,112 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
|
||||
}, nil
|
||||
}
|
||||
|
||||
const perServerDiscoveryTimeout = 5 * time.Second
|
||||
|
||||
func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.ServerDescriptor) ([]RuntimeServer, []RuntimeFailure) {
|
||||
results := make([]RuntimeServer, 0, len(servers))
|
||||
failures := make([]RuntimeFailure, 0)
|
||||
for _, server := range servers {
|
||||
if server.CLI.Skip {
|
||||
continue
|
||||
type discoveryResult struct {
|
||||
server RuntimeServer
|
||||
failure *RuntimeFailure
|
||||
}
|
||||
|
||||
filtered := make([]market.ServerDescriptor, 0, len(servers))
|
||||
for _, srv := range servers {
|
||||
if !srv.CLI.Skip {
|
||||
filtered = append(filtered, srv)
|
||||
}
|
||||
runtimeServer, err := s.DiscoverServerRuntime(ctx, server)
|
||||
if err != nil {
|
||||
if errors.Is(err, errCLIServerSkipped) {
|
||||
continue
|
||||
}
|
||||
if len(filtered) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
ch := make(chan discoveryResult, len(filtered))
|
||||
var wg sync.WaitGroup
|
||||
for _, srv := range filtered {
|
||||
wg.Add(1)
|
||||
go func(server market.ServerDescriptor) {
|
||||
defer wg.Done()
|
||||
serverCtx, cancel := context.WithTimeout(ctx, perServerDiscoveryTimeout)
|
||||
defer cancel()
|
||||
start := time.Now()
|
||||
rs, err := s.DiscoverServerRuntime(serverCtx, server)
|
||||
elapsed := time.Since(start)
|
||||
if err != nil {
|
||||
if errors.Is(err, errCLIServerSkipped) {
|
||||
return
|
||||
}
|
||||
if s.Logger != nil {
|
||||
s.Logger.Warn("server_discovery_failed",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
|
||||
slog.String("error", err.Error()),
|
||||
slog.Bool("is_timeout", errors.Is(err, context.DeadlineExceeded)),
|
||||
)
|
||||
}
|
||||
// Per-server sub-context timed out but parent is still alive:
|
||||
// try cache fallback instead of reporting a hard failure.
|
||||
if errors.Is(err, context.DeadlineExceeded) && ctx.Err() == nil {
|
||||
if cached, cacheErr := s.loadServerFromCache(server); cacheErr == nil {
|
||||
if s.Logger != nil {
|
||||
s.Logger.Info("server_discovery_cache_fallback",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("source", cached.Source),
|
||||
)
|
||||
}
|
||||
ch <- discoveryResult{server: cached}
|
||||
return
|
||||
}
|
||||
}
|
||||
ch <- discoveryResult{failure: &RuntimeFailure{ServerKey: server.Key, Err: err}}
|
||||
return
|
||||
}
|
||||
failures = append(failures, RuntimeFailure{
|
||||
ServerKey: server.Key,
|
||||
Err: err,
|
||||
})
|
||||
continue
|
||||
if s.Logger != nil {
|
||||
s.Logger.Debug("server_discovery_ok",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
|
||||
slog.String("source", rs.Source),
|
||||
)
|
||||
}
|
||||
ch <- discoveryResult{server: rs}
|
||||
}(srv)
|
||||
}
|
||||
go func() {
|
||||
wg.Wait()
|
||||
close(ch)
|
||||
}()
|
||||
|
||||
results := make([]RuntimeServer, 0, len(filtered))
|
||||
failures := make([]RuntimeFailure, 0)
|
||||
for dr := range ch {
|
||||
if dr.failure != nil {
|
||||
failures = append(failures, *dr.failure)
|
||||
} else {
|
||||
results = append(results, dr.server)
|
||||
}
|
||||
results = append(results, runtimeServer)
|
||||
}
|
||||
return results, failures
|
||||
}
|
||||
|
||||
// loadServerFromCache tries to load a server's tools from cache, returning a
|
||||
// degraded RuntimeServer. Used as fallback when a per-server discovery timeout
|
||||
// fires but the parent context is still alive.
|
||||
func (s *Service) loadServerFromCache(server market.ServerDescriptor) (RuntimeServer, error) {
|
||||
partition := s.partition()
|
||||
snapshot, freshness, err := s.Cache.LoadTools(partition, server.Key)
|
||||
if err != nil {
|
||||
return RuntimeServer{}, err
|
||||
}
|
||||
server.NegotiatedProtocolVersion = snapshot.ProtocolVersion
|
||||
server.Source = string(freshness) + "_cache"
|
||||
server.Degraded = true
|
||||
return RuntimeServer{
|
||||
Server: server,
|
||||
NegotiatedProtocolVersion: snapshot.ProtocolVersion,
|
||||
Tools: snapshot.Tools,
|
||||
Source: string(freshness) + "_cache",
|
||||
Degraded: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) DiscoverDetail(ctx context.Context, server market.ServerDescriptor) (market.DetailResponse, error) {
|
||||
partition := s.partition()
|
||||
var fetchErr error
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package errors
|
||||
|
||||
import "strings"
|
||||
|
||||
// ServerDiagnostics holds server-side diagnostic fields extracted from
|
||||
// MCP response bodies or HTTP response headers. Fields are populated
|
||||
// on a best-effort basis during error construction.
|
||||
type ServerDiagnostics struct {
|
||||
TraceID string `json:"trace_id,omitempty"`
|
||||
ServerErrorCode string `json:"server_error_code,omitempty"`
|
||||
TechnicalDetail string `json:"technical_detail,omitempty"`
|
||||
ServerRetryable *bool `json:"server_retryable,omitempty"`
|
||||
}
|
||||
|
||||
// IsEmpty returns true when no diagnostic field has been populated.
|
||||
func (d ServerDiagnostics) IsEmpty() bool {
|
||||
return d.TraceID == "" && d.ServerErrorCode == "" &&
|
||||
d.TechnicalDetail == "" && d.ServerRetryable == nil
|
||||
}
|
||||
|
||||
// WithServerDiag attaches server diagnostics to the error.
|
||||
func WithServerDiag(diag ServerDiagnostics) Option {
|
||||
if diag.IsEmpty() {
|
||||
return func(*Error) {}
|
||||
}
|
||||
return func(e *Error) {
|
||||
e.ServerDiag = diag
|
||||
// Override retryable if server explicitly specified.
|
||||
if diag.ServerRetryable != nil {
|
||||
e.Retryable = *diag.ServerRetryable
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// WithTraceID records the server-provided trace identifier.
|
||||
// Used when only the trace ID is available (e.g. from HTTP headers)
|
||||
// without a full ServerDiagnostics struct.
|
||||
func WithTraceID(id string) Option {
|
||||
id = strings.TrimSpace(id)
|
||||
if id == "" {
|
||||
return func(*Error) {}
|
||||
}
|
||||
return func(e *Error) {
|
||||
e.ServerDiag.TraceID = id
|
||||
}
|
||||
}
|
||||
+119
-36
@@ -14,13 +14,12 @@
|
||||
package errors
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"bytes"
|
||||
)
|
||||
|
||||
// Category represents a stable error class with a documented exit code.
|
||||
@@ -36,18 +35,19 @@ const (
|
||||
|
||||
// Error is the structured repository-local error model for the Go rewrite.
|
||||
type Error struct {
|
||||
Category Category
|
||||
Message string
|
||||
Operation string
|
||||
ServerKey string
|
||||
Retryable bool
|
||||
Reason string
|
||||
Hint string
|
||||
Actions []string
|
||||
Snapshot string
|
||||
RPCCode int `json:"rpc_code,omitempty"`
|
||||
RPCData json.RawMessage `json:"rpc_data,omitempty"`
|
||||
Cause error `json:"-"`
|
||||
Category Category
|
||||
Message string
|
||||
Operation string
|
||||
ServerKey string
|
||||
Retryable bool
|
||||
Reason string
|
||||
Hint string
|
||||
Actions []string
|
||||
Snapshot string
|
||||
RPCCode int `json:"rpc_code,omitempty"`
|
||||
RPCData json.RawMessage `json:"rpc_data,omitempty"`
|
||||
ServerDiag ServerDiagnostics `json:"-"`
|
||||
Cause error `json:"-"`
|
||||
}
|
||||
|
||||
func (e *Error) Error() string {
|
||||
@@ -198,12 +198,31 @@ func NewInternal(message string, opts ...Option) error {
|
||||
return newError(CategoryInternal, message, opts...)
|
||||
}
|
||||
|
||||
// ExitCoder is implemented by errors that provide their own exit code.
|
||||
// Edition-specific error types (e.g. PATError, CLIError) implement this
|
||||
// so the framework can resolve exit codes without importing edition packages.
|
||||
type ExitCoder interface {
|
||||
ExitCode() int
|
||||
}
|
||||
|
||||
// RawStderrError is implemented by errors that must output raw content
|
||||
// directly to stderr, bypassing all CLI formatting (e.g. "Error:" prefix).
|
||||
// PAT authorization errors use this to pass JSON through to the desktop runtime.
|
||||
type RawStderrError interface {
|
||||
error
|
||||
RawStderr() string
|
||||
}
|
||||
|
||||
// ExitCode maps any error to a stable exit code.
|
||||
func ExitCode(err error) int {
|
||||
var typed *Error
|
||||
if stderrors.As(err, &typed) {
|
||||
return typed.ExitCode()
|
||||
}
|
||||
var ec ExitCoder
|
||||
if stderrors.As(err, &ec) {
|
||||
return ec.ExitCode()
|
||||
}
|
||||
return 5
|
||||
}
|
||||
|
||||
@@ -247,6 +266,23 @@ func PrintJSON(w io.Writer, err error) error {
|
||||
errorPayload["rpc_data"] = parsed
|
||||
}
|
||||
}
|
||||
if !typed.ServerDiag.IsEmpty() {
|
||||
if typed.ServerDiag.TraceID != "" {
|
||||
errorPayload["trace_id"] = typed.ServerDiag.TraceID
|
||||
}
|
||||
if typed.ServerDiag.ServerErrorCode != "" {
|
||||
errorPayload["server_error_code"] = typed.ServerDiag.ServerErrorCode
|
||||
// Add user-friendly hint for specific server error codes
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
|
||||
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
|
||||
}
|
||||
}
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
errorPayload["technical_detail"] = typed.ServerDiag.TechnicalDetail
|
||||
}
|
||||
}
|
||||
if typed.Cause != nil {
|
||||
errorPayload["cause"] = typed.Cause.Error()
|
||||
}
|
||||
@@ -263,8 +299,25 @@ func PrintJSON(w io.Writer, err error) error {
|
||||
return writeErr
|
||||
}
|
||||
|
||||
// PrintHuman writes a concise human-readable error rendering.
|
||||
// Verbosity controls how much detail PrintHuman includes.
|
||||
type Verbosity int
|
||||
|
||||
const (
|
||||
// VerbosityNormal shows essential info: error, hint, actions, trace_id, server_code.
|
||||
VerbosityNormal Verbosity = 0
|
||||
// VerbosityVerbose adds technical_detail, snapshot, execution context.
|
||||
VerbosityVerbose Verbosity = 1
|
||||
// VerbosityDebug adds all internal diagnostics (category, operation, reason, rpc_code).
|
||||
VerbosityDebug Verbosity = 2
|
||||
)
|
||||
|
||||
// PrintHuman writes a concise human-readable error rendering at normal verbosity.
|
||||
func PrintHuman(w io.Writer, err error) error {
|
||||
return PrintHumanAt(w, err, VerbosityNormal)
|
||||
}
|
||||
|
||||
// PrintHumanAt writes a human-readable error rendering at the given verbosity level.
|
||||
func PrintHumanAt(w io.Writer, err error, v Verbosity) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -275,21 +328,23 @@ func PrintHuman(w io.Writer, err error) error {
|
||||
return writeErr
|
||||
}
|
||||
|
||||
// Line 1: Error summary
|
||||
lines := []string{
|
||||
fmt.Sprintf("Error: [%s] %s", strings.ToUpper(string(typed.Category)), typed.Message),
|
||||
}
|
||||
if typed.Reason != "" {
|
||||
lines = append(lines, fmt.Sprintf("Reason: %s", typed.Reason))
|
||||
}
|
||||
if typed.Operation != "" {
|
||||
lines = append(lines, fmt.Sprintf("Operation: %s", typed.Operation))
|
||||
}
|
||||
if typed.ServerKey != "" {
|
||||
lines = append(lines, fmt.Sprintf("Server: %s", typed.ServerKey))
|
||||
}
|
||||
|
||||
// Always shown: hint, actions, retryable
|
||||
if typed.Hint != "" {
|
||||
lines = append(lines, fmt.Sprintf("Hint: %s", typed.Hint))
|
||||
}
|
||||
|
||||
// Add user-friendly hint for specific server error codes
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
|
||||
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
|
||||
}
|
||||
|
||||
if len(typed.Actions) > 0 {
|
||||
for _, action := range typed.Actions {
|
||||
if strings.TrimSpace(action) == "" {
|
||||
@@ -298,22 +353,50 @@ func PrintHuman(w io.Writer, err error) error {
|
||||
lines = append(lines, fmt.Sprintf("Action: %s", action))
|
||||
}
|
||||
}
|
||||
if typed.Snapshot != "" {
|
||||
lines = append(lines, fmt.Sprintf("Snapshot: %s", typed.Snapshot))
|
||||
}
|
||||
if typed.RPCCode != 0 {
|
||||
lines = append(lines, fmt.Sprintf("RPC Code: %d", typed.RPCCode))
|
||||
}
|
||||
if len(typed.RPCData) > 0 {
|
||||
lines = append(lines, fmt.Sprintf("RPC Data: %s", string(typed.RPCData)))
|
||||
}
|
||||
if typed.Cause != nil {
|
||||
lines = append(lines, fmt.Sprintf("Cause: %s", typed.Cause.Error()))
|
||||
}
|
||||
if typed.Retryable {
|
||||
lines = append(lines, "Retryable: true")
|
||||
}
|
||||
|
||||
// Always shown when present: Trace ID, Server Code
|
||||
if typed.ServerDiag.TraceID != "" {
|
||||
lines = append(lines, fmt.Sprintf("Trace ID: %s", typed.ServerDiag.TraceID))
|
||||
}
|
||||
if typed.ServerDiag.ServerErrorCode != "" {
|
||||
lines = append(lines, fmt.Sprintf("Server Code: %s", typed.ServerDiag.ServerErrorCode))
|
||||
}
|
||||
|
||||
// Verbose+: technical detail, snapshot, reason, server key
|
||||
if v >= VerbosityVerbose {
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
lines = append(lines, fmt.Sprintf("Detail: %s", typed.ServerDiag.TechnicalDetail))
|
||||
}
|
||||
if typed.Reason != "" {
|
||||
lines = append(lines, fmt.Sprintf("Reason: %s", typed.Reason))
|
||||
}
|
||||
if typed.ServerKey != "" {
|
||||
lines = append(lines, fmt.Sprintf("Server: %s", typed.ServerKey))
|
||||
}
|
||||
if typed.Snapshot != "" {
|
||||
lines = append(lines, fmt.Sprintf("Snapshot: %s", typed.Snapshot))
|
||||
}
|
||||
if typed.Cause != nil {
|
||||
lines = append(lines, fmt.Sprintf("Cause: %s", typed.Cause.Error()))
|
||||
}
|
||||
}
|
||||
|
||||
// Debug: all internal diagnostics
|
||||
if v >= VerbosityDebug {
|
||||
if typed.Operation != "" {
|
||||
lines = append(lines, fmt.Sprintf("Operation: %s", typed.Operation))
|
||||
}
|
||||
if typed.RPCCode != 0 {
|
||||
lines = append(lines, fmt.Sprintf("RPC Code: %d", typed.RPCCode))
|
||||
}
|
||||
if len(typed.RPCData) > 0 {
|
||||
lines = append(lines, fmt.Sprintf("RPC Data: %s", string(typed.RPCData)))
|
||||
}
|
||||
}
|
||||
|
||||
_, writeErr := fmt.Fprintln(w, strings.Join(lines, "\n"))
|
||||
return writeErr
|
||||
}
|
||||
|
||||
@@ -94,18 +94,29 @@ func TestPrintJSON_AllFields(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHuman_WithCause(t *testing.T) {
|
||||
func TestPrintHuman_WithCause_Verbose(t *testing.T) {
|
||||
t.Parallel()
|
||||
cause := fmt.Errorf("timeout")
|
||||
e := &Error{Category: CategoryDiscovery, Message: "discovery failed", Cause: cause}
|
||||
var buf bytes.Buffer
|
||||
PrintHumanAt(&buf, e, VerbosityVerbose)
|
||||
if !strings.Contains(buf.String(), "timeout") {
|
||||
t.Fatalf("expected cause in verbose human output: %s", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHuman_WithCause_NormalHidesCause(t *testing.T) {
|
||||
t.Parallel()
|
||||
cause := fmt.Errorf("timeout")
|
||||
e := &Error{Category: CategoryDiscovery, Message: "discovery failed", Cause: cause}
|
||||
var buf bytes.Buffer
|
||||
PrintHuman(&buf, e)
|
||||
if !strings.Contains(buf.String(), "timeout") {
|
||||
t.Fatalf("expected cause in human output: %s", buf.String())
|
||||
if strings.Contains(buf.String(), "Cause:") {
|
||||
t.Fatalf("normal mode should not show Cause: %s", buf.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHuman_AllFields(t *testing.T) {
|
||||
func TestPrintHuman_AllFields_Debug(t *testing.T) {
|
||||
t.Parallel()
|
||||
e := NewAPI("api error",
|
||||
WithOperation("initialize"),
|
||||
@@ -120,11 +131,11 @@ func TestPrintHuman_AllFields(t *testing.T) {
|
||||
WithCause(fmt.Errorf("network")),
|
||||
)
|
||||
var buf bytes.Buffer
|
||||
PrintHuman(&buf, e)
|
||||
PrintHumanAt(&buf, e, VerbosityDebug)
|
||||
out := buf.String()
|
||||
for _, expected := range []string{"API", "initialize", "connection_refused", "doc", "check network", "run again", "snap", "-32601", "network", "Retryable"} {
|
||||
if !strings.Contains(out, expected) {
|
||||
t.Fatalf("missing %q in human output: %s", expected, out)
|
||||
t.Fatalf("missing %q in debug human output: %s", expected, out)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -81,7 +81,7 @@ func TestPrintHuman(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var b strings.Builder
|
||||
if err := PrintHuman(&b, NewValidation(
|
||||
if err := PrintHumanAt(&b, NewValidation(
|
||||
"bad flag",
|
||||
WithReason("missing_required_flag"),
|
||||
WithOperation("calendar.list"),
|
||||
@@ -90,7 +90,7 @@ func TestPrintHuman(t *testing.T) {
|
||||
WithRetryable(true),
|
||||
WithActions("retry command"),
|
||||
WithSnapshot("/tmp/dws-recovery/snapshot.json"),
|
||||
)); err != nil {
|
||||
), VerbosityVerbose); err != nil {
|
||||
t.Fatalf("PrintHuman() error = %v", err)
|
||||
}
|
||||
|
||||
@@ -108,13 +108,64 @@ func TestPrintHuman(t *testing.T) {
|
||||
t.Fatalf("expected action in output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "Snapshot: /tmp/dws-recovery/snapshot.json") {
|
||||
t.Fatalf("expected snapshot in output, got %q", got)
|
||||
t.Fatalf("expected snapshot in verbose output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "Retryable: true") {
|
||||
t.Fatalf("expected retryable marker in output, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHuman_NormalMode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var b strings.Builder
|
||||
PrintHuman(&b, NewValidation(
|
||||
"bad flag",
|
||||
WithHint("fix it"),
|
||||
WithRetryable(true),
|
||||
WithActions("retry"),
|
||||
WithServerDiag(ServerDiagnostics{TraceID: "trace-abc", ServerErrorCode: "PARAM_ERROR"}),
|
||||
))
|
||||
|
||||
got := b.String()
|
||||
if !strings.Contains(got, "Error: [VALIDATION] bad flag") {
|
||||
t.Fatalf("expected header, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "Trace ID: trace-abc") {
|
||||
t.Fatalf("expected trace id in normal output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "Server Code: PARAM_ERROR") {
|
||||
t.Fatalf("expected server code in normal output, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintJSONIncludesServerDiag(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var b strings.Builder
|
||||
if err := PrintJSON(&b, NewAPI(
|
||||
"server error",
|
||||
WithServerDiag(ServerDiagnostics{
|
||||
TraceID: "trace-xyz",
|
||||
ServerErrorCode: "TIMEOUT_ERROR",
|
||||
TechnicalDetail: "deadline exceeded",
|
||||
}),
|
||||
)); err != nil {
|
||||
t.Fatalf("PrintJSON() error = %v", err)
|
||||
}
|
||||
|
||||
got := b.String()
|
||||
if !strings.Contains(got, `"trace_id": "trace-xyz"`) {
|
||||
t.Fatalf("expected trace_id in output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, `"server_error_code": "TIMEOUT_ERROR"`) {
|
||||
t.Fatalf("expected server_error_code in output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, `"technical_detail": "deadline exceeded"`) {
|
||||
t.Fatalf("expected technical_detail in output, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintJSONIncludesRPCCodeAndData(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -137,23 +188,38 @@ func TestPrintJSONIncludesRPCCodeAndData(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHumanIncludesRPCCode(t *testing.T) {
|
||||
func TestPrintHumanIncludesRPCCode_Debug(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var b strings.Builder
|
||||
if err := PrintHuman(&b, NewValidation(
|
||||
if err := PrintHumanAt(&b, NewValidation(
|
||||
"invalid params",
|
||||
WithRPCCode(-32602),
|
||||
WithRPCData([]byte(`"missing field"`)),
|
||||
)); err != nil {
|
||||
), VerbosityDebug); err != nil {
|
||||
t.Fatalf("PrintHuman() error = %v", err)
|
||||
}
|
||||
|
||||
got := b.String()
|
||||
if !strings.Contains(got, "RPC Code: -32602") {
|
||||
t.Fatalf("expected RPC Code in output, got %q", got)
|
||||
t.Fatalf("expected RPC Code in debug output, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "RPC Data:") {
|
||||
t.Fatalf("expected RPC Data in output, got %q", got)
|
||||
t.Fatalf("expected RPC Data in debug output, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintHumanHidesRPCCode_Normal(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var b strings.Builder
|
||||
PrintHuman(&b, NewValidation(
|
||||
"invalid params",
|
||||
WithRPCCode(-32602),
|
||||
))
|
||||
|
||||
got := b.String()
|
||||
if strings.Contains(got, "RPC Code:") {
|
||||
t.Fatalf("normal mode should not show RPC Code, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -0,0 +1,515 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
RegisterPublic(func() Handler {
|
||||
return reportHandler{}
|
||||
})
|
||||
}
|
||||
|
||||
type reportHandler struct{}
|
||||
|
||||
func (reportHandler) Name() string {
|
||||
return "report"
|
||||
}
|
||||
|
||||
func (reportHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
root := &cobra.Command{
|
||||
Use: "report",
|
||||
Aliases: []string{"log"},
|
||||
Short: "日志 / 模版 / 统计",
|
||||
Long: `钉钉日志:模版、创建、详情、列表、统计。
|
||||
|
||||
子命令:
|
||||
template 日志模版(list / detail)
|
||||
create 创建日志
|
||||
detail 获取日志详情
|
||||
list 查询收到的日志列表
|
||||
stats 获取日志统计数据
|
||||
sent 查询已发送的日志列表`,
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
template := &cobra.Command{
|
||||
Use: "template",
|
||||
Short: "日志模版",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
template.AddCommand(
|
||||
newReportTemplateListCommand(runner),
|
||||
newReportTemplateDetailCommand(runner),
|
||||
)
|
||||
|
||||
root.AddCommand(
|
||||
template,
|
||||
newReportCreateCommand(runner),
|
||||
newReportDetailCommand(runner),
|
||||
newReportListCommand(runner),
|
||||
newReportStatsCommand(runner),
|
||||
newReportSentCommand(runner),
|
||||
)
|
||||
return root
|
||||
}
|
||||
|
||||
// ── flexTimeLayouts: supported date formats, most specific first ──
|
||||
|
||||
var flexTimeLayouts = []string{
|
||||
time.RFC3339, // 2006-01-02T15:04:05+08:00
|
||||
"2006-01-02T15:04:05Z", // UTC Z suffix
|
||||
"2006-01-02T15:04:05-07:00", // with offset but no colon
|
||||
"2006-01-02T15:04:05", // no timezone
|
||||
"2006-01-02 15:04:05", // space-separated
|
||||
"2006-01-02T15:04", // no seconds
|
||||
"2006-01-02 15:04", // no seconds, space
|
||||
"2006-01-02", // date only
|
||||
"2006/01/02 15:04:05", // slash + time
|
||||
"2006/01/02", // slash date
|
||||
"20060102", // compact YYYYMMDD
|
||||
}
|
||||
|
||||
// parseFlexTimeToMillis parses a date string using multiple formats and returns Unix milliseconds.
|
||||
// Supports 11 formats for maximum compatibility with user input.
|
||||
func parseFlexTimeToMillis(flagName, value string) (int64, error) {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
return 0, apperrors.NewValidation(fmt.Sprintf(
|
||||
"--%s is required\n hint: example: 2026-03-10T14:00:00+08:00", flagName))
|
||||
}
|
||||
loc, _ := time.LoadLocation("Asia/Shanghai")
|
||||
if loc == nil {
|
||||
loc = time.Local
|
||||
}
|
||||
for _, layout := range flexTimeLayouts {
|
||||
t, err := time.ParseInLocation(layout, value, loc)
|
||||
if err == nil {
|
||||
return t.UnixMilli(), nil
|
||||
}
|
||||
}
|
||||
return 0, apperrors.NewValidation(fmt.Sprintf(
|
||||
"cannot parse time for --%s (input: %q)\n hint: supported formats: 2026-03-23T14:00:00+08:00, 2026-03-23 14:00:00, 2026-03-23",
|
||||
flagName, value))
|
||||
}
|
||||
|
||||
// validateTimeRange checks that endMs is strictly after startMs.
|
||||
func validateTimeRange(startMs, endMs int64) error {
|
||||
if endMs <= startMs {
|
||||
return apperrors.NewValidation("--end must be after --start\n hint: swap the values or adjust the time range")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── template list ──────────────────────────────────────────
|
||||
|
||||
func newReportTemplateListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "获取当前用户可用的日志模版列表",
|
||||
Example: " dws report template list",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
params := map[string]any{}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── template detail ────────────────────────────────────────
|
||||
|
||||
func newReportTemplateDetailCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "detail",
|
||||
Short: "获取日志模版详情",
|
||||
Example: " dws report template detail --name <templateName>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name, _ := cmd.Flags().GetString("name")
|
||||
if name == "" {
|
||||
return apperrors.NewValidation("--name is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_template_name": name,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("name", "", "模版名称 (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── create ─────────────────────────────────────────────────
|
||||
|
||||
func newReportCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create",
|
||||
Short: "创建日志",
|
||||
Long: `按模版创建一条日志。--contents 为 JSON 数组,每项需含 key、sort、content、contentType、type,
|
||||
与远程 create_report 一致;可先通过 report template list / template detail 取得 templateId 与控件定义。`,
|
||||
Example: ` dws report create --template-id TPL_ID --contents '[{"content":"完成开发","sort":"0","key":"今日完成","contentType":"markdown","type":"1"}]'
|
||||
dws report create --template-id TPL_ID --contents '[...]' --to-chat --to-user-ids userId1,userId2`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
tplID, _ := cmd.Flags().GetString("template-id")
|
||||
if tplID == "" {
|
||||
return apperrors.NewValidation("--template-id is required")
|
||||
}
|
||||
contentsJSON, _ := cmd.Flags().GetString("contents")
|
||||
if contentsJSON == "" {
|
||||
return apperrors.NewValidation("--contents is required")
|
||||
}
|
||||
var contents []map[string]any
|
||||
if err := json.Unmarshal([]byte(contentsJSON), &contents); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("--contents JSON parse failed: %v", err))
|
||||
}
|
||||
ddFrom, _ := cmd.Flags().GetString("dd-from")
|
||||
if ddFrom == "" {
|
||||
ddFrom = "dws"
|
||||
}
|
||||
toChat, _ := cmd.Flags().GetBool("to-chat")
|
||||
params := map[string]any{
|
||||
"templateId": tplID,
|
||||
"contents": contents,
|
||||
"ddFrom": ddFrom,
|
||||
"toChat": toChat,
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("to-user-ids"); v != "" {
|
||||
params["toUserIds"] = parseUserIDs(v)
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("template-id", "", "日志模版 ID (必填)")
|
||||
cmd.Flags().String("contents", "", "日志内容 JSON 数组 (必填),每项含 key/sort/content/contentType/type")
|
||||
cmd.Flags().String("dd-from", "dws", "创建来源标识")
|
||||
cmd.Flags().Bool("to-chat", false, "是否发送到日志接收人单聊")
|
||||
cmd.Flags().String("to-user-ids", "", "接收人 userId,逗号分隔 (可选)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── detail ─────────────────────────────────────────────────
|
||||
|
||||
func newReportDetailCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "detail",
|
||||
Short: "获取日志详情",
|
||||
Example: " dws report detail --report-id <reportId>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
reportID, _ := cmd.Flags().GetString("report-id")
|
||||
if reportID == "" {
|
||||
return apperrors.NewValidation("--report-id is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_id": reportID,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("report-id", "", "日志 ID (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── list (received reports) ────────────────────────────────
|
||||
// Key fix: cursor defaults to 0, size defaults to 20, flexible date parsing
|
||||
|
||||
func newReportListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "查询当前人收到的日志列表",
|
||||
Example: ` dws report list --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
|
||||
dws report list --start "2026-03-10 00:00:00" --end "2026-03-10 23:59:59" --cursor 0 --size 20`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
startStr, _ := cmd.Flags().GetString("start")
|
||||
endStr, _ := cmd.Flags().GetString("end")
|
||||
|
||||
startMs, err := parseFlexTimeToMillis("start", startStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
endMs, err := parseFlexTimeToMillis("end", endStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateTimeRange(startMs, endMs); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// cursor defaults to 0, size defaults to 20
|
||||
cursor, _ := cmd.Flags().GetInt("cursor")
|
||||
size, _ := cmd.Flags().GetInt("size")
|
||||
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
|
||||
size = v
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"startTime": float64(startMs),
|
||||
"endTime": float64(endMs),
|
||||
"cursor": float64(cursor),
|
||||
"size": float64(size),
|
||||
}
|
||||
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("start", "", "开始时间 ISO-8601 (如 2026-03-10T00:00:00+08:00) (必填)")
|
||||
cmd.Flags().String("end", "", "结束时间 ISO-8601 (如 2026-03-10T23:59:59+08:00) (必填)")
|
||||
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
|
||||
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
|
||||
cmd.Flags().Int("limit", 0, "--size 的别名")
|
||||
_ = cmd.Flags().MarkHidden("limit")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── stats ──────────────────────────────────────────────────
|
||||
|
||||
func newReportStatsCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "stats",
|
||||
Short: "获取日志统计数据",
|
||||
Example: " dws report stats --report-id <reportId>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
reportID, _ := cmd.Flags().GetString("report-id")
|
||||
if reportID == "" {
|
||||
return apperrors.NewValidation("--report-id is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_id": reportID,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("report-id", "", "日志 ID (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── sent (my created reports) ──────────────────────────────
|
||||
// Key fix: cursor defaults to 0, size defaults to 20, start/end default to last 30 days
|
||||
|
||||
func newReportSentCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "sent",
|
||||
Short: "查询当前人创建的日志列表",
|
||||
Example: ` dws report sent
|
||||
dws report sent --cursor 0 --size 20
|
||||
dws report sent --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
|
||||
dws report sent --template-name "日报"`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
// cursor defaults to 0, size defaults to 20
|
||||
cursor, _ := cmd.Flags().GetInt("cursor")
|
||||
size, _ := cmd.Flags().GetInt("size")
|
||||
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
|
||||
size = v
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"cursor": float64(cursor),
|
||||
"size": float64(size),
|
||||
}
|
||||
|
||||
// Default time range: last 30 days
|
||||
now := time.Now()
|
||||
startDefault := now.AddDate(0, 0, -30).Truncate(24 * time.Hour).Format(time.RFC3339)
|
||||
endDefault := time.Date(now.Year(), now.Month(), now.Day(), 23, 59, 59, 0, now.Location()).Format(time.RFC3339)
|
||||
|
||||
startStr, _ := cmd.Flags().GetString("start")
|
||||
if startStr == "" {
|
||||
startStr = startDefault
|
||||
}
|
||||
endStr, _ := cmd.Flags().GetString("end")
|
||||
if endStr == "" {
|
||||
endStr = endDefault
|
||||
}
|
||||
|
||||
startMs, err := parseFlexTimeToMillis("start", startStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["startTime"] = float64(startMs)
|
||||
|
||||
endMs, err := parseFlexTimeToMillis("end", endStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["endTime"] = float64(endMs)
|
||||
|
||||
if err := validateTimeRange(startMs, endMs); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Optional modified time filters
|
||||
if v, _ := cmd.Flags().GetString("modified-start"); v != "" {
|
||||
ms, err := parseFlexTimeToMillis("modified-start", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["modifiedStartTime"] = float64(ms)
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("modified-end"); v != "" {
|
||||
ms, err := parseFlexTimeToMillis("modified-end", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["modifiedEndTime"] = float64(ms)
|
||||
}
|
||||
|
||||
// Optional template name filter
|
||||
if v, _ := cmd.Flags().GetString("template-name"); v != "" {
|
||||
params["report_template_name"] = v
|
||||
}
|
||||
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
|
||||
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
|
||||
cmd.Flags().Int("limit", 0, "--size 的别名")
|
||||
_ = cmd.Flags().MarkHidden("limit")
|
||||
cmd.Flags().String("start", "", "创建开始时间 ISO-8601 (默认最近 30 天)")
|
||||
cmd.Flags().String("end", "", "创建结束时间 ISO-8601 (默认最近 30 天)")
|
||||
cmd.Flags().String("modified-start", "", "修改开始时间 ISO-8601 (可选)")
|
||||
cmd.Flags().String("modified-end", "", "修改结束时间 ISO-8601 (可选)")
|
||||
cmd.Flags().String("template-name", "", "日志模板名称 (可选,不传查全部)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── helpers ────────────────────────────────────────────────
|
||||
|
||||
func parseUserIDs(s string) []string {
|
||||
parts := strings.Split(s, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
out = append(out, p)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestParseFlexTimeToMillis_RFC3339(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-10T00:00:00+08:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_SpaceSeparated(t *testing.T) {
|
||||
t.Parallel()
|
||||
// This is the format that was causing the HTTP 400 error
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-01 00:00:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_DateOnly(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-01")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_NoTimezone(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-10T14:00:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_SlashFormat(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026/03/10")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_CompactFormat(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "20260310")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := parseFlexTimeToMillis("start", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_Invalid(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := parseFlexTimeToMillis("start", "not-a-date")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for invalid date")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Valid(t *testing.T) {
|
||||
t.Parallel()
|
||||
start := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
end := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(start, end); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Invalid(t *testing.T) {
|
||||
t.Parallel()
|
||||
start := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
end := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(start, end); err == nil {
|
||||
t.Fatal("expected error when end is before start")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Equal(t *testing.T) {
|
||||
t.Parallel()
|
||||
ts := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(ts, ts); err == nil {
|
||||
t.Fatal("expected error when start equals end")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseUserIDs(t *testing.T) {
|
||||
t.Parallel()
|
||||
tests := []struct {
|
||||
input string
|
||||
want int
|
||||
}{
|
||||
{"user1,user2,user3", 3},
|
||||
{"user1", 1},
|
||||
{"user1, user2, user3", 3},
|
||||
{"user1,,user2", 2},
|
||||
{"", 0},
|
||||
{" , , ", 0},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := parseUserIDs(tt.input)
|
||||
if len(got) != tt.want {
|
||||
t.Errorf("parseUserIDs(%q) = %d items, want %d", tt.input, len(got), tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,20 @@
|
||||
"返回数据缺少 uploadUrl 或 fileToken": "Response data missing uploadUrl or fileToken",
|
||||
"附件工作流": "Attachment workflow",
|
||||
"页码 (必填)": "page number (required)",
|
||||
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ Unable to verify CLI data access permission status",
|
||||
" 请检查网络连接后重试。": " Please check your network connection and retry.",
|
||||
"检查 CLI 授权状态失败": "Failed to check CLI auth status",
|
||||
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ CLI data access is not enabled for this organization",
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " The organization admin has not enabled \"Allow members to access their personal data via CLI\".",
|
||||
" 组织主管理员:": " Organization super admins: ",
|
||||
" 请联系组织主管理员开启后重新登录。": " Please contact the organization super admin to enable it and re-login.",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " Admin settings: https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "CLI data access is not enabled for this organization, please contact admin to enable it",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ CLI data access is not enabled for this organization, please submit an authorization request in the browser...",
|
||||
"✅ 权限已开启,继续登录...": "✅ Permission enabled, continuing login...",
|
||||
"等待管理员审批中": "Waiting for admin approval",
|
||||
"等待提交申请中": "Waiting to submit request",
|
||||
"操作超时,请重新登录": "Operation timed out, please re-login",
|
||||
"检查组织 CLI 授权状态...": "Checking organization CLI auth status...",
|
||||
"🔐 登录钉钉": "🔐 Login to DingTalk"
|
||||
}
|
||||
|
||||
@@ -155,5 +155,20 @@
|
||||
"返回数据缺少 uploadUrl 或 fileToken": "返回数据缺少 uploadUrl 或 fileToken",
|
||||
"附件工作流": "附件工作流",
|
||||
"页码 (必填)": "页码 (必填)",
|
||||
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ 无法检查 CLI 数据访问权限状态",
|
||||
" 请检查网络连接后重试。": " 请检查网络连接后重试。",
|
||||
"检查 CLI 授权状态失败": "检查 CLI 授权状态失败",
|
||||
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ 该组织尚未开启 CLI 数据访问权限",
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。",
|
||||
" 组织主管理员:": " 组织主管理员:",
|
||||
" 请联系组织主管理员开启后重新登录。": " 请联系组织主管理员开启后重新登录。",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "该组织尚未开启 CLI 数据访问权限,请联系管理员开启",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...",
|
||||
"✅ 权限已开启,继续登录...": "✅ 权限已开启,继续登录...",
|
||||
"等待管理员审批中": "等待管理员审批中",
|
||||
"等待提交申请中": "等待提交申请中",
|
||||
"操作超时,请重新登录": "操作超时,请重新登录",
|
||||
"检查组织 CLI 授权状态...": "检查组织 CLI 授权状态...",
|
||||
"🔐 登录钉钉": "🔐 登录钉钉"
|
||||
}
|
||||
|
||||
@@ -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 (
|
||||
|
||||
@@ -13,7 +13,14 @@
|
||||
|
||||
package logging
|
||||
|
||||
import "strings"
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
// sensitiveKeys are header/field names whose values must be redacted in logs.
|
||||
var sensitiveKeys = map[string]bool{
|
||||
@@ -25,12 +32,30 @@ var sensitiveKeys = map[string]bool{
|
||||
"secret": true,
|
||||
"password": true,
|
||||
"cookie": true,
|
||||
"api_key": true,
|
||||
"api-key": true,
|
||||
"access_token": true,
|
||||
"credential": true,
|
||||
}
|
||||
|
||||
// sensitiveSubstrings are substrings that mark a key as sensitive.
|
||||
var sensitiveSubstrings = []string{
|
||||
"password", "secret", "token", "credential",
|
||||
}
|
||||
|
||||
// IsSensitiveKey returns true if the key (case-insensitive) refers to a
|
||||
// credential or secret that must not appear in log files.
|
||||
func IsSensitiveKey(key string) bool {
|
||||
return sensitiveKeys[strings.ToLower(key)]
|
||||
lower := strings.ToLower(key)
|
||||
if sensitiveKeys[lower] {
|
||||
return true
|
||||
}
|
||||
for _, sub := range sensitiveSubstrings {
|
||||
if strings.Contains(lower, sub) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// RedactValue replaces a sensitive value with a safe placeholder.
|
||||
@@ -42,3 +67,65 @@ func RedactValue(value string) string {
|
||||
}
|
||||
return value[:4] + "***"
|
||||
}
|
||||
|
||||
// TruncateBody returns the body truncated to maxBytes with a UTF-8 safe
|
||||
// boundary. If truncated, appends a marker showing the original size.
|
||||
func TruncateBody(body []byte, maxBytes int) string {
|
||||
if len(body) <= maxBytes {
|
||||
return string(body)
|
||||
}
|
||||
safe := body[:maxBytes]
|
||||
// Walk back to a valid UTF-8 boundary.
|
||||
for len(safe) > 0 && !utf8.Valid(safe) {
|
||||
safe = safe[:len(safe)-1]
|
||||
}
|
||||
return fmt.Sprintf("%s...(truncated, total=%d bytes)", string(safe), len(body))
|
||||
}
|
||||
|
||||
// SanitizeArguments returns a JSON string of the arguments map with
|
||||
// sensitive-looking values replaced by "***". Truncates to maxBytes.
|
||||
func SanitizeArguments(args map[string]any, maxBytes int) string {
|
||||
if len(args) == 0 {
|
||||
return "{}"
|
||||
}
|
||||
sanitized := make(map[string]any, len(args))
|
||||
for k, v := range args {
|
||||
sanitized[k] = v
|
||||
}
|
||||
redactMapValues(sanitized)
|
||||
data, err := json.Marshal(sanitized)
|
||||
if err != nil {
|
||||
return "{}"
|
||||
}
|
||||
return TruncateBody(data, maxBytes)
|
||||
}
|
||||
|
||||
// redactMapValues replaces values of sensitive keys with "***" in-place.
|
||||
func redactMapValues(m map[string]any) {
|
||||
for k, v := range m {
|
||||
if IsSensitiveKey(k) {
|
||||
m[k] = "***"
|
||||
continue
|
||||
}
|
||||
if nested, ok := v.(map[string]any); ok {
|
||||
redactMapValues(nested)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RedactHeaders returns slog attributes for HTTP headers with sensitive
|
||||
// values redacted.
|
||||
func RedactHeaders(headers http.Header) []slog.Attr {
|
||||
if len(headers) == 0 {
|
||||
return nil
|
||||
}
|
||||
attrs := make([]slog.Attr, 0, len(headers))
|
||||
for key := range headers {
|
||||
value := headers.Get(key)
|
||||
if IsSensitiveKey(key) {
|
||||
value = RedactValue(value)
|
||||
}
|
||||
attrs = append(attrs, slog.String("header."+strings.ToLower(key), value))
|
||||
}
|
||||
return attrs
|
||||
}
|
||||
|
||||
@@ -13,7 +13,11 @@
|
||||
|
||||
package logging
|
||||
|
||||
import "testing"
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestIsSensitiveKey(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -69,3 +73,97 @@ func TestRedactValue(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
short := []byte("hello")
|
||||
if got := TruncateBody(short, 100); got != "hello" {
|
||||
t.Fatalf("expected no truncation, got %q", got)
|
||||
}
|
||||
long := []byte(strings.Repeat("a", 200))
|
||||
got := TruncateBody(long, 50)
|
||||
if !strings.Contains(got, "truncated") {
|
||||
t.Fatalf("expected truncation marker, got %q", got)
|
||||
}
|
||||
if !strings.Contains(got, "total=200") {
|
||||
t.Fatalf("expected total size, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateBody_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := TruncateBody(nil, 100); got != "" {
|
||||
t.Fatalf("expected empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeArguments(t *testing.T) {
|
||||
t.Parallel()
|
||||
args := map[string]any{
|
||||
"name": "test",
|
||||
"password": "secret123",
|
||||
"nested": map[string]any{
|
||||
"api_key": "key-value",
|
||||
"safe": "ok",
|
||||
},
|
||||
}
|
||||
got := SanitizeArguments(args, 4096)
|
||||
if strings.Contains(got, "secret123") {
|
||||
t.Fatalf("password should be redacted: %s", got)
|
||||
}
|
||||
if strings.Contains(got, "key-value") {
|
||||
t.Fatalf("api_key should be redacted: %s", got)
|
||||
}
|
||||
if !strings.Contains(got, "test") {
|
||||
t.Fatalf("non-sensitive value should remain: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSanitizeArguments_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := SanitizeArguments(nil, 100); got != "{}" {
|
||||
t.Fatalf("expected {}, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRedactHeaders(t *testing.T) {
|
||||
t.Parallel()
|
||||
headers := http.Header{
|
||||
"Authorization": {"Bearer token123456"},
|
||||
"Content-Type": {"application/json"},
|
||||
}
|
||||
attrs := RedactHeaders(headers)
|
||||
if len(attrs) != 2 {
|
||||
t.Fatalf("expected 2 attrs, got %d", len(attrs))
|
||||
}
|
||||
for _, attr := range attrs {
|
||||
if attr.Key == "header.authorization" && !strings.Contains(attr.Value.String(), "***") {
|
||||
t.Fatalf("authorization should be redacted: %s", attr.Value.String())
|
||||
}
|
||||
if attr.Key == "header.content-type" && attr.Value.String() != "application/json" {
|
||||
t.Fatalf("content-type should not be redacted: %s", attr.Value.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsSensitiveKey_Substrings(t *testing.T) {
|
||||
t.Parallel()
|
||||
tests := []struct {
|
||||
key string
|
||||
want bool
|
||||
}{
|
||||
{"x-api-token", true},
|
||||
{"user_password_hash", true},
|
||||
{"my_secret_key", true},
|
||||
{"x-credential-id", true},
|
||||
{"safe-header", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.key, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := IsSensitiveKey(tt.key); got != tt.want {
|
||||
t.Errorf("IsSensitiveKey(%q) = %v, want %v", tt.key, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,9 +16,17 @@ package logging
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"runtime"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxBodyLogSize is the maximum bytes of request/response body to log.
|
||||
maxBodyLogSize = 4096
|
||||
// maxArgLogSize is the maximum bytes for sanitized argument summaries.
|
||||
maxArgLogSize = 1024
|
||||
)
|
||||
|
||||
// LogRequest logs a JSON-RPC request at Debug level.
|
||||
func LogRequest(logger *slog.Logger, method, endpoint, executionId string, bodySize int) {
|
||||
if logger == nil {
|
||||
@@ -32,14 +40,28 @@ func LogRequest(logger *slog.Logger, method, endpoint, executionId string, bodyS
|
||||
)
|
||||
}
|
||||
|
||||
// LogResponse logs a JSON-RPC response at Debug level.
|
||||
func LogResponse(logger *slog.Logger, method, endpoint string, statusCode int, respSize int, duration time.Duration, err error) {
|
||||
// LogRequestBody logs a truncated, redacted request body for tools/call.
|
||||
func LogRequestBody(logger *slog.Logger, method, executionId string, toolName string, arguments map[string]any) {
|
||||
if logger == nil || method != "tools/call" {
|
||||
return
|
||||
}
|
||||
logger.Debug("jsonrpc_request_body",
|
||||
slog.String("method", method),
|
||||
slog.String("execution_id", executionId),
|
||||
slog.String("tool_name", toolName),
|
||||
slog.String("arguments_summary", SanitizeArguments(arguments, maxArgLogSize)),
|
||||
)
|
||||
}
|
||||
|
||||
// LogResponse logs a JSON-RPC response at Debug level (Warn on error).
|
||||
func LogResponse(logger *slog.Logger, method, endpoint, executionId string, statusCode int, respSize int, duration time.Duration, err error) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("method", method),
|
||||
slog.String("endpoint", redactEndpoint(endpoint)),
|
||||
slog.String("execution_id", executionId),
|
||||
slog.Int("status", statusCode),
|
||||
slog.Int("resp_size", respSize),
|
||||
slog.String("duration", duration.Truncate(time.Millisecond).String()),
|
||||
@@ -52,6 +74,110 @@ func LogResponse(logger *slog.Logger, method, endpoint string, statusCode int, r
|
||||
logger.LogAttrs(context.TODO(), slog.LevelDebug, "jsonrpc_response", attrs...)
|
||||
}
|
||||
|
||||
// LogResponseBody logs a truncated response body on error paths.
|
||||
func LogResponseBody(logger *slog.Logger, method, executionId string, statusCode int, body []byte, traceID string) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("method", method),
|
||||
slog.String("execution_id", executionId),
|
||||
slog.Int("status", statusCode),
|
||||
slog.String("body", TruncateBody(body, maxBodyLogSize)),
|
||||
}
|
||||
if traceID != "" {
|
||||
attrs = append(attrs, slog.String("trace_id", traceID))
|
||||
}
|
||||
level := slog.LevelDebug
|
||||
if statusCode >= 400 {
|
||||
level = slog.LevelWarn
|
||||
}
|
||||
logger.LogAttrs(context.TODO(), level, "jsonrpc_response_body", attrs...)
|
||||
}
|
||||
|
||||
// LogRetryAttempt logs a retry attempt at Warn level.
|
||||
func LogRetryAttempt(logger *slog.Logger, method, executionId string, attempt, maxRetries int, statusCode int, delay time.Duration, lastErr error) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("method", method),
|
||||
slog.String("execution_id", executionId),
|
||||
slog.Int("attempt", attempt+1),
|
||||
slog.Int("max_attempts", maxRetries+1),
|
||||
slog.Int("status", statusCode),
|
||||
slog.String("delay", delay.String()),
|
||||
}
|
||||
if lastErr != nil {
|
||||
attrs = append(attrs, slog.String("error", lastErr.Error()))
|
||||
}
|
||||
logger.LogAttrs(context.TODO(), slog.LevelWarn, "jsonrpc_retry", attrs...)
|
||||
}
|
||||
|
||||
// LogErrorClassified logs the final error classification at Warn level.
|
||||
func LogErrorClassified(logger *slog.Logger, method, executionId, category, reason string, httpStatus, rpcCode int, retryable bool, traceID string) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("method", method),
|
||||
slog.String("execution_id", executionId),
|
||||
slog.String("category", category),
|
||||
slog.String("reason", reason),
|
||||
slog.Bool("retryable", retryable),
|
||||
}
|
||||
if httpStatus != 0 {
|
||||
attrs = append(attrs, slog.Int("http_status", httpStatus))
|
||||
}
|
||||
if rpcCode != 0 {
|
||||
attrs = append(attrs, slog.Int("rpc_code", rpcCode))
|
||||
}
|
||||
if traceID != "" {
|
||||
attrs = append(attrs, slog.String("trace_id", traceID))
|
||||
}
|
||||
logger.LogAttrs(context.TODO(), slog.LevelWarn, "error_classified", attrs...)
|
||||
}
|
||||
|
||||
// LogCommandStart logs the beginning of a command execution.
|
||||
func LogCommandStart(logger *slog.Logger, executionId, product, tool, endpoint, version string, authPresent bool, timeoutSec int) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("execution_id", executionId),
|
||||
slog.String("product", product),
|
||||
slog.String("tool", tool),
|
||||
slog.String("endpoint", redactEndpoint(endpoint)),
|
||||
slog.String("cli_version", version),
|
||||
slog.String("os", runtime.GOOS),
|
||||
slog.String("arch", runtime.GOARCH),
|
||||
slog.Bool("auth_token_present", authPresent),
|
||||
}
|
||||
if timeoutSec > 0 {
|
||||
attrs = append(attrs, slog.Int("timeout_sec", timeoutSec))
|
||||
}
|
||||
logger.LogAttrs(context.TODO(), slog.LevelInfo, "command_start", attrs...)
|
||||
}
|
||||
|
||||
// LogCommandEnd logs the end of a command execution.
|
||||
func LogCommandEnd(logger *slog.Logger, executionId, product, tool string, success bool, duration time.Duration, errCategory, errReason string) {
|
||||
if logger == nil {
|
||||
return
|
||||
}
|
||||
attrs := []slog.Attr{
|
||||
slog.String("execution_id", executionId),
|
||||
slog.String("product", product),
|
||||
slog.String("tool", tool),
|
||||
slog.Bool("success", success),
|
||||
slog.String("duration", duration.Truncate(time.Millisecond).String()),
|
||||
}
|
||||
if !success {
|
||||
attrs = append(attrs, slog.String("error_category", errCategory))
|
||||
attrs = append(attrs, slog.String("error_reason", errReason))
|
||||
}
|
||||
logger.LogAttrs(context.TODO(), slog.LevelInfo, "command_end", attrs...)
|
||||
}
|
||||
|
||||
// redactEndpoint removes query parameters from endpoint URLs in logs.
|
||||
func redactEndpoint(endpoint string) string {
|
||||
for i := 0; i < len(endpoint); i++ {
|
||||
|
||||
@@ -52,7 +52,7 @@ func TestLogResponseSuccess(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", 200, 1024, 150*time.Millisecond, nil)
|
||||
LogResponse(logger, "tools/call", "https://mcp.dingtalk.com/api", "exec-1", 200, 1024, 150*time.Millisecond, nil)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "jsonrpc_response") {
|
||||
@@ -72,7 +72,7 @@ func TestLogResponseError(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", 500, 0, 2*time.Second, errors.New("connection refused"))
|
||||
LogResponse(logger, "initialize", "https://mcp.dingtalk.com", "exec-2", 500, 0, 2*time.Second, errors.New("connection refused"))
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "WARN") {
|
||||
@@ -87,7 +87,95 @@ func TestLogRequestNilLogger(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Should not panic
|
||||
LogRequest(nil, "test", "http://localhost", "", 0)
|
||||
LogResponse(nil, "test", "http://localhost", 200, 0, 0, nil)
|
||||
LogResponse(nil, "test", "http://localhost", "", 200, 0, 0, nil)
|
||||
LogRequestBody(nil, "tools/call", "exec-1", "tool", nil)
|
||||
LogResponseBody(nil, "tools/call", "exec-1", 200, nil, "")
|
||||
LogRetryAttempt(nil, "tools/call", "exec-1", 0, 2, 429, 0, nil)
|
||||
LogErrorClassified(nil, "tools/call", "exec-1", "api", "timeout", 0, 0, true, "")
|
||||
LogCommandStart(nil, "exec-1", "doc", "list", "https://mcp.example.com", "1.0.0", false, 0)
|
||||
LogCommandEnd(nil, "exec-1", "doc", "list", true, 0, "", "")
|
||||
}
|
||||
|
||||
func TestLogRequestBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
args := map[string]any{"name": "test", "limit": 10}
|
||||
LogRequestBody(logger, "tools/call", "exec-1", "doc.list", args)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "jsonrpc_request_body") {
|
||||
t.Error("missing message")
|
||||
}
|
||||
if !strings.Contains(out, "doc.list") {
|
||||
t.Error("missing tool_name")
|
||||
}
|
||||
if !strings.Contains(out, "exec-1") {
|
||||
t.Error("missing execution_id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogRequestBody_SkipsNonToolsCall(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogRequestBody(logger, "initialize", "exec-1", "", nil)
|
||||
if buf.Len() != 0 {
|
||||
t.Error("should not log body for non-tools/call methods")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogResponseBody(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogResponseBody(logger, "tools/call", "exec-1", 500, []byte(`{"error":"fail"}`), "trace-abc")
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "jsonrpc_response_body") {
|
||||
t.Error("missing message")
|
||||
}
|
||||
if !strings.Contains(out, "trace-abc") {
|
||||
t.Error("missing trace_id")
|
||||
}
|
||||
if !strings.Contains(out, "WARN") {
|
||||
t.Error("expected WARN level for 500 status")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogRetryAttempt(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogRetryAttempt(logger, "tools/call", "exec-1", 0, 2, 429, 10*time.Millisecond, errors.New("rate limited"))
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "jsonrpc_retry") {
|
||||
t.Error("missing message")
|
||||
}
|
||||
if !strings.Contains(out, "rate limited") {
|
||||
t.Error("missing error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogErrorClassified(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf bytes.Buffer
|
||||
logger := slog.New(slog.NewJSONHandler(&buf, &slog.HandlerOptions{Level: slog.LevelDebug}))
|
||||
|
||||
LogErrorClassified(logger, "tools/call", "exec-1", "auth", "http_401", 401, 0, false, "trace-xyz")
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "error_classified") {
|
||||
t.Error("missing message")
|
||||
}
|
||||
if !strings.Contains(out, "trace-xyz") {
|
||||
t.Error("missing trace_id")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRedactEndpoint(t *testing.T) {
|
||||
|
||||
+19
-17
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user