Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
af000a8dfa | ||
|
|
f19a3ccfa5 | ||
|
|
89c5038446 | ||
|
|
4f915e4e2c | ||
|
|
ec03b7cca3 | ||
|
|
e0544579d2 | ||
|
|
ce43280c11 | ||
|
|
74ca40c197 | ||
|
|
cfaa673863 | ||
|
|
3bc6c31a2d | ||
|
|
383aeefaf6 | ||
|
|
a5bede3a19 | ||
|
|
bbf66e23d6 | ||
|
|
5e168c92cf | ||
|
|
725577103d | ||
|
|
f762117d4e | ||
|
|
750b6c04d6 | ||
|
|
59e51c348a | ||
|
|
a056a9abfb | ||
|
|
33ae780103 | ||
|
|
daf56514f7 | ||
|
|
8bcbceb971 | ||
|
|
df01f36442 | ||
|
|
b0024aa669 | ||
|
|
7b7aeadbbe | ||
|
|
94ad422a9f | ||
|
|
4f1ee37508 | ||
|
|
fec0347cd6 | ||
|
|
93318f4a83 | ||
|
|
a14fd0250c | ||
|
|
c99e228669 | ||
|
|
95d495f290 | ||
|
|
4bf300d862 | ||
|
|
1a1fc531f5 | ||
|
|
9fc570607f | ||
|
|
4851d19141 | ||
|
|
42fb25d150 | ||
|
|
416ad6571d | ||
|
|
9a119fbd64 | ||
|
|
c5decb2f90 | ||
|
|
fae2a4f5f0 | ||
|
|
d25b106e4f | ||
|
|
9f78e51ae7 | ||
|
|
d2752d8b5b | ||
|
|
8ecbff391c | ||
|
|
d259864a2b | ||
|
|
408098bdc1 | ||
|
|
658ec1676c | ||
|
|
e36d3b3474 | ||
|
|
0b9952c58d | ||
|
|
56af1ea091 | ||
|
|
ea5859b92b | ||
|
|
19f2ed5c69 | ||
|
|
efbaf7a49d | ||
|
|
374a9e9b13 | ||
|
|
d7d85c9e67 | ||
|
|
0fa982fe91 | ||
|
|
c4fb1bbd3e | ||
|
|
26d7d8f946 | ||
|
|
05ac342c4b | ||
|
|
5e491aef8f | ||
|
|
202187d5e2 | ||
|
|
13877b1c3a | ||
|
|
0e72e89ba3 | ||
|
|
f1b68271cc | ||
|
|
83efff21cd | ||
|
|
e6a4b35921 | ||
|
|
cc2d97ddba | ||
|
|
b78dd19cf9 | ||
|
|
1f0a75f836 | ||
|
|
16202c83a3 | ||
|
|
f4cc76c77d | ||
|
|
9fef6a9c43 | ||
|
|
810985b03a | ||
|
|
02633c6bd3 | ||
|
|
eb9416aa16 | ||
|
|
65b64af213 | ||
|
|
f1d160a481 | ||
|
|
f8c7f012a1 | ||
|
|
45618a55e6 | ||
|
|
c49583836b | ||
|
|
9dc8dc7065 | ||
|
|
f978e306cc | ||
|
|
aec852f971 | ||
|
|
143f781064 | ||
|
|
953b422295 | ||
|
|
da1a0f1299 | ||
|
|
bc7d19cfd8 | ||
|
|
df3122090f | ||
|
|
713fdf6188 | ||
|
|
70e21b58b4 | ||
|
|
18ebba1bb2 | ||
|
|
937404e6df | ||
|
|
88e155dd23 | ||
|
|
2e2cea0973 | ||
|
|
9b8c13a8b6 | ||
|
|
8238cc9f41 | ||
|
|
e59c4f30b8 | ||
|
|
fd7ef5edc2 | ||
|
|
a8e1acec09 | ||
|
|
31eb10985e | ||
|
|
1436b62a80 | ||
|
|
ec6a27635b | ||
|
|
1727744691 | ||
|
|
afdd47b5a5 | ||
|
|
d968e8e551 | ||
|
|
c649d1a762 | ||
|
|
a1f5d97345 | ||
|
|
58062515a5 | ||
|
|
5614b508f2 | ||
|
|
5e003a41b1 | ||
|
|
4eaeb1dd4a | ||
|
|
84471bd6f0 | ||
|
|
c8e3ac21c2 | ||
|
|
c38892b7cf | ||
|
|
1a0a5324f0 | ||
|
|
c1e9e9e0d6 | ||
|
|
57c93243a0 | ||
|
|
4259336e6d | ||
|
|
a18ce2e54d | ||
|
|
a0dc5d6183 | ||
|
|
5149f6808f | ||
|
|
918db33a8b | ||
|
|
95cbde9187 | ||
|
|
5e194393ff | ||
|
|
3da572a76a | ||
|
|
a116cba8ba | ||
|
|
fe5952fe14 | ||
|
|
2ab45ffd90 | ||
|
|
7c4932154c | ||
|
|
e2dbaa7c78 | ||
|
|
11a0dbc84a | ||
|
|
ee441643dd | ||
|
|
4c5affba99 | ||
|
|
c8148ef2cc | ||
|
|
b89740bad6 | ||
|
|
5c0d2b274c | ||
|
|
0489cd67c8 | ||
|
|
15d495e56e | ||
|
|
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 | ||
|
|
e64eb5db95 | ||
|
|
e175115de1 | ||
|
|
9dc4f95f60 | ||
|
|
54145b65ec | ||
|
|
15f53c0a6b | ||
|
|
2fad9c95db | ||
|
|
6bffbba817 | ||
|
|
ab11f5d583 | ||
|
|
c1d2210c95 | ||
|
|
b12ae13047 | ||
|
|
8a61c038ff | ||
|
|
b5f3603be5 | ||
|
|
802713aafe | ||
|
|
937fb80ee9 | ||
|
|
e14d7f9846 | ||
|
|
ad3c3a22bf | ||
|
|
e57b9ac6f3 | ||
|
|
36ea9b6b3c | ||
|
|
dbfd785af3 | ||
|
|
2cb55d7cb6 | ||
|
|
758da4d292 | ||
|
|
9c69dcc3e5 | ||
|
|
c4bc460dc7 | ||
|
|
c1893f40fc | ||
|
|
13f248b285 | ||
|
|
9bf51322cc | ||
|
|
f75227aad1 | ||
|
|
4a4633d188 | ||
|
|
9eb85a22bb | ||
|
|
d19a145472 | ||
|
|
4fab1610a2 | ||
|
|
82a2f7f77e | ||
|
|
50d2e75fd6 | ||
|
|
c8f1a9a912 | ||
|
|
60086a8aaf | ||
|
|
9d90474fd3 | ||
|
|
7fdabb6230 | ||
|
|
d0ab1ff641 | ||
|
|
53bd013899 | ||
|
|
464d571eb3 | ||
|
|
0faf6c1b50 | ||
|
|
e1e70e2137 | ||
|
|
372869d5e6 | ||
|
|
f3eb2bcb83 | ||
|
|
e7d20a1246 | ||
|
|
6d0917ff73 | ||
|
|
3c104fcd03 | ||
|
|
ee9e5671c4 | ||
|
|
4446a3ad7a | ||
|
|
f7df083ac0 | ||
|
|
b994aea18a | ||
|
|
944e136fe8 | ||
|
|
ad30854b59 | ||
|
|
a610cdb80d | ||
|
|
aa3f8f3990 | ||
|
|
cb67d98e4f | ||
|
|
0075cba438 | ||
|
|
077039bbf3 | ||
|
|
3fc9f9608b | ||
|
|
a1dffe454d | ||
|
|
56412b12e4 | ||
|
|
76df9bfcf1 | ||
|
|
f3a4bcfc1f | ||
|
|
045a3dcf31 | ||
|
|
75be1e3eca | ||
|
|
cc04a082b1 | ||
|
|
78fdec17e5 | ||
|
|
f3418a04de | ||
|
|
2836a0ba1c | ||
|
|
58e807b81e |
@@ -0,0 +1,21 @@
|
||||
# DingTalk Workspace CLI Environment Variables
|
||||
# Copy this file to .env and fill in your values
|
||||
|
||||
# DingTalk App Credentials (required for OAuth authentication)
|
||||
# DWS_CLIENT_ID=<your-dingtalk-app-client-id>
|
||||
# DWS_CLIENT_SECRET=<your-dingtalk-app-client-secret>
|
||||
|
||||
# Configuration directory (optional, defaults to ~/.dws)
|
||||
# DWS_CONFIG_DIR=
|
||||
|
||||
# Language setting (optional, defaults to system locale)
|
||||
# DWS_LANG=
|
||||
|
||||
# Trusted domains for API requests (optional)
|
||||
# DWS_TRUSTED_DOMAINS=*.dingtalk.com
|
||||
|
||||
# Allow HTTP endpoints (0=no, 1=yes; default: 0)
|
||||
# DWS_ALLOW_HTTP_ENDPOINTS=0
|
||||
|
||||
# Cache directory (optional, defaults to ~/.dws/cache)
|
||||
# DWS_CACHE_DIR=
|
||||
@@ -0,0 +1,4 @@
|
||||
# Default code owners for all files
|
||||
# These users will be automatically requested for review on PRs.
|
||||
|
||||
* @DingTalk-Real-AI/cli-maintainers
|
||||
@@ -0,0 +1,36 @@
|
||||
---
|
||||
name: Bug Report
|
||||
about: Report a bug to help us improve
|
||||
title: "[Bug] "
|
||||
labels: bug
|
||||
assignees: ''
|
||||
---
|
||||
|
||||
## Describe the Bug
|
||||
|
||||
A clear and concise description of what the bug is.
|
||||
|
||||
## Steps to Reproduce
|
||||
|
||||
1. Run `dws ...`
|
||||
2. ...
|
||||
3. See error
|
||||
|
||||
## Expected Behavior
|
||||
|
||||
A clear and concise description of what you expected to happen.
|
||||
|
||||
## Actual Behavior
|
||||
|
||||
What actually happened, including any error messages or unexpected output.
|
||||
|
||||
## Environment
|
||||
|
||||
- **OS**: [e.g., macOS 15.2, Ubuntu 24.04, Windows 11]
|
||||
- **Architecture**: [e.g., arm64, amd64]
|
||||
- **CLI Version**: [output of `dws version`]
|
||||
- **Go Version** (if building from source): [output of `go version`]
|
||||
|
||||
## Additional Context
|
||||
|
||||
Add any other context about the problem here (logs, screenshots, etc.).
|
||||
@@ -0,0 +1,27 @@
|
||||
---
|
||||
name: Feature Request
|
||||
about: Suggest an idea for this project
|
||||
title: "[Feature] "
|
||||
labels: enhancement
|
||||
assignees: ''
|
||||
---
|
||||
|
||||
## Problem Statement
|
||||
|
||||
A clear and concise description of the problem or limitation you are experiencing.
|
||||
|
||||
## Proposed Solution
|
||||
|
||||
Describe the solution you'd like. Include any specific CLI commands, flags, or behaviors you envision.
|
||||
|
||||
## Alternatives Considered
|
||||
|
||||
A clear and concise description of any alternative solutions or features you've considered.
|
||||
|
||||
## Use Case
|
||||
|
||||
Describe the use case(s) that would benefit from this feature.
|
||||
|
||||
## Additional Context
|
||||
|
||||
Add any other context, mockups, or examples about the feature request here.
|
||||
@@ -0,0 +1,17 @@
|
||||
## Summary
|
||||
|
||||
- What changed?
|
||||
- Why is this change needed?
|
||||
|
||||
## Verification
|
||||
|
||||
- [ ] `make build`
|
||||
- [ ] `make lint`
|
||||
- [ ] `make test`
|
||||
- [ ] `make policy`
|
||||
- [ ] `./scripts/policy/check-generated-drift.sh`
|
||||
- [ ] `./scripts/policy/check-command-surface.sh --strict` (if command surface changed)
|
||||
|
||||
## Notes
|
||||
|
||||
- Any risks, follow-up work, or intentional scope cuts
|
||||
@@ -0,0 +1 @@
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="108" height="20" role="img" aria-label="coverage: 49.8%"><title>coverage: 49.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">49.8%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">49.8%</text></g></svg>
|
||||
|
After Width: | Height: | Size: 1.1 KiB |
@@ -0,0 +1,181 @@
|
||||
name: CI
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
pull_request:
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
lint:
|
||||
name: Lint
|
||||
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: Format Check
|
||||
run: |
|
||||
unformatted="$(find cmd internal test -name '*.go' -print0 | xargs -0r gofmt -l)"
|
||||
test -z "$unformatted" || (printf '%s\n' "$unformatted" && exit 1)
|
||||
|
||||
- name: Go Vet
|
||||
run: go vet ./...
|
||||
|
||||
# golangci-lint temporarily disabled: v1.64.8 built with Go 1.24 is incompatible with Go 1.25
|
||||
# - name: golangci-lint
|
||||
# uses: golangci/golangci-lint-action@v6
|
||||
# with:
|
||||
# version: v1.64.8
|
||||
# args: ./...
|
||||
|
||||
test:
|
||||
name: Test
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
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: Install archive tooling
|
||||
run: sudo apt-get update && sudo apt-get install -y zip unzip
|
||||
|
||||
- name: Build
|
||||
run: make build
|
||||
|
||||
- name: Test with Race Detection
|
||||
run: go test -v -race -count=1 -timeout=5m ./cmd/... ./internal/...
|
||||
|
||||
coverage:
|
||||
name: Coverage
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
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: Install archive tooling
|
||||
run: sudo apt-get update && sudo apt-get install -y zip unzip
|
||||
|
||||
- name: Build
|
||||
run: make build
|
||||
|
||||
- name: Run tests with coverage
|
||||
run: |
|
||||
go test -coverprofile=coverage.txt -covermode=atomic ./cmd/... ./internal/...
|
||||
go tool cover -func=coverage.txt
|
||||
|
||||
- name: Generate coverage report
|
||||
run: go tool cover -html=coverage.txt -o coverage.html
|
||||
|
||||
- name: Upload coverage artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: coverage-report
|
||||
path: |
|
||||
coverage.txt
|
||||
coverage.html
|
||||
|
||||
- name: Update coverage badge
|
||||
if: github.ref == 'refs/heads/main'
|
||||
run: |
|
||||
COVERAGE=$(go tool cover -func=coverage.txt | grep total | awk '{print $3}' | sed 's/%//')
|
||||
echo "Coverage: ${COVERAGE}%"
|
||||
if (( $(echo "$COVERAGE >= 80" | bc -l) )); then
|
||||
COLOR="brightgreen"
|
||||
elif (( $(echo "$COVERAGE >= 60" | bc -l) )); then
|
||||
COLOR="yellow"
|
||||
else
|
||||
COLOR="red"
|
||||
fi
|
||||
mkdir -p .github/badges
|
||||
curl -s "https://img.shields.io/badge/coverage-${COVERAGE}%25-${COLOR}" > .github/badges/coverage.svg
|
||||
|
||||
- name: Commit badge
|
||||
if: github.ref == 'refs/heads/main'
|
||||
run: |
|
||||
git config --local user.email "github-actions[bot]@users.noreply.github.com"
|
||||
git config --local user.name "github-actions[bot]"
|
||||
git add .github/badges/coverage.svg || true
|
||||
git diff --staged --quiet || git commit -m "chore: update coverage badge [skip ci]"
|
||||
git push || true
|
||||
|
||||
policy:
|
||||
name: Policy Check
|
||||
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: Build
|
||||
run: make build
|
||||
|
||||
- name: Policy
|
||||
run: make policy
|
||||
|
||||
- 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 }}
|
||||
@@ -0,0 +1,59 @@
|
||||
name: Release
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- "v*"
|
||||
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
jobs:
|
||||
release:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Install archive tooling
|
||||
run: sudo apt-get update && sudo apt-get install -y zip unzip
|
||||
|
||||
- name: Run GoReleaser
|
||||
uses: goreleaser/goreleaser-action@v6
|
||||
with:
|
||||
version: "~> v2"
|
||||
args: release --clean
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- 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
|
||||
with:
|
||||
node-version: "20"
|
||||
registry-url: "https://registry.npmjs.org"
|
||||
|
||||
- name: Publish to npm
|
||||
working-directory: dist/npm/dingtalk-workspace-cli
|
||||
run: npm publish --access public
|
||||
env:
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
+22
-41
@@ -1,51 +1,32 @@
|
||||
# Python
|
||||
__pycache__/
|
||||
.pytest_cache/
|
||||
.venv*/
|
||||
*.pyc
|
||||
*.pyo
|
||||
|
||||
# Build outputs
|
||||
dist/
|
||||
.tmp-bin/
|
||||
dws
|
||||
|
||||
# IDE
|
||||
.idea/
|
||||
.vscode/
|
||||
*.swp
|
||||
*.swo
|
||||
*~
|
||||
|
||||
# OS
|
||||
.worktrees/
|
||||
.pytest_cache/
|
||||
.venv*/
|
||||
var/
|
||||
.DS_Store
|
||||
Thumbs.db
|
||||
.agents/
|
||||
/.idea/
|
||||
/CLAUDE.md
|
||||
/.claude/
|
||||
/.idea/vcs.xml
|
||||
/.idea/.gitignore
|
||||
dws
|
||||
test/cli/testdata/
|
||||
tmp/
|
||||
test/cli_compat/testdata/
|
||||
.gitignore
|
||||
.worktrees/
|
||||
|
||||
# Environment & Secrets
|
||||
# Secrets & credentials
|
||||
.env
|
||||
.env.*
|
||||
*.pem
|
||||
*.key
|
||||
credentials*
|
||||
|
||||
# Test artifacts
|
||||
*.log
|
||||
coverage/
|
||||
test/cli/testdata/
|
||||
test/cli_compat/testdata/
|
||||
|
||||
# Local working directories
|
||||
.worktrees/
|
||||
.agents/
|
||||
var/
|
||||
|
||||
# Node (if applicable)
|
||||
node_modules/
|
||||
npm-debug.log*
|
||||
|
||||
# Plans (local design docs)
|
||||
docs/plans/
|
||||
|
||||
# Claude
|
||||
/CLAUDE.md
|
||||
/.claude/
|
||||
plans
|
||||
_docs
|
||||
dws.zip
|
||||
*.code-workspace
|
||||
/dingtalk-workspace.zip
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
# GoReleaser configuration for dws
|
||||
# Docs: https://goreleaser.com
|
||||
#
|
||||
# To release:
|
||||
# git tag -a v0.1.0 -m "Release v0.1.0"
|
||||
# git push origin v0.1.0
|
||||
#
|
||||
# To test locally (no publish):
|
||||
# goreleaser release --snapshot --clean
|
||||
|
||||
version: 2
|
||||
|
||||
before:
|
||||
hooks:
|
||||
- go mod tidy
|
||||
|
||||
builds:
|
||||
- main: ./cmd
|
||||
binary: dws
|
||||
env:
|
||||
- CGO_ENABLED=0
|
||||
flags:
|
||||
- -buildmode=pie
|
||||
- -trimpath
|
||||
ldflags:
|
||||
- -s -w
|
||||
- -X github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app.version=v{{.Version}}
|
||||
- -X github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app.gitCommit={{.ShortCommit}}
|
||||
- -X github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/app.buildTime={{.Date}}
|
||||
goos:
|
||||
- darwin
|
||||
- linux
|
||||
- windows
|
||||
goarch:
|
||||
- amd64
|
||||
- arm64
|
||||
|
||||
archives:
|
||||
- formats:
|
||||
- tar.gz
|
||||
name_template: "dws-{{ .Os }}-{{ .Arch }}"
|
||||
format_overrides:
|
||||
- goos: windows
|
||||
formats:
|
||||
- zip
|
||||
files:
|
||||
- LICENSE
|
||||
- NOTICE
|
||||
- README.md
|
||||
- CHANGELOG.md
|
||||
|
||||
checksum:
|
||||
name_template: "checksums.txt"
|
||||
algorithm: sha256
|
||||
|
||||
changelog:
|
||||
sort: asc
|
||||
filters:
|
||||
exclude:
|
||||
- "^docs:"
|
||||
- "^test:"
|
||||
- "^ci:"
|
||||
- "^chore:"
|
||||
|
||||
release:
|
||||
github:
|
||||
owner: DingTalk-Real-AI
|
||||
name: dingtalk-workspace-cli
|
||||
draft: false
|
||||
prerelease: auto
|
||||
name_template: "v{{.Version}}"
|
||||
mode: replace
|
||||
+341
@@ -4,6 +4,347 @@ All notable changes to this project will be documented in this file.
|
||||
|
||||
The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and this project follows [Semantic Versioning](https://semver.org/).
|
||||
|
||||
## [1.0.13] - 2026-04-22
|
||||
|
||||
IM / Messaging capability expansion: the `chat` (aka `im`) product surface grows from "group + bot messaging" into a full conversational layer — user-identity messaging, message reading & search, personal messages, topic replies, mentions, focused contacts, unread/top/common conversations, org-wide group creation, and first-class bot lifecycle.
|
||||
|
||||
### Added
|
||||
|
||||
- **`dws im` alias** — `dws im` is now registered as an alias of `dws chat` for intent clarity
|
||||
- **User-identity messaging** (`chat message send`) — send group or 1-on-1 messages as the current user
|
||||
- Recipient selection is mutually exclusive: `--group <openConversationId>` / `--user <userId>` / `--open-dingtalk-id <openDingTalkId>`
|
||||
- Markdown text via `--text` (or positional arg), optional `--title`
|
||||
- Group-only: `--at-all` to @everyone, `--at-users` for per-member @mentions
|
||||
- Image messages via `--media-id` (obtained from `dt_media_upload`)
|
||||
- **Personal messages** (`chat message send-personal`) — sensitive personal-channel send (⚠️ destructive/dangerous op, requires confirmation)
|
||||
- **Conversation read paths**:
|
||||
- `chat message list` — pull group / 1-on-1 conversation messages
|
||||
- `chat message list-all` — pull all conversations for the current user in a time range
|
||||
- `chat message list-topic-replies` — pull group topic reply threads
|
||||
- `chat message list-by-sender` — messages by a specific sender
|
||||
- `chat message list-mentions` — messages where the current user was @-mentioned
|
||||
- `chat message list-focused` — messages from focused / starred contacts
|
||||
- `chat message list-unread-conversations` — unread conversation list
|
||||
- `chat message search` — keyword search across conversations
|
||||
- `chat message info` — conversation metadata
|
||||
- `chat list-top-conversations` — pinned conversation list
|
||||
- **Group creation & discovery**:
|
||||
- `chat group create-org` — create an organization-wide group
|
||||
- `chat search-common` — search groups shared with a nickname list (`--nicks`, `--match-mode AND|OR`, cursor-based pagination)
|
||||
- **Bot lifecycle**:
|
||||
- `chat bot create` — create an enterprise bot
|
||||
- `chat bot search-groups` — search the groups a bot is present in
|
||||
|
||||
### Changed
|
||||
|
||||
- **`chat` skill reference** (`skills/references/products/chat.md`, #148) restructured into three sub-groups — `group` (9) / `message` (15) / `bot` (3) — with refreshed intent-routing table, workflow examples, and context-passing rules aligned with `dws-service-endpoints.json` (16 new group-chat tool overrides + 2 new bot tool overrides)
|
||||
- **README Key Services** sync:
|
||||
- `Chat` row: 10 → 20 commands; subcommand tags expanded to `message` `group` `search` `list-top-conversations`
|
||||
- `Bot` row: 6 → 7 commands; subcommand tags expanded with `create` `search-groups`
|
||||
- Total raised to **152 commands across 14 products**
|
||||
|
||||
## [1.0.12] - 2026-04-21
|
||||
|
||||
Product-surface expansion: first-class `doc` (DingTalk Docs) and `minutes` (AI Minutes) skill references, refreshed `aitable` guide aligned with the shipped binary (including dashboard / chart / export), and a README sync that brings the full command catalog to **141 commands across 14 products**.
|
||||
|
||||
### Added
|
||||
|
||||
- **`doc` skill reference** (`skills/references/products/doc.md`) — 16-command coverage of DingTalk Docs:
|
||||
- Discovery: `search`, `list`, `info`, `read`
|
||||
- Authoring: `create`, `update`, `folder create`
|
||||
- Files: `upload`, `download`
|
||||
- Block-level editing: block `query`, `insert`, `update`, `delete`
|
||||
- Comments: `comment list`, `create`, `reply`
|
||||
- URL → `doc_id` extraction rules and nodeId dual-format notes
|
||||
- **`minutes` skill reference** (`skills/references/products/minutes.md`) — coverage of AI Minutes:
|
||||
- Lists: personal / shared-with-me / all-accessible
|
||||
- Content: basic info, AI summary, keywords, transcription, extracted todos, batch detail
|
||||
- Editing: title update
|
||||
- Recording control: start, pause, resume, stop
|
||||
- **SKILL.md routing**:
|
||||
- Product overview table rows for `doc` and `minutes`
|
||||
- Intent decision tree routes — `钉钉文档/云文档/知识库/块级编辑/文档评论` → `doc`; `听记/AI听记/会议纪要/转写/摘要/思维导图/发言人/热词` → `minutes`
|
||||
- Danger-op table entries: `doc delete`, `doc block delete`
|
||||
- `aitable` description completed with the `附件` (attachment) group
|
||||
- **`aitable` skill enhancements**:
|
||||
- `field create` single-field mode (`--name` / `--type` / `--config`) with examples
|
||||
- `base get` URL → `baseId` quick-tip
|
||||
- Dedicated "URL → baseId 提取" chapter
|
||||
- "`--filters` 筛选语法排错与使用规范" chapter
|
||||
- "相关产品" cross-link section pointing to `doc`
|
||||
- **"复杂操作" chapter** (#141) — dashboard / chart workflow (with two-call sequencing and `chart share get` vs `dashboard share get` error semantics) and two-stage `export data` polling (`scope=all/table/view` parameter constraints)
|
||||
- **README Key Services sync** (#140):
|
||||
- New rows: `doc` (16 commands), `minutes` (22 commands — adds `hot-word`, `mind-graph`, `replace-text`, `speaker`, `upload` subgroups)
|
||||
- `aitable` expanded from 20 → 37 commands; surfaces `chart`, `dashboard`, `export`, `import`, `view` subgroups
|
||||
- Total command count updated from **86 → 141 across 14 products**
|
||||
- "Coming soon" list drops `doc` and `minutes`
|
||||
|
||||
### Changed
|
||||
|
||||
- `aitable record query` docs rename `--keyword` → `--query` to match the shipped binary
|
||||
- `aitable record query` docs clarify `--sort` direction semantics (avoids misuse of `order`)
|
||||
- `aitable base list` guidance strengthened — "only for recent browsing; use `base search` for lookups"; intent decision prioritizes `base search` for base discovery
|
||||
|
||||
## [1.0.11] - 2026-04-20
|
||||
|
||||
Plugin subsystem hardening: faster cold startup, cleaner lifecycle, stricter isolation, and polished UX for PAT / i18n / error routing.
|
||||
|
||||
### Added
|
||||
|
||||
- `feat: supports claw-like products` — overlay path for Claw-style embedded editions
|
||||
- `feat(plugin): inject user identity (UserID, CorpID) into stdio plugin subprocesses`
|
||||
- `feat(auth): improve login UX for terminal auth denial cases` — clearer messaging + retry affordance
|
||||
- `feat: PAT scope error visualization and auto-retry with authorization polling` (#113)
|
||||
- Human-readable error output (lark-cli style) with type/message/hint/authorization command
|
||||
- JSON payload also available via `--format json`
|
||||
- Auto-retry once the user completes scope authorization
|
||||
|
||||
### Changed
|
||||
|
||||
- `perf(plugin): serve plugin MCP tool list from disk cache on startup` — hot path skips Initialize+ListTools when snapshot exists
|
||||
- `perf(plugin): parallelize all plugin discovery and tighten cold timeouts` — HTTP cold budget 4s → 700ms (auth) / 500ms (plain); stdio and HTTP fan out concurrently
|
||||
- `perf(plugin): share cache.Store across discovery` — single `*cache.Store` above the fan-out instead of per-goroutine instances
|
||||
- `refactor(plugin): remove default/managed plugin privileged mechanism` (#124) — third-party plugins install on an equal footing via `dws plugin install`
|
||||
- `refactor(plugin): purge removed plugin settings instead of merely disabling` — `RemovePlugin` now deletes `EnabledPlugins` and `PluginConfigs` entries
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(transport): cap plugin MCP startup at ~4s when endpoints are unreachable` (#119) — eliminates the 10s `dws --help` stall caused by compounding transport timeouts
|
||||
- `fix(plugin): stop stdio child processes on exit and before removal` — no more orphaned plugin subprocesses
|
||||
- `fix(pat): avoid shared PAT command state in root registration` (#129)
|
||||
- `fix: -f json 模式下错误 JSON 从 stdout 改为输出到 stderr` (#133) — restores CI stderr-based failure assertions
|
||||
- `fix(cli): localize plugin/help command strings via i18n` (#118, #134) — zh locale now shows consistent Chinese `--help`; wraps plugin module, help command, and OAuth client-id/secret flag descriptions
|
||||
- `chore: remove workspace and bundled artifacts` (#127) — clean local-only repository leftovers
|
||||
|
||||
## [1.0.9] - 2026-04-16
|
||||
|
||||
Plugin system launch + execution-pipeline overhaul. This is the largest release since 1.0.0: third-party MCP servers become first-class commands, the command pipeline grows to five stages, and the edition overlay gains the hooks needed for embedded hosts.
|
||||
|
||||
### Added
|
||||
|
||||
#### Plugin system (new)
|
||||
|
||||
- `plugin` command family: `install`, `list`, `info`, `enable`, `disable`, `remove`, `create`, `dev`, `config set/get/list/unset`
|
||||
- Plugin manifest parsing/validation, managed/user directory-based identity
|
||||
- MCP server conversion and injection into the dynamic routing registry
|
||||
- Pipeline hook adapter for shell-based hooks
|
||||
- Stdio transport: subprocess lifecycle, `DWS_PLUGIN_ROOT` / `DWS_PLUGIN_DATA` variable expansion
|
||||
- Stdio server tools automatically registered as CLI subcommands (e.g. `dws hello greet --name Peter`)
|
||||
- Streamable-HTTP MCP tool discovery via `registerHTTPServer`
|
||||
- Updater: managed plugin update check on CLI startup (10 s timeout, best-effort)
|
||||
- `dws plugin create` scaffold (plugin.json, SKILL.md, hooks.json); `dws plugin dev` source-dir registration without copy
|
||||
- `SyncSkills` — copies plugin skills to agent directories on startup
|
||||
- **Auth Token Registry**: per-server HTTP headers declared in `plugin.json` for third-party MCP servers (e.g. Alibaba Cloud Bailian) independent from DingTalk OAuth
|
||||
- **Persistent plugin config** (`dws plugin config ...`): values persisted to `~/.dws/settings.json`, auto-injected as env vars; `${KEY}` in `plugin.json` resolves without manual `export`
|
||||
- **Build lifecycle**: `build` field compiles stdio servers to native binaries at install time
|
||||
- **Command-name conflict protection**: reserved built-in names (`auth`, `plugin`, `cache`, …) and plugin-vs-plugin duplicate detection
|
||||
- Parallel service discovery (`sync.WaitGroup`) — startup reduced from sequential `N*10s` to parallel `max(10s)`
|
||||
|
||||
#### Core commands & diagnostics
|
||||
|
||||
- `dws doctor` — one-stop environment/auth/network diagnostics
|
||||
- `dws config list` — centralized view of scattered configuration
|
||||
- Structured perf tracing (upgraded from debug tool to diagnostics output)
|
||||
- `feat(skill): restore find/get for legacy skill market API` — `skill find`, `skill get`; `skill add` still uses aihub download
|
||||
|
||||
#### Edition / overlay hooks
|
||||
|
||||
- `edition.Hooks.SaveToken` / `LoadToken` / `DeleteToken` — delegate token persistence with keychain fallback
|
||||
- `edition.Hooks.AuthClientID` / `AuthClientFromMCP` — overlay can override the OAuth client ID and route auth through MCP endpoints
|
||||
- `edition.Hooks.AfterPersistentPreRun` — wire non-MCP clients (e.g. A2A gateway) after root setup
|
||||
- `edition.Hooks.ClassifyToolResult` — custom MCP result classification before the default business-error detection
|
||||
- Token marker file (`token.json`) for embedded hosts to detect auth state without keychain access
|
||||
- `pkg/runtimetoken.ResolveAccessToken` mirroring MCP auth resolution; MCP identity headers exported via `pkg/cli` for auxiliary HTTP transports
|
||||
- `ExitCoder` interface — edition-specific errors carry custom exit codes
|
||||
- `RawStderrError` interface — errors that bypass CLI formatting and emit raw stderr (for desktop runtimes)
|
||||
|
||||
### Changed
|
||||
|
||||
- **Command execution pipeline: 3 → 5 stages** (`Register → PreParse → PostParse → PreRequest → PostResponse`)
|
||||
- `feat(schema): return structured degraded errors instead of silent empty catalog` — new `CatalogDegraded` error with reasons `unauthenticated` / `market_unreachable` / `runtime_all_failed`; auth pre-check short-circuits doomed MCP connections
|
||||
- `refactor(auth): unify auxiliary token resolution with MCP cached path` — shared `resolveAccessTokenFromDir`; overlays reuse the process-level token cache
|
||||
- `feat(plugin): improve CLI overlay resolution and plugin install robustness`
|
||||
- `plugin.json` `cli` field now accepts a file path (e.g. `"cli": "overlay.json"`) in addition to inline JSON
|
||||
- `description` field on `CLIToolOverride` for static fallback when MCP `tools/list` is unavailable
|
||||
- Windows install uses `cmd /C` instead of `sh -c` for build commands
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(plugin): harden plugin system security boundaries`
|
||||
- Reject `file://` / local paths in git URLs; allow only `https` / `ssh`
|
||||
- Reject symlink entries during ZIP extraction (path-traversal defense)
|
||||
- `build.output` must be a relative path within the plugin directory
|
||||
- Reject absolute paths in stdio command declarations
|
||||
- Block dangerous env var names (`PATH`, `LD_PRELOAD`, …) from plugin config injection
|
||||
- `fix(plugin): schema flag params, HTTP tool discovery, and integration tests`
|
||||
- `fix(plugin): skip min version check in dev mode`
|
||||
|
||||
## [1.0.8] - 2026-04-07
|
||||
|
||||
AITable command surface expansion, installer alignment with npm conventions, and execution-timeout hardening.
|
||||
|
||||
### Added
|
||||
|
||||
- **AITable static helper commands** (20 commands in total) replacing dynamic routing:
|
||||
- `base`: `list`, `search`, `get`, `create`, `update`
|
||||
- `table`: `get`, `create`, `update`
|
||||
- `field`: `get`, `create`, `update`
|
||||
- `record`: `query`, `create`, `update`
|
||||
- `template`: `search`
|
||||
- `attachment`: `upload`
|
||||
- `feat(install): align skill dirs with npm and add OpenClaw` — skill install paths follow npm conventions; OpenClaw added to supported agents
|
||||
- Label rendering optimization for AITable records (`to #73551688`)
|
||||
- README: npm install method documented
|
||||
- README: note that `dws upgrade` requires v1.0.7+
|
||||
|
||||
### Changed
|
||||
|
||||
- `perf: optimize command timeout handling, instrumentation, and diagnostics`
|
||||
|
||||
## [1.0.7] - 2026-04-02
|
||||
|
||||
Self-upgrade, edition overlay foundation, and fail-closed auth enforcement.
|
||||
|
||||
### Added
|
||||
|
||||
- **`dws upgrade`** — self-upgrade via GitHub Releases; atomic replace; cross-platform (macOS/Linux/Windows)
|
||||
- `feat: edition layer for Wukong overlay` — build-time edition hook lets downstream overlays customize auth UX, config dir, static server list, visible products, and extra root commands
|
||||
- `pkg/edition` defaults + `pkg/editiontest` contract tests
|
||||
- `Makefile` target `edition-test`; CI job `edition-tests`
|
||||
- Static server injection skips market discovery when configured
|
||||
- Deduplicates top-level commands so overlay wins
|
||||
- `hideNonDirectRuntimeCommands` respects edition `VisibleProducts`
|
||||
- Gated `auth login` subcommand + hints for embedded editions
|
||||
- Optional token auto-purge; edition `ConfigDir` override
|
||||
- `dws version` — human-readable multi-line output plus JSON with edition, architecture, build, commit
|
||||
- Tag reporting for case suites (`to #73551688`)
|
||||
- `feat(auth): unify MCP retry constant and add retry to remaining endpoints`
|
||||
|
||||
### Changed
|
||||
|
||||
- `style(auth): redesign OAuth authorization pages UI`
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix(auth): switch CLI auth check from fail-open to fail-closed`
|
||||
- When `/cli/cliAuthEnabled` is unreachable (network error/timeout/5xx), OAuth callback now routes to the permission request page instead of silently marking "enabled"
|
||||
- Device Flow blocks login and asks the user to verify network connectivity
|
||||
- `CheckCLIAuthEnabled` retries with backoff (3 attempts, 0s/1s/2s) to tolerate transient issues
|
||||
|
||||
## [1.0.6] - 2026-04-01
|
||||
|
||||
Error diagnostics overhaul, destructive-command confirmation, and credential auto-persistence.
|
||||
|
||||
### Added
|
||||
|
||||
- **Interactive confirmation for destructive dynamic commands** — prompts before delete/remove operations unless `--yes` is set
|
||||
- **Enhanced error diagnostics**
|
||||
- `ServerDiagnostics` struct extracts `trace_id`, `server_error_code`, `technical_detail`, `server_retryable` from MCP responses
|
||||
- Pulls diagnostics from JSON-RPC `error.data`, tool call result content, and HTTP headers (`X-Trace-Id`, `X-Request-Id`, `x-dingtalk-trace-id`)
|
||||
- Three verbosity levels for `PrintHuman`: Normal (trace ID + server code), Verbose (+ technical detail), Debug (+ RPC code / operation / reason)
|
||||
- Local logging now includes sanitized request body, response body on error, retry attempts, and classification events
|
||||
- `TruncateBody` / `SanitizeArguments` / `RedactHeaders` helpers with sensitive-key substring detection
|
||||
- **Auth credential persistence**
|
||||
- `feat(auth): enhance device flow with CLI auth check and admin guidance`
|
||||
- `feat(auth): persist OAuth credentials for reliable token refresh`
|
||||
- `feat(auth): persist client credentials and optimize keychain access` — auto-persist `--client-id` / `--client-secret`; keychain credential cache to avoid repeated reads; enhanced logout cleans `app.json` + keychain secrets + `token.json`
|
||||
- `add report helper with flexible date parsing and defaults`
|
||||
- `feat: to #73551688 支持消息通知`
|
||||
- README: Official App mode (recommended, direct login without creating an app) + Custom App mode; admin guide for enabling CLI access
|
||||
|
||||
### Changed
|
||||
|
||||
- Getting Started simplified with inline login commands; whitelist references removed from the IMPORTANT banner
|
||||
- Version bump documentation updated to v1.0.5 internal; co-creation group QR code refreshed
|
||||
|
||||
### Fixed
|
||||
|
||||
- `fix: resolve verbosity flag lookup, FileLogger lazy binding, and business error logging`
|
||||
- `resolveVerbosity` uses `cmd.Flags()` instead of `PersistentFlags()` so subcommands inherit `--verbose` / `--debug`
|
||||
- `FileLogger` lazy-binds in `executeInvocation` (after `configureLogLevel` init)
|
||||
- Business errors (HTTP 200 + `success=false`) now written to the file logger for offline diagnosis
|
||||
- OAuth callback race condition (write response before sending code)
|
||||
- `import path for errors package in skill_command.go`
|
||||
|
||||
## [1.0.4] - 2026-03-30
|
||||
|
||||
Token-refresh reliability and onboarding clarity.
|
||||
|
||||
### Added
|
||||
|
||||
- `feat(auth): persist client credentials for token refresh` — `--client-id` / `--client-secret` are stored for automatic refresh after expiration; client secret lives in the system Keychain with a file reference
|
||||
- README onboarding flow rewrite with step-by-step first-time setup and more realistic examples
|
||||
- Agent skill reference polish: clearer examples, updated intent routing patterns, expanded `simple.md` onboarding, cross-skill reference fixes
|
||||
|
||||
## [1.0.3] - 2026-03-29
|
||||
|
||||
Filtering power, schema rendering, and a native `todo` command family.
|
||||
|
||||
### Added
|
||||
|
||||
- **Nested / array-indexed output filtering**
|
||||
- `--fields` now accepts dot-notation (e.g. `--fields response.content`) and array index access (e.g. `response.items[0]`)
|
||||
- New field-path parser with recursive extraction logic
|
||||
- **`schema` command enhancements**
|
||||
- Table format output for human consumption
|
||||
- Product-level endpoint loading in the CLI loader
|
||||
- Schema-text rendering wired into the runner output pipeline
|
||||
- **`todo` task helper family** — static `create` / `update` / `done` / `get` / `delete` with `preferLegacyLeaf` replacing dynamic commands
|
||||
- MCP tool alignment: `create_personal_todo`, `update_todo_task`, `update_todo_done_status`, `query_todo_detail`, `delete_todo`
|
||||
- ISO-8601 due-time parsing
|
||||
- Hidden title aliases and delete confirmation
|
||||
- Priority field on `todo` helper
|
||||
- Expanded zh / en i18n coverage (fixes `en.json` spacing/wording issues)
|
||||
- README restructured with collapsible feature sections
|
||||
|
||||
## [1.0.2] - 2026-03-29
|
||||
|
||||
Deep workspace tooling upgrade: pipeline-based input correction, output filtering, enhanced stdin handling, and multi-endpoint routing.
|
||||
|
||||
### Added
|
||||
|
||||
- Pipeline engine (`internal/pipeline`) for pre-parse and post-parse input correction
|
||||
- `AliasHandler`: normalises model-generated flag casing (e.g. `--userId` → `--user-id`)
|
||||
- `StickyHandler`: splits glued flag values (e.g. `--limit100` → `--limit 100`)
|
||||
- `ParamNameHandler`: fixes near-miss flag typos (e.g. `--limt` → `--limit`)
|
||||
- `ParamValueHandler`: normalises structured parameter values after parsing
|
||||
- Output filtering via `--fields` and `--jq` global flags (`internal/output/filter.go`)
|
||||
- `--fields`: comma-separated field selection for top-level keys (case-insensitive)
|
||||
- `--jq`: jq expression filtering powered by `gojq` library
|
||||
- `StdinGuard` for safe single-read stdin across multiple flags in one invocation
|
||||
- `ResolveInputSource` unified resolver supporting `@file`, `@-` (explicit stdin), and implicit pipe fallback
|
||||
- `@file` / `@-` syntax support for all string-typed override flags in tool commands
|
||||
- Chat helper support for `@file` input to read message content from files
|
||||
- Tool-level endpoint routing (`dynamicToolEndpoints`) for multi-endpoint products
|
||||
- Comprehensive test suites for pipeline handlers, stdin guard, canonical commands, and chat input
|
||||
|
||||
### Changed
|
||||
|
||||
- `directRuntimeEndpoint` now accepts tool name for finer-grained endpoint resolution
|
||||
- `collectOverrides` resolves `@file` / `@-` for all string-typed flags
|
||||
- `NewRootCommand` refactored to `NewRootCommandWithEngine` with optional pipeline engine
|
||||
- `schema` command no longer hidden (visible in help output)
|
||||
- Default output format changed from `table` to `json`
|
||||
|
||||
## [1.0.1] - 2026-03-28
|
||||
|
||||
Backward-compatible feature and security update after the initial 1.0.0 release.
|
||||
|
||||
### Added
|
||||
|
||||
- JSON output support for `dws auth login` and `dws auth status`
|
||||
- Cross-platform keychain-backed secure storage and migration helpers
|
||||
- Atomic file write helpers to avoid partial config and download writes
|
||||
- Stronger path and input validation helpers for local file operations
|
||||
- Install-script coverage for local-source installs
|
||||
|
||||
### Changed
|
||||
|
||||
- Improved `auth login` help text, hidden compatibility flags, and interactive UX
|
||||
- Added root-level flag suggestions for common compatibility mistakes such as `--json` and legacy auth flags
|
||||
- Updated AITable upload parsing to accept nested `content` payloads
|
||||
- Refreshed bundled skills metadata for the new CLI version
|
||||
|
||||
## [1.0.0] - 2026-03-27
|
||||
|
||||
First public release of DingTalk Workspace CLI.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -1,360 +1,431 @@
|
||||
<h1 align="center">DingTalk Workspace CLI (dws)</h1>
|
||||
|
||||
**一个 CLI 搞定钉钉工作台所有功能 — 为人类和 AI Agent 而生。**<br>
|
||||
覆盖通讯录、日历、待办、考勤、智能表格等核心能力,无需样板代码即可调用,所有响应均为结构化 JSON 输出,并内置 Agent Skills 让 AI 开箱即用。
|
||||
<p align="center"><code>dws</code> — DingTalk Workspace on the command line, built for humans and AI agents.</p>
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i1/O1CN01oKAc2r28jOyyspcQt_!!6000000007968-2-tps-4096-1701.png" alt="DWS Product Overview" width="100%">
|
||||
</p>
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。当前为灰度共创阶段,请加入钉钉 DWS 共创群,提供以下材料给官方人员完成白名单配置:① 钉钉应用的 Client ID;② 企业主管理员确认开通的凭证。后续将支持企业管理员自助审批开通。
|
||||
>
|
||||
> <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>
|
||||
|
||||
<p>
|
||||
<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.0.0-red" alt="v1.0.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>
|
||||
|
||||
## 目录
|
||||
<p align="center">
|
||||
<a href="./README_zh.md">中文版</a> · <a href="./README.md">English</a> · <a href="./docs/reference.md">Reference</a> · <a href="./CHANGELOG.md">Changelog</a>
|
||||
</p>
|
||||
|
||||
- [为什么选择 dws?](#why-dws)
|
||||
- [核心服务](#核心服务)
|
||||
- [安装](#安装)
|
||||
- [开始使用](#开始使用)
|
||||
- [快速开始](#快速开始)
|
||||
- [AI Agent Skills](#ai-agent-skills)
|
||||
- [高级用法](#高级用法)
|
||||
- [环境变量](#环境变量)
|
||||
- [退出码](#退出码)
|
||||
- [架构设计](#架构设计)
|
||||
- [开发指南](#开发指南)
|
||||
- [测试](#测试)
|
||||
- [更新日志](#更新日志)
|
||||
- [安全策略](#安全策略)
|
||||
- [贡献指南](#贡献指南)
|
||||
> [!IMPORTANT]
|
||||
> **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/i4/O1CN01Rijgk81gKqVSKMzdx_!!6000000004124-2-tps-654-644.png" alt="DingTalk Group QR Code" width="150"></a>
|
||||
|
||||
<h2 id="why-dws">为什么选择 dws?</h2>
|
||||
<details>
|
||||
<summary><strong>Table of Contents</strong></summary>
|
||||
|
||||
**为人类而设计** — 告别手写 API 调用。`dws` 为每个资源提供 `--help`,用 `--dry-run` 预览请求,支持表格/JSON/原始格式的结构化输出。
|
||||
- [Why dws?](#why-dws)
|
||||
- [Installation](#installation)
|
||||
- [Upgrade](#upgrade)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Quick Start](#quick-start)
|
||||
- [Using with Agents](#using-with-agents)
|
||||
- [Features](#features)
|
||||
- [Key Services](#key-services)
|
||||
- [Security by Design](#security-by-design)
|
||||
- [Reference & Docs](#reference--docs)
|
||||
- [Contributing](#contributing)
|
||||
|
||||
**为 AI Agent 而设计** — 每个响应都是结构化 JSON。配合内置的 agent skills,您的 LLM 无需自定义工具即可管理钉钉工作台。
|
||||
</details>
|
||||
|
||||
```bash
|
||||
# 搜索联系人
|
||||
dws contact user search --keyword "悟空"
|
||||
|
||||
# 创建待办事项
|
||||
dws todo task create --title "准备季度汇报材料" --executors "<userId>"
|
||||
---
|
||||
|
||||
# 预览操作但不执行
|
||||
dws todo task list --dry-run
|
||||
<h2 id="why-dws">Why dws?</h2>
|
||||
|
||||
# JSON 输出供 agent 使用
|
||||
dws contact user search --keyword "悟空" -f json
|
||||
```
|
||||
- **For humans** — `--help` for usage, `--dry-run` to preview requests, `-f table/json/raw` for output formats.
|
||||
- **For AI agents** — structured JSON responses + built-in Agent Skills, ready out of the box.
|
||||
- **For enterprise admins** — zero-trust architecture: OAuth device-flow auth + domain allowlisting + least-privilege scoping. **Not a single byte can bypass authentication and audit.**
|
||||
|
||||
## 核心服务
|
||||
## Installation
|
||||
|
||||
`dws` 通过统一的命令界面覆盖钉钉产品:
|
||||
|
||||
| 服务 | 命令 | 描述 |
|
||||
|---------|---------|-------------|
|
||||
| 通讯录 | `contact` | 通讯录 / 用户 / 部门 |
|
||||
| 群聊 | `chat` | 机器人消息 / Webhook / 机器人管理 |
|
||||
| 智能表格 | `aitable` | AI 表格操作 |
|
||||
| 日历 | `calendar` | 日历日程 / 会议室 / 闲忙 |
|
||||
| 待办 | `todo` | 待办任务管理 |
|
||||
| 审批 | `approval` | 审批流程 / 表单 / 实例 |
|
||||
| 考勤 | `attendance` | 考勤打卡 / 排班 / 统计 |
|
||||
| DING | `ding` | DING 消息 / 发送 / 撤回 |
|
||||
| 日志 | `report` | 日志 / 模版 / 统计 |
|
||||
| 工作台 | `workbench` | 工作台应用查询 |
|
||||
| 开发者文档 | `devdoc` | 开放平台文档搜索 |
|
||||
| 文档 | `doc` | 文档操作(即将推出) |
|
||||
| 邮箱 | `mail` | 邮件管理(即将推出) |
|
||||
| AI 听记 | `minutes` | AI 听记 / 会议纪要(即将推出) |
|
||||
| 钉盘 | `drive` | 云盘 / 文件存储(即将推出) |
|
||||
| 视频会议 | `conference` | 视频会议(即将推出) |
|
||||
| Teambition | `tb` | 项目管理(即将推出) |
|
||||
| AI 应用 | `aiapp` | AI 应用管理(即将推出) |
|
||||
| 直播 | `live` | 直播管理(即将推出) |
|
||||
| 技能市场 | `skill` | 技能搜索与下载(即将推出) |
|
||||
|
||||
运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看特定服务的命令。
|
||||
|
||||
## 安装
|
||||
|
||||
### 一键安装(推荐)
|
||||
|
||||
**macOS / Linux:**
|
||||
**macOS / Linux:**
|
||||
|
||||
```bash
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.sh | sh
|
||||
```
|
||||
|
||||
**Windows (PowerShell):**
|
||||
**Windows (PowerShell):**
|
||||
|
||||
```powershell
|
||||
irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.ps1 | iex
|
||||
```
|
||||
|
||||
> 自动检测操作系统和架构,从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 下载预编译二进制文件,并安装 Agent Skills 到 `~/.agents/skills/dws` — 无需 Go、Node.js 或其他依赖。大多数 AI Agent(Claude Code、Cursor、Windsurf 等)可自动发现 `.agents/skills/` 目录下的技能。
|
||||
<details>
|
||||
<summary>Other install methods</summary>
|
||||
|
||||
> [!TIP]
|
||||
> 二进制文件默认安装到 `~/.local/bin`。如果安装后找不到 `dws` 命令,请将其添加到 PATH:
|
||||
**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
|
||||
> export PATH="$HOME/.local/bin:$PATH"
|
||||
> xattr -d com.apple.quarantine /path/to/dws
|
||||
> ```
|
||||
> 将此行添加到 `~/.bashrc` 或 `~/.zshrc` 以永久生效。
|
||||
|
||||
### 预编译二进制文件(手动)
|
||||
|
||||
从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 下载适合您平台的最新二进制文件。
|
||||
|
||||
### 从源码构建
|
||||
**Build from source**:
|
||||
|
||||
```bash
|
||||
git clone https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli.git
|
||||
cd dingtalk-workspace-cli
|
||||
make build
|
||||
./dws version
|
||||
go build -o dws ./cmd # build to current directory
|
||||
cp dws ~/.local/bin/ # install to PATH
|
||||
```
|
||||
|
||||
这只会构建二进制文件。如需同时将 agent skills 安装到主目录:
|
||||
> Requires Go 1.25+. Use `make package` to cross-compile for all platforms (macOS / Linux / Windows x amd64 / arm64).
|
||||
|
||||
</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
|
||||
sh scripts/install.sh
|
||||
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
|
||||
```
|
||||
|
||||
这会检测本地源码目录,无需从 GitHub 下载即可安装二进制文件和 skills。
|
||||
<details>
|
||||
<summary><strong>How it works</strong></summary>
|
||||
|
||||
## 开始使用
|
||||
The upgrade process follows a two-phase atomic flow to ensure consistency:
|
||||
|
||||
### 步骤 1:创建钉钉应用
|
||||
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.).
|
||||
|
||||
进入 [开放平台应用开发后台](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app),在「企业内部应用 - 钉钉应用」点击右上角的**创建应用**,新建一个应用。
|
||||
A backup of the current version is automatically created before each upgrade. Use `dws upgrade --rollback` to restore the previous version if needed.
|
||||
|
||||
| Flag | Description |
|
||||
|------|-------------|
|
||||
| `--check` | Check for updates without installing |
|
||||
| `--list` | List all available versions with changelogs |
|
||||
| `--version` | Upgrade to a specific version (e.g. `v1.0.7`) |
|
||||
| `--rollback` | Rollback to the previous backed-up version |
|
||||
| `--force` | Force reinstall even if already on the latest version |
|
||||
| `--skip-skills` | Skip skill package update |
|
||||
| `-y` | Skip confirmation prompt |
|
||||
|
||||
</details>
|
||||
|
||||
## Getting Started
|
||||
|
||||
```bash
|
||||
dws auth login # browser opens automatically
|
||||
dws auth login --device # for headless environments (Docker, SSH, CI)
|
||||
```
|
||||
|
||||
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><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/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="创建应用" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="Apply for Access" width="600">
|
||||
</p>
|
||||
|
||||
### 步骤 2:配置重定向 URL
|
||||
</details>
|
||||
|
||||
创建应用后,进入应用内,点击**安全设置**。在「重定向 URL(回调设置)」里,输入 `http://127.0.0.1` 并保存。
|
||||
<details>
|
||||
<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/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="配置重定向URL" width="600">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI Access Management" width="600">
|
||||
</p>
|
||||
|
||||
### 步骤 3:发布应用
|
||||
</details>
|
||||
|
||||
点击「应用发布 - 版本管理与发布」,发布版本,使应用变成上线状态。
|
||||
<details>
|
||||
<summary><strong>Custom App mode (CI/CD, ISV integration)</strong></summary>
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="发布应用" width="600">
|
||||
</p>
|
||||
For enterprise-managed scenarios, create your own DingTalk app:
|
||||
|
||||
### 步骤 4:申请白名单
|
||||
|
||||
参照页面顶部的 [共创阶段说明](#important),加入钉钉 DWS 共创群完成白名单配置。
|
||||
|
||||
### 步骤 5:使用凭证登录
|
||||
|
||||
获取 Client ID(AppKey)和 Client Secret(AppSecret)后,可通过 CLI 参数指定:
|
||||
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>
|
||||
```
|
||||
|
||||
或者通过环境变量设置:
|
||||
Credentials are securely persisted after first login (Keychain). Subsequent runs auto-refresh tokens.
|
||||
|
||||
</details>
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
dws contact user search --keyword "engineering" # search contacts
|
||||
dws calendar event list # list calendar events
|
||||
dws todo task create --title "Quarterly report" --executors "<your-userId>" # create a todo (replace <your-userId>)
|
||||
dws todo task list --dry-run # preview without executing
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> CLI 参数优先级高于环境变量。这些凭证用于钉钉的 OAuth 设备流认证。
|
||||
## Using with Agents
|
||||
|
||||
### Token 加密
|
||||
dws is designed as an AI-native CLI. Complete [Installation](#installation) and [Getting Started](#getting-started) first, then configure your agent:
|
||||
|
||||
Token 使用 **PBKDF2(600,000 次迭代)+ AES-256-GCM** 加密存储,密钥由您的设备 MAC 地址生成。
|
||||
|
||||
## 快速开始
|
||||
### Agent Invocation Patterns
|
||||
|
||||
```bash
|
||||
dws auth login # 钉钉身份认证
|
||||
dws contact user search --keyword "悟空" # 搜索联系人
|
||||
dws calendar event list # 列出日历事件
|
||||
dws todo task create --title "准备季度汇报材料" --executors "<userId>" # 创建待办
|
||||
# Use --yes to skip confirmation prompts (required for agents)
|
||||
dws todo task create --title "Review PR" --executors "<your-userId>" --yes
|
||||
|
||||
# Use --dry-run to preview operations (safe execution)
|
||||
dws contact user search --keyword "engineering" --dry-run
|
||||
|
||||
# Use --jq to extract precisely (save tokens)
|
||||
dws contact user get-self --jq '.result[0].orgEmployeeModel | {name: .orgUserName, dept: .depts[0].deptName, userId}'
|
||||
```
|
||||
|
||||
## AI Agent Skills
|
||||
### Schema Discovery
|
||||
|
||||
仓库为每个支持的钉钉产品提供 agent skills(`SKILL.md` 文件)。
|
||||
|
||||
Skills 由[安装](#安装)脚本自动安装。如需单独将 skills 安装到现有项目:
|
||||
Agents don't need pre-built knowledge of every command. Use `dws schema` to dynamically discover capabilities:
|
||||
|
||||
```bash
|
||||
# macOS / Linux — 仅将 skills 安装到当前项目
|
||||
# Step 1: Discover all available products
|
||||
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
|
||||
|
||||
# Step 2: Inspect target tool's parameter 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
|
||||
```
|
||||
|
||||
### Agent Skills
|
||||
|
||||
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
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-skills.sh | sh
|
||||
```
|
||||
|
||||
一键安装器(`install.sh`)将 skills 安装到 `~/.agents/skills/dws`(主目录)。
|
||||
当您想要为特定项目仓库添加 skills 时,请使用 `install-skills.sh`,它会安装到 `./.agents/skills/dws`(当前工作目录)。
|
||||
> `install.sh` installs to `$HOME/.agents/skills/dws` (global); `install-skills.sh` installs to `./.agents/skills/dws` (current project).
|
||||
|
||||
> [!NOTE]
|
||||
> **主目录 vs. 项目 skills**:`install.sh` 将 skills 放在 `$HOME/.agents/skills/dws`。`install-skills.sh` 安装到**当前工作目录**(`./.agents/skills/dws`),适用于为特定项目仓库添加 skills。
|
||||
**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</summary>
|
||||
|
||||
Built-in pipeline engine that normalizes flag names, splits sticky arguments, and fuzzy-matches typos:
|
||||
|
||||
```bash
|
||||
# 表格(默认,适合人类阅读)
|
||||
dws contact user search --keyword "悟空" -f table
|
||||
# Naming convention auto-conversion (camelCase / snake_case / UPPER -> kebab-case)
|
||||
dws aitable record query --baseId BASE_ID --tableId TABLE_ID # auto-corrected to --base-id --table-id
|
||||
|
||||
# JSON(适合 agent 和管道处理)
|
||||
dws contact user search --keyword "悟空" -f json
|
||||
# Sticky argument splitting
|
||||
dws contact user search --keyword "engineering" --timeout30 # auto-split to --timeout 30
|
||||
|
||||
# 原始 API 响应
|
||||
dws contact user search --keyword "悟空" -f raw
|
||||
# Fuzzy flag name matching
|
||||
dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-id -> --table-id
|
||||
|
||||
# Value normalization (boolean / number / date / enum)
|
||||
# "yes" -> true, "1,000" -> 1000, "2024/03/29" -> "2024-03-29", "ACTIVE" -> "active"
|
||||
```
|
||||
|
||||
### 试运行
|
||||
| Agent Output | dws Auto-Corrects To |
|
||||
|-----------|--------------|
|
||||
| `--userId` | `--user-id` |
|
||||
| `--limit100` | `--limit 100` |
|
||||
| `--tabel-id` | `--table-id` |
|
||||
| `--USER-ID` | `--user-id` |
|
||||
| `--user_name` | `--user-name` |
|
||||
|
||||
预览 MCP 工具调用但不执行:
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>jq Filtering & Field Selection</strong> — fine-grained output control to reduce token consumption</summary>
|
||||
|
||||
```bash
|
||||
dws todo task list --dry-run
|
||||
# Built-in jq expressions
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --jq '.invocation.params'
|
||||
dws schema --jq '.products[] | {id, tools: (.tools | length)}'
|
||||
|
||||
# Return only specific fields
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocation,response
|
||||
```
|
||||
|
||||
### 输出到文件
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Schema Introspection</strong> — query parameter schemas before making calls</summary>
|
||||
|
||||
```bash
|
||||
dws contact user search --keyword "李明" -o result.json
|
||||
dws schema # list all products and tools
|
||||
dws schema aitable.query_records # view parameter schema
|
||||
dws schema aitable.query_records --jq '.tool.required' # view required fields
|
||||
dws schema --jq '.products[].id' # extract all product IDs
|
||||
```
|
||||
|
||||
### Shell 自动补全
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Pipe & File Input</strong> — read flag values from files or stdin</summary>
|
||||
|
||||
```bash
|
||||
# Bash
|
||||
dws completion bash > /etc/bash_completion.d/dws
|
||||
# Read message body from a file
|
||||
dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "Weekly Report" --text @report.md
|
||||
|
||||
# Zsh
|
||||
dws completion zsh > "${fpath[1]}/_dws"
|
||||
# Pipe content via stdin
|
||||
cat report.md | dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "Weekly Report"
|
||||
|
||||
# Fish
|
||||
dws completion fish > ~/.config/fish/completions/dws.fish
|
||||
# Read from stdin explicitly
|
||||
dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "Weekly Report" --text @-
|
||||
```
|
||||
|
||||
## 环境变量
|
||||
</details>
|
||||
|
||||
常用的运行时和开发覆盖项:
|
||||
## Key Services
|
||||
|
||||
| 变量 | 用途 |
|
||||
|---------|---------|
|
||||
| `DWS_CONFIG_DIR` | 覆盖默认配置目录 |
|
||||
| `DWS_SERVERS_URL` | 将服务发现指向自定义服务器注册端点 |
|
||||
| `DWS_CLIENT_ID` | OAuth client ID(钉钉 AppKey) |
|
||||
| `DWS_CLIENT_SECRET` | OAuth client secret(钉钉 AppSecret) |
|
||||
| `DWS_TRUSTED_DOMAINS` | Bearer token 允许发送的域名列表,逗号分隔(默认 `*.dingtalk.com`)。仅开发环境可设为 `*` |
|
||||
| `DWS_ALLOW_HTTP_ENDPOINTS` | 设为 `1` 允许对回环地址使用 HTTP(非 TLS),仅用于开发调试 |
|
||||
| Service | Command | Commands | Subcommands | Description |
|
||||
|---------|---------|:--------:|-------------|-------------|
|
||||
| Contact | `contact` | 6 | `user` `dept` | Search users by name/mobile, batch query, departments, current user profile |
|
||||
| Chat / IM | `chat` (alias `im`) | 20 | `message` `group` `search` `list-top-conversations` | User-identity send (group / 1-on-1 / open-dingtalk-id), Markdown + image, @mentions; read & search conversations (list, list-all, topic replies, by-sender, mentions, focused, unread, search, info, top / common groups); group CRUD + member management |
|
||||
| Bot | `chat bot` | 7 | `bot` `group` `message` `search` `create` `search-groups` | Bot create / search, search bot groups; bot-identity group & batch-1:1 messaging, Webhook, message recall; add bot to group |
|
||||
| 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` | 37 | `base` `table` `record` `field` `attachment` `template` `chart` `dashboard` `export` `import` `view` | Full CRUD for bases/tables/records/fields; charts/dashboards; data import/export; views; templates |
|
||||
| Doc | `doc` | 16 | `search` `list` `info` `read` `create` `update` `upload` `download` `folder` `block` `comment` | Search, read, create/update documents; block-level editing; file upload/download; comments |
|
||||
| Minutes | `minutes` | 22 | `list` `get` `update` `record` `hot-word` `mind-graph` `replace-text` `speaker` `upload` | List/search AI meeting transcripts; summaries, transcriptions, todos, mind-maps; recording control; speaker management, hot-words, file upload |
|
||||
| Workbench | `workbench` | 2 | `app` | Batch query app details |
|
||||
| DevDoc | `devdoc` | 1 | `article` | Search platform docs and error codes |
|
||||
|
||||
## 退出码
|
||||
> 152 commands across 14 products. Run `dws --help` for the full list, or `dws <service> --help` for subcommands.
|
||||
|
||||
| 退出码 | 类别 | 描述 |
|
||||
|--------|------|------|
|
||||
| 0 | 成功 | 命令执行成功 |
|
||||
| 1 | API | MCP 工具调用或上游 API 失败 |
|
||||
| 2 | 认证 | 身份认证或授权失败 |
|
||||
| 3 | 校验 | 输入参数、命令行标志或参数 schema 不匹配 |
|
||||
| 4 | 发现 | 服务发现、缓存或协议协商失败 |
|
||||
| 5 | 内部 | 未预期的内部错误 |
|
||||
<details>
|
||||
<summary>Coming soon</summary>
|
||||
|
||||
使用 `-f json` 时,错误响应包含结构化信息(`category`、`reason`、`hint`、`actions` 字段),便于机器消费。
|
||||
`mail` (email) · `drive` (cloud drive) · `conference` (video) · `tb` (Teambition) · `aiapp` (AI apps) · `live` (streaming) · `skill` (marketplace)
|
||||
|
||||
## 架构设计
|
||||
</details>
|
||||
|
||||
`dws` 使用 **发现驱动的管道** — 不硬编码任何产品命令:
|
||||
<h2 id="security-by-design">Security by Design</h2>
|
||||
|
||||
```
|
||||
Market Registry ──► Discovery ──► IR (规范化目录) ──► CLI (Cobra) ──► Transport (MCP JSON-RPC)
|
||||
│ │
|
||||
▼ ▼
|
||||
mcp.dingtalk.com 缓存(TTL + 过期降级)
|
||||
```
|
||||
`dws` treats security as a first-class architectural concern, not an afterthought. **Credentials never touch disk, tokens never leave trusted domains, permissions never exceed grants, operations never escape audit** — every API call must pass through DingTalk Open Platform's authentication and audit chain, no exceptions.
|
||||
|
||||
1. **Market** — 从 `mcp.dingtalk.com` 获取 MCP 服务注册表
|
||||
2. **Discovery** — 解析服务运行时能力,支持磁盘缓存和过期降级保证离线可用
|
||||
3. **IR** — 将服务规范化为统一的产品/工具目录
|
||||
4. **CLI** — 将目录挂载到 Cobra 命令树,映射 flag 到 MCP 输入参数
|
||||
5. **Transport** — 执行 MCP JSON-RPC 调用,支持重试、认证注入和响应大小限制
|
||||
<details>
|
||||
<summary><strong>For Developers</strong></summary>
|
||||
|
||||
使用 `-f json` 时,所有输出 — 成功、错误和元数据 — 都是结构化 JSON。
|
||||
| Mechanism | Details |
|
||||
|-----------|----------|
|
||||
| **Encrypted token storage** | **PBKDF2 + AES-256-GCM** encryption, keyed by device physical MAC address; cross-platform Keychain/DPAPI integration provides additional protection — tokens cannot be decrypted on another machine |
|
||||
| **Input security** | Path traversal protection (symlink resolution + working directory containment), CRLF injection blocking, Unicode visual spoofing filtering — prevents AI Agents from being tricked by malicious instructions |
|
||||
| **Domain allowlist** | `DWS_TRUSTED_DOMAINS` defaults to `*.dingtalk.com`; bearer tokens are never sent to non-allowlisted domains |
|
||||
| **HTTPS enforced** | All requests require TLS; HTTP only permitted for loopback during development |
|
||||
| **Dry-run preview** | `--dry-run` shows call parameters without executing, preventing accidental mutations |
|
||||
| **Zero credential persistence** | Client ID / Secret used in memory only — never written to config files or logs |
|
||||
|
||||
## 开发指南
|
||||
</details>
|
||||
|
||||
```bash
|
||||
make build # 开发构建
|
||||
make test # 单元测试
|
||||
make lint # 格式化 + lint 检查
|
||||
make package # 本地构建所有发布产物(goreleaser snapshot)
|
||||
make release # 通过 goreleaser 构建和发布
|
||||
make publish-homebrew-formula # 将 dist/homebrew/dingtalk-workspace-cli.rb 推送到 tap 仓库
|
||||
```
|
||||
<details>
|
||||
<summary><strong>For Enterprise Admins</strong></summary>
|
||||
|
||||
### 包管理器产物
|
||||
| Mechanism | Details |
|
||||
|-----------|---------|
|
||||
| **OAuth device-flow auth** | Users must authenticate through an admin-authorized DingTalk application |
|
||||
| **Least-privilege scoping** | CLI can only invoke APIs granted to the application — no privilege escalation |
|
||||
| **Allowlist gating** | Admin confirmation required during co-creation phase; self-service approval planned |
|
||||
| **Full-chain audit** | Every data read/write passes through the DingTalk Open Platform API — enterprise admins can trace complete call logs in real time; no anomalous operation can hide |
|
||||
|
||||
构建并验证本地包管理器产物:
|
||||
</details>
|
||||
|
||||
```bash
|
||||
make package # 生成所有平台归档、npm 资源、Homebrew formula
|
||||
./scripts/release/verify-package-managers.sh # 验证 dws 二进制文件和 skills 包含在内
|
||||
```
|
||||
<details>
|
||||
<summary><strong>For ISVs</strong></summary>
|
||||
|
||||
## 测试
|
||||
| Mechanism | Details |
|
||||
|-----------|---------|
|
||||
| **Tenant data isolation** | Operates under authorized app identity; cross-tenant access is impossible |
|
||||
| **Skill sandbox** | Agent Skills are Markdown documents (`SKILL.md`) — prompt descriptions only, no arbitrary code execution |
|
||||
| **Zero blind spots** | Every API call during ISV–dws skill orchestration is forced through DingTalk Open Platform authentication — full call chain is traceable with no bypass path |
|
||||
|
||||
### CLI 测试
|
||||
</details>
|
||||
|
||||
运行完整的 CLI 测试套件(单元测试、golden 测试和集成测试):
|
||||
> Found a vulnerability? Report via [GitHub Security Advisories](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/security/advisories/new). See [SECURITY.md](./SECURITY.md).
|
||||
|
||||
```bash
|
||||
bash test/scripts/run_all_tests.sh --jobs 8
|
||||
```
|
||||
## Reference & Docs
|
||||
|
||||
### 打包测试
|
||||
- [Reference](./docs/reference.md) — environment variables, exit codes, output formats, shell completion
|
||||
- [Architecture](./docs/architecture.md) — discovery-driven pipeline, IR, transport layer
|
||||
- [Changelog](./CHANGELOG.md) — release history and migration notes
|
||||
|
||||
运行打包契约测试和本地包管理器验证:
|
||||
## Contributing
|
||||
|
||||
```bash
|
||||
go test ./test/scripts/... -count=1
|
||||
make package
|
||||
./scripts/release/verify-package-managers.sh
|
||||
```
|
||||
See [CONTRIBUTING.md](./CONTRIBUTING.md) for build instructions, testing, and development workflow.
|
||||
|
||||
### Skills 测试
|
||||
|
||||
安装 skills 后,使用 [`test/skill_tests.md`](./test/skill_tests.md) 进行验证。将该文件中的测试提示输入您的 AI agent 并确认预期输出。
|
||||
|
||||
## 更新日志
|
||||
|
||||
参见 [CHANGELOG.md](./CHANGELOG.md) 了解版本历史和迁移说明。
|
||||
|
||||
## 安全策略
|
||||
|
||||
报告安全漏洞请参见 [SECURITY.md](./SECURITY.md)。
|
||||
|
||||
## 贡献指南
|
||||
|
||||
参见 [CONTRIBUTING.md](./CONTRIBUTING.md) 了解开发工作流和本地验证步骤。
|
||||
|
||||
## 许可证
|
||||
## License
|
||||
|
||||
Apache-2.0
|
||||
|
||||
-360
@@ -1,360 +0,0 @@
|
||||
<h1 align="center">DingTalk Workspace CLI (dws)</h1>
|
||||
|
||||
**One CLI for all of DingTalk Workspace — built for humans and AI agents.**<br>
|
||||
Access contacts, calendar, todos, attendance, AI tables and more with zero boilerplate, get structured JSON responses ready for automation, and leverage built-in Agent Skills for seamless AI integration.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i1/O1CN01oKAc2r28jOyyspcQt_!!6000000007968-2-tps-4096-1701.png" alt="DWS Product Overview" width="100%">
|
||||
</p>
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **Co-creation Phase**: This project accesses DingTalk enterprise data and requires enterprise admin authorization. We are currently in a gray-scale co-creation phase. Please join the DingTalk DWS co-creation group and provide the following materials to the official staff for whitelist configuration: ① Your DingTalk application's Client ID; ② Confirmation from the enterprise admin to enable access. Self-service approval by enterprise admins will be supported in the future.
|
||||
>
|
||||
> <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>
|
||||
|
||||
<p>
|
||||
<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.0.0-red" alt="v1.0.0"></a>
|
||||
</p>
|
||||
|
||||
## Contents
|
||||
|
||||
- [Why dws?](#why-dws)
|
||||
- [Key Services](#key-services)
|
||||
- [Installation](#installation)
|
||||
- [Getting Started](#getting-started)
|
||||
- [Quick Start](#quick-start)
|
||||
- [AI Agent Skills](#ai-agent-skills)
|
||||
- [Advanced Usage](#advanced-usage)
|
||||
- [Environment Variables](#environment-variables)
|
||||
- [Exit Codes](#exit-codes)
|
||||
- [Architecture](#architecture)
|
||||
- [Development](#development)
|
||||
- [Testing](#testing)
|
||||
- [Changelog](#changelog)
|
||||
- [Security](#security)
|
||||
- [Contributing](#contributing)
|
||||
|
||||
<h2 id="why-dws">Why dws?</h2>
|
||||
|
||||
**For humans** — stop writing raw API calls. `dws` gives you `--help` on every resource, `--dry-run` to preview requests, and structured output in table/JSON/raw formats.
|
||||
|
||||
**For AI agents** — every response is structured JSON. Pair it with the included agent skills and your LLM can manage DingTalk Workspace without custom tooling.
|
||||
|
||||
```bash
|
||||
# Search for a contact
|
||||
dws contact user search --keyword "Alice"
|
||||
|
||||
# Create a todo item
|
||||
dws todo task create --title "Prepare quarterly report" --executors "<userId>"
|
||||
|
||||
# Preview an operation without executing
|
||||
dws todo task list --dry-run
|
||||
|
||||
# JSON output for agent consumption
|
||||
dws contact user search --keyword "Alice" -f json
|
||||
```
|
||||
|
||||
## Key Services
|
||||
|
||||
`dws` covers DingTalk products through a unified command surface:
|
||||
|
||||
| Service | Command | Description |
|
||||
|---------|---------|-------------|
|
||||
| Contact | `contact` | Contacts / users / departments |
|
||||
| Chat | `chat` | Bot messaging / webhook / bot management |
|
||||
| Calendar | `calendar` | Calendar events / meeting rooms / free-busy |
|
||||
| Todo | `todo` | Todo task management |
|
||||
| Approval | `approval` | Approval processes / forms / instances |
|
||||
| Attendance | `attendance` | Attendance / shifts / statistics |
|
||||
| Ding | `ding` | DING messages / send / recall |
|
||||
| Report | `report` | Report / template / statistics |
|
||||
| AITable | `aitable` | AI table operations |
|
||||
| Workbench | `workbench` | Workbench app query |
|
||||
| DevDoc | `devdoc` | Open platform docs search |
|
||||
| Doc | `doc` | Document operations (coming soon) |
|
||||
| Mail | `mail` | Email management (coming soon) |
|
||||
| Minutes | `minutes` | AI meeting transcription (coming soon) |
|
||||
| Drive | `drive` | Cloud drive / file storage (coming soon) |
|
||||
| Conference | `conference` | Video conferencing (coming soon) |
|
||||
| Teambition | `tb` | Project management (coming soon) |
|
||||
| AI App | `aiapp` | AI application management (coming soon) |
|
||||
| Live | `live` | Live streaming (coming soon) |
|
||||
| Skill | `skill` | Skill marketplace (coming soon) |
|
||||
|
||||
Run `dws --help` for the complete list, or `dws <service> --help` for service-specific commands.
|
||||
|
||||
## Installation
|
||||
|
||||
### One-line install (recommended)
|
||||
|
||||
**macOS / Linux:**
|
||||
|
||||
```bash
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.sh | sh
|
||||
```
|
||||
|
||||
**Windows (PowerShell):**
|
||||
|
||||
```powershell
|
||||
irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.ps1 | iex
|
||||
```
|
||||
|
||||
> Auto-detects OS and architecture, downloads the pre-built binary from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases), and installs Agent Skills to `~/.agents/skills/dws` — no Go, Node.js, or other dependencies required. Most AI agents (Claude Code, Cursor, Windsurf, etc.) can discover skills from the `.agents/skills/` directory.
|
||||
|
||||
> [!TIP]
|
||||
> The binary is installed to `~/.local/bin` by default. If `dws` is not found after installation, add it to your PATH:
|
||||
> ```bash
|
||||
> export PATH="$HOME/.local/bin:$PATH"
|
||||
> ```
|
||||
> Add this line to your `~/.bashrc` or `~/.zshrc` to make it permanent.
|
||||
|
||||
### Pre-built binary (manual)
|
||||
|
||||
Download the latest binary for your platform from [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases).
|
||||
|
||||
### Build from source
|
||||
|
||||
```bash
|
||||
git clone https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli.git
|
||||
cd dingtalk-workspace-cli
|
||||
make build
|
||||
./dws version
|
||||
```
|
||||
|
||||
This builds the binary only. To also install agent skills into your home directory:
|
||||
|
||||
```bash
|
||||
sh scripts/install.sh
|
||||
```
|
||||
|
||||
This detects the local source checkout and installs both the binary and skills without downloading from GitHub.
|
||||
|
||||
## Getting Started
|
||||
|
||||
### Step 1: Create a DingTalk Application
|
||||
|
||||
Go to the [Open Platform App Development Console](https://open-dev.dingtalk.com/fe/app?hash=%23%2Fcorp%2Fapp#/corp/app). Under "Internal Enterprise Apps - DingTalk Apps", click **Create App** in the top right corner to create a new application.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01VIkwvV1a5NQzCIFO0_!!6000000003278-2-tps-2690-1462.png" alt="Create Application" width="600">
|
||||
</p>
|
||||
|
||||
### Step 2: Configure Redirect URL
|
||||
|
||||
After creating the app, go into the app settings and click **Security Settings**. In the "Redirect URL (Callback Settings)" section, enter `http://127.0.0.1` and save.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN017xQGWb1ycrAG0uxBO_!!6000000006600-2-tps-2000-1032.png" alt="Configure Redirect URL" width="600">
|
||||
</p>
|
||||
|
||||
### Step 3: Publish the Application
|
||||
|
||||
Click "App Release - Version Management & Release", publish a version to make the app go live.
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01WOLZFz244P46B3FPu_!!6000000007337-2-tps-2000-1100.png" alt="Publish Application" width="600">
|
||||
</p>
|
||||
|
||||
### Step 4: Request Whitelist Access
|
||||
|
||||
Refer to the [Co-creation Phase notice](#important) at the top of this page to join the DingTalk DWS co-creation group and complete whitelist configuration.
|
||||
|
||||
### Step 5: Login with Credentials
|
||||
|
||||
Once you have the AppKey and AppSecret, specify them via CLI flags:
|
||||
|
||||
```bash
|
||||
dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
|
||||
```
|
||||
|
||||
Alternatively, set via environment variables:
|
||||
|
||||
```bash
|
||||
export DWS_CLIENT_ID=<your-app-key>
|
||||
export DWS_CLIENT_SECRET=<your-app-secret>
|
||||
dws auth login
|
||||
```
|
||||
|
||||
> [!NOTE]
|
||||
> CLI flags take precedence over environment variables. These credentials are used for the OAuth device flow authentication with DingTalk.
|
||||
|
||||
### Token Encryption
|
||||
|
||||
Tokens are encrypted at rest using **PBKDF2 (600,000 iterations) + AES-256-GCM**, keyed by your device MAC address.
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
dws auth login # authenticate with DingTalk
|
||||
dws contact user search --keyword "Alice" # search contacts
|
||||
dws calendar event list # list calendar events
|
||||
dws todo task create --title "Prepare quarterly report" --executors "<userId>" # create a todo
|
||||
```
|
||||
|
||||
## AI Agent Skills
|
||||
|
||||
The repo ships agent skills (`SKILL.md` files) for every supported DingTalk product.
|
||||
|
||||
Skills are installed automatically by the [Installation](#installation) scripts. To install skills separately into an existing project:
|
||||
|
||||
```bash
|
||||
# macOS / Linux — install only skills into the current project
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-skills.sh | sh
|
||||
```
|
||||
|
||||
The one-line installer (`install.sh`) installs skills to `~/.agents/skills/dws` (home directory).
|
||||
Use `install-skills.sh` when you want to seed a specific project repository with `./.agents/skills/dws` (current working directory).
|
||||
|
||||
> [!NOTE]
|
||||
> **Home vs. project skills**: `install.sh` places skills in `$HOME/.agents/skills/dws`. `install-skills.sh` installs into the **current working directory** (`./.agents/skills/dws`), which is useful for seeding a specific project repository.
|
||||
|
||||
## Advanced Usage
|
||||
|
||||
### Output Formats
|
||||
|
||||
All commands support multiple output formats:
|
||||
|
||||
```bash
|
||||
# Table (default, human-friendly)
|
||||
dws contact user search --keyword "Alice" -f table
|
||||
|
||||
# JSON (for agents and piping)
|
||||
dws contact user search --keyword "Alice" -f json
|
||||
|
||||
# Raw API response
|
||||
dws contact user search --keyword "Alice" -f raw
|
||||
```
|
||||
|
||||
### Dry Run
|
||||
|
||||
Preview the MCP tool invocation without executing:
|
||||
|
||||
```bash
|
||||
dws todo task list --dry-run
|
||||
```
|
||||
|
||||
### Output to File
|
||||
|
||||
```bash
|
||||
dws contact user search --keyword "Alice" -o result.json
|
||||
```
|
||||
|
||||
### Shell Completion
|
||||
|
||||
```bash
|
||||
# Bash
|
||||
dws completion bash > /etc/bash_completion.d/dws
|
||||
|
||||
# Zsh
|
||||
dws completion zsh > "${fpath[1]}/_dws"
|
||||
|
||||
# Fish
|
||||
dws completion fish > ~/.config/fish/completions/dws.fish
|
||||
```
|
||||
|
||||
## Environment Variables
|
||||
|
||||
Common runtime and development overrides:
|
||||
|
||||
| Variable | Purpose |
|
||||
|---------|---------|
|
||||
| `DWS_CONFIG_DIR` | Overrides the default config directory |
|
||||
| `DWS_SERVERS_URL` | Points discovery at a custom server registry endpoint |
|
||||
| `DWS_CLIENT_ID` | OAuth client ID (DingTalk AppKey) |
|
||||
| `DWS_CLIENT_SECRET` | OAuth client secret (DingTalk AppSecret) |
|
||||
| `DWS_TRUSTED_DOMAINS` | Comma-separated list of trusted domains for bearer token injection (default: `*.dingtalk.com`). Set to `*` for development only |
|
||||
| `DWS_ALLOW_HTTP_ENDPOINTS` | Set to `1` to allow HTTP (non-TLS) for loopback addresses during development |
|
||||
|
||||
## Exit Codes
|
||||
|
||||
| Code | Category | Description |
|
||||
|------|----------|-------------|
|
||||
| 0 | Success | Command completed successfully |
|
||||
| 1 | API | MCP tool call or upstream API failure |
|
||||
| 2 | Auth | Authentication or authorization failure |
|
||||
| 3 | Validation | Invalid input, flags, or parameter schema mismatch |
|
||||
| 4 | Discovery | Server discovery, cache, or protocol negotiation failure |
|
||||
| 5 | Internal | Unexpected internal error |
|
||||
|
||||
When `-f json` is used, error responses include structured payloads with `category`, `reason`, `hint`, and optional `actions` fields for machine consumption.
|
||||
|
||||
## Architecture
|
||||
|
||||
`dws` uses a **discovery-driven pipeline** — no product commands are hardcoded:
|
||||
|
||||
```
|
||||
Market Registry ──► Discovery ──► IR (Canonical Catalog) ──► CLI (Cobra) ──► Transport (MCP JSON-RPC)
|
||||
│ │
|
||||
▼ ▼
|
||||
mcp.dingtalk.com Cache (TTL + stale fallback)
|
||||
```
|
||||
|
||||
1. **Market** — fetches the MCP server registry from `mcp.dingtalk.com`
|
||||
2. **Discovery** — resolves runtime server capabilities with disk cache and stale-fallback for offline resilience
|
||||
3. **IR** — normalizes servers into a canonical product/tool catalog
|
||||
4. **CLI** — mounts the catalog onto a Cobra command tree, maps flags to MCP input parameters
|
||||
5. **Transport** — executes MCP JSON-RPC calls with retries, auth injection, and response size limits
|
||||
|
||||
All output — success, errors, and metadata — is structured JSON when using `-f json`.
|
||||
|
||||
## Development
|
||||
|
||||
```bash
|
||||
make build # dev build
|
||||
make test # unit tests
|
||||
make lint # formatting + lint checks
|
||||
make package # build all release artifacts locally (goreleaser snapshot)
|
||||
make release # build and publish a release via goreleaser
|
||||
make publish-homebrew-formula # push dist/homebrew/dingtalk-workspace-cli.rb to a tap repo
|
||||
```
|
||||
|
||||
### Package Manager Artifacts
|
||||
|
||||
Build and verify local package-manager artifacts:
|
||||
|
||||
```bash
|
||||
make package # generates all platform archives, npm assets, Homebrew formulas
|
||||
./scripts/release/verify-package-managers.sh # verifies dws binary + skills are included
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
### CLI Tests
|
||||
|
||||
Run the full CLI test suite (unit, golden, and integration tests):
|
||||
|
||||
```bash
|
||||
bash test/scripts/run_all_tests.sh --jobs 8
|
||||
```
|
||||
|
||||
### Packaging Tests
|
||||
|
||||
Run packaging contract tests and local package-manager verification:
|
||||
|
||||
```bash
|
||||
go test ./test/scripts/... -count=1
|
||||
make package
|
||||
./scripts/release/verify-package-managers.sh
|
||||
```
|
||||
|
||||
### Skill Tests
|
||||
|
||||
After installing the skills, use [`test/skill_tests.md`](./test/skill_tests.md) to verify them. Feed the test prompts from that file to your AI agent and confirm the expected outputs.
|
||||
|
||||
## Changelog
|
||||
|
||||
See [CHANGELOG.md](./CHANGELOG.md) for release history and migration notes.
|
||||
|
||||
## Security
|
||||
|
||||
To report a vulnerability, see [SECURITY.md](./SECURITY.md).
|
||||
|
||||
## Contributing
|
||||
|
||||
See [CONTRIBUTING.md](./CONTRIBUTING.md) for development workflow and local verification steps.
|
||||
|
||||
## License
|
||||
|
||||
Apache-2.0
|
||||
+431
@@ -0,0 +1,431 @@
|
||||
<h1 align="center">DingTalk Workspace CLI (dws)</h1>
|
||||
|
||||
<p align="center"><code>dws</code> — 钉钉工作台命令行工具,为人类和 AI Agent 而生。</p>
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i1/O1CN01oKAc2r28jOyyspcQt_!!6000000007968-2-tps-4096-1701.png" alt="DWS Product Overview" width="100%">
|
||||
</p>
|
||||
|
||||
<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/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>
|
||||
|
||||
<p align="center">
|
||||
<a href="./README_zh.md">中文版</a> · <a href="./README.md">English</a> · <a href="./docs/reference.md">参考手册</a> · <a href="./CHANGELOG.md">更新日志</a>
|
||||
</p>
|
||||
|
||||
> [!IMPORTANT]
|
||||
> **共创阶段**:本项目涉及钉钉企业数据访问,需企业管理员授权后方可使用。欢迎加入钉钉 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/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-中使用)
|
||||
- [功能特性](#功能特性)
|
||||
- [核心服务](#核心服务)
|
||||
- [安全设计](#安全设计)
|
||||
- [参考与文档](#参考与文档)
|
||||
- [贡献指南](#贡献指南)
|
||||
|
||||
</details>
|
||||
|
||||
|
||||
---
|
||||
|
||||
<h2 id="why-dws">为什么选择 dws?</h2>
|
||||
|
||||
- **为人类而设计** — `--help` 查看用法,`--dry-run` 预览请求,`-f table/json/raw` 切换格式。
|
||||
- **为 AI Agent 而设计** — 结构化 JSON 响应 + 内置 Agent Skills,开箱即用。
|
||||
- **为企业管理员而设计** — 零信任架构:OAuth 设备流认证 + 域名白名单 + 权限最小化。**没有一个字节能绕过安全鉴权和审计。**
|
||||
|
||||
## 安装
|
||||
|
||||
**macOS / Linux:**
|
||||
|
||||
```bash
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.sh | sh
|
||||
```
|
||||
|
||||
**Windows(PowerShell):**
|
||||
|
||||
```powershell
|
||||
irm https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install.ps1 | iex
|
||||
```
|
||||
|
||||
<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
|
||||
git clone https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli.git
|
||||
cd dingtalk-workspace-cli
|
||||
go build -o dws ./cmd # 编译到当前目录
|
||||
cp dws ~/.local/bin/ # 安装到 PATH
|
||||
```
|
||||
|
||||
> 需要 Go 1.25+。也可以用 `make package` 构建所有平台产物(macOS / Linux / Windows × amd64 / arm64)。
|
||||
|
||||
</details>
|
||||
|
||||
## 升级
|
||||
|
||||
> 需要 **v1.0.7** 及以上版本。更早版本请重新执行[安装脚本](#安装)进行升级。
|
||||
|
||||
dws 内置自升级能力,直接从 [GitHub Releases](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases) 拉取更新,支持 SHA256 完整性校验和自动备份。
|
||||
|
||||
```bash
|
||||
dws upgrade # 交互式升级到最新版本
|
||||
dws upgrade --check # 仅检查是否有新版本
|
||||
dws upgrade --list # 列出所有可用版本
|
||||
dws upgrade --version v1.0.7 # 升级到指定版本
|
||||
dws upgrade --rollback # 回滚到上一版本
|
||||
dws upgrade -y # 跳过确认直接升级
|
||||
```
|
||||
|
||||
<details>
|
||||
<summary><strong>工作原理</strong></summary>
|
||||
|
||||
升级过程采用两阶段原子流程,确保一致性:
|
||||
|
||||
1. **准备阶段** — 将平台对应的二进制文件和技能包下载到临时目录,校验 SHA256 校验和,解压并验证所有文件。任何步骤失败则立即中止,不会修改现有安装。
|
||||
2. **执行阶段** — 仅在所有准备工作成功后,替换二进制文件并将技能包安装到所有已检测到的 Agent 目录(`~/.agents/skills/dws`、`~/.claude/skills/dws`、`~/.cursor/skills/dws` 等)。
|
||||
|
||||
每次升级前自动备份当前版本,可通过 `dws upgrade --rollback` 随时回滚。
|
||||
|
||||
| Flag | 说明 |
|
||||
|------|------|
|
||||
| `--check` | 仅检查更新,不安装 |
|
||||
| `--list` | 列出所有可用版本及更新日志 |
|
||||
| `--version` | 升级到指定版本(如 `v1.0.7`) |
|
||||
| `--rollback` | 回滚到上一个备份版本 |
|
||||
| `--force` | 强制重新安装,即使已是最新版本 |
|
||||
| `--skip-skills` | 跳过技能包更新 |
|
||||
| `-y` | 跳过确认提示 |
|
||||
|
||||
</details>
|
||||
|
||||
## 开始使用
|
||||
|
||||
```bash
|
||||
dws auth login # 自动唤起浏览器
|
||||
dws auth login --device # 无浏览器环境(Docker、SSH、CI)
|
||||
```
|
||||
|
||||
选择组织并授权即可。
|
||||
|
||||
> 如果组织尚未开启 CLI 访问权限,系统会引导你向管理员发送申请。审批通过后重新执行 `dws auth login` 即可。
|
||||
|
||||
<details>
|
||||
<summary><strong>组织未开启 CLI 访问权限?</strong></summary>
|
||||
|
||||
1. 选择组织后,点击「立即申请」通知管理员
|
||||
2. 管理员收到申请卡片,一键审批
|
||||
3. 审批通过后,重新执行 `dws auth login`
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i2/O1CN01wtsYuQ1CTbboVTlsD_!!6000000000082-2-tps-2696-1544.png" alt="申请权限" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>管理员:为组织开启 CLI 访问权限</strong></summary>
|
||||
|
||||
进入 [开发者平台](https://open-dev.dingtalk.com) →「CLI 访问管理」→ 开启。
|
||||
|
||||
<p align="center">
|
||||
<img src="https://img.alicdn.com/imgextra/i4/O1CN01M8K7Wj1rZ0WikrZby_!!6000000005644-2-tps-2940-1596.png" alt="CLI访问管理" width="600">
|
||||
</p>
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>自建应用模式(CI/CD、ISV 集成)</strong></summary>
|
||||
|
||||
企业自主管控场景,可创建自有钉钉应用:
|
||||
|
||||
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。
|
||||
|
||||
</details>
|
||||
|
||||
## 快速开始
|
||||
|
||||
```bash
|
||||
dws contact user search --keyword "悟空" # 搜索联系人
|
||||
dws calendar event list # 查看日历日程
|
||||
dws todo task create --title "季度汇报" --executors "<your-userId>" # 创建待办(请替换为真实 userId)
|
||||
dws todo task list --dry-run # 预览操作但不执行
|
||||
```
|
||||
|
||||
## 在 Agent 中使用
|
||||
|
||||
dws 是为 AI Agent 设计的 CLI 工具。请先完成[安装](#安装)和[开始使用](#开始使用),然后配置 Agent 环境:
|
||||
|
||||
### Agent 调用模式
|
||||
|
||||
```bash
|
||||
# 使用 --yes 跳过确认提示(Agent 必须)
|
||||
dws todo task create --title "Review PR" --executors "<your-userId>" --yes
|
||||
|
||||
# 使用 --dry-run 预览操作(安全执行)
|
||||
dws contact user search --keyword "张三" --dry-run
|
||||
|
||||
# 使用 --jq 精确提取(节省 token)
|
||||
dws contact user get-self --jq '.result[0].orgEmployeeModel | {name: .orgUserName, dept: .depts[0].deptName, userId}'
|
||||
```
|
||||
|
||||
### Schema 发现
|
||||
|
||||
Agent 无需预置所有命令知识,通过 `dws schema` 动态发现可用能力:
|
||||
|
||||
```bash
|
||||
# 第一步:发现所有可用产品
|
||||
dws schema --jq '.products[] | {id, tool_count: (.tools | length)}'
|
||||
|
||||
# 第二步:查看目标工具的参数结构
|
||||
dws schema aitable.query_records --jq '.tool.parameters'
|
||||
|
||||
# 第三步:构造正确的调用
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --limit 10
|
||||
```
|
||||
|
||||
### Agent Skills
|
||||
|
||||
仓库内置完整的 Agent Skill 体系(`skills/`),安装后 Claude Code / Cursor 等 AI 工具可通过自然语言直接操作钉钉:
|
||||
|
||||
```bash
|
||||
# 安装 skills 到当前项目
|
||||
curl -fsSL https://raw.githubusercontent.com/DingTalk-Real-AI/dingtalk-workspace-cli/main/scripts/install-skills.sh | sh
|
||||
```
|
||||
|
||||
> `install.sh` 安装到 `$HOME/.agents/skills/dws`(全局);`install-skills.sh` 安装到 `./.agents/skills/dws`(当前项目)。
|
||||
|
||||
**包含内容:**
|
||||
|
||||
| 组件 | 路径 | 说明 |
|
||||
|------|------|------|
|
||||
| 主 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 模型常见的参数错误</summary>
|
||||
|
||||
内置 Pipeline 纠错引擎,支持命名风格转换、粘连参数拆分、拼写模糊匹配:
|
||||
|
||||
```bash
|
||||
# 命名风格自动转换 (camelCase / snake_case / UPPER → kebab-case)
|
||||
dws aitable record query --baseId BASE_ID --tableId TABLE_ID # 自动纠正为 --base-id --table-id
|
||||
|
||||
# 粘连参数自动拆分
|
||||
dws contact user search --keyword "张三" --timeout30 # 自动拆分为 --timeout 30
|
||||
|
||||
# 拼写错误模糊匹配
|
||||
dws aitable record query --base-id BASE_ID --tabel-id TABLE_ID # --tabel-id → --table-id
|
||||
|
||||
# 参数值归一化 (布尔 / 数字 / 日期 / 枚举)
|
||||
# "yes" → true, "1,000" → 1000, "2024/03/29" → "2024-03-29", "ACTIVE" → "active"
|
||||
```
|
||||
|
||||
| Agent 输出 | dws 自动纠正为 |
|
||||
|-----------|--------------|
|
||||
| `--userId` | `--user-id` |
|
||||
| `--limit100` | `--limit 100` |
|
||||
| `--tabel-id` | `--table-id` |
|
||||
| `--USER-ID` | `--user-id` |
|
||||
| `--user_name` | `--user-name` |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>jq 过滤 & 字段筛选</strong> — 精确控制输出,减少 token 消耗</summary>
|
||||
|
||||
```bash
|
||||
# 内置 jq 表达式
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --jq '.invocation.params'
|
||||
dws schema --jq '.products[] | {id, tools: (.tools | length)}'
|
||||
|
||||
# 只返回指定字段
|
||||
dws aitable record query --base-id BASE_ID --table-id TABLE_ID --fields invocation,response
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>Schema 自省</strong> — 调用前查询任意工具的参数结构</summary>
|
||||
|
||||
```bash
|
||||
dws schema # 列出所有产品和工具
|
||||
dws schema aitable.query_records # 查看参数 Schema
|
||||
dws schema aitable.query_records --jq '.tool.required' # 查看必填字段
|
||||
dws schema --jq '.products[].id' # 提取所有产品 ID
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>管道 & 文件输入</strong> — 从文件或 stdin 读取 flag 值</summary>
|
||||
|
||||
```bash
|
||||
# 从文件读取消息内容
|
||||
dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "周报" --text @report.md
|
||||
|
||||
# 通过管道传入内容
|
||||
cat report.md | dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "周报"
|
||||
|
||||
# 显式从 stdin 读取
|
||||
dws chat message send-by-bot --robot-code BOT_CODE --group GROUP_ID \
|
||||
--title "周报" --text @-
|
||||
```
|
||||
|
||||
</details>
|
||||
|
||||
## 核心服务
|
||||
|
||||
| 服务 | 命令 | 命令数 | 子命令 | 描述 |
|
||||
|------|------|:------:|--------|------|
|
||||
| 通讯录 | `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` | 搜索开放平台文档与错误码 |
|
||||
|
||||
> 12 个产品,86 个命令。运行 `dws --help` 查看完整列表,或 `dws <service> --help` 查看子命令。
|
||||
|
||||
<details>
|
||||
<summary>即将推出</summary>
|
||||
|
||||
`doc`(文档)· `mail`(邮箱)· `minutes`(AI 听记)· `drive`(钉盘)· `conference`(视频会议)· `tb`(Teambition)· `aiapp`(AI 应用)· `live`(直播)· `skill`(技能市场)
|
||||
|
||||
</details>
|
||||
|
||||
## 安全设计
|
||||
|
||||
`dws` 从架构层面将安全作为一等公民,而非事后补丁。**凭证不落盘、Token 不出域、权限不越界、操作不脱审** — 每一次 API 调用都必须经过钉钉开放平台的鉴权和审计链路,无例外。
|
||||
|
||||
<details>
|
||||
<summary><strong>开发者安全机制</strong></summary>
|
||||
|
||||
| 机制 | 说明 |
|
||||
|------|------|
|
||||
| **Token 加密存储** | **PBKDF2(600,000 次迭代 + SHA-256)+ AES-256-GCM** 加密,密钥绑定设备物理 MAC 地址;macOS 集成系统 Keychain、Windows 集成 DPAPI 提供额外保护,跨设备无法解密 |
|
||||
| **输入安全防护** | 路径遍历防护(符号链接解析 + 工作目录约束)、CRLF 注入拦截、Unicode 视觉欺骗字符过滤,防止 AI Agent 被恶意指令诱导 |
|
||||
| **域名白名单** | `DWS_TRUSTED_DOMAINS` 默认仅信任 `*.dingtalk.com`,Bearer Token 不会发送到非白名单域 |
|
||||
| **并发安全** | 双层锁机制(进程内 + 跨进程文件锁)保障 Token 刷新原子性,适配高并发 MCP Server 场景 |
|
||||
| **数据完整性** | 所有配置写入采用原子操作(temp + fsync + rename),确保进程中断时数据不损坏 |
|
||||
| **HTTPS 强制** | 除 loopback 开发调试外,所有请求强制 TLS |
|
||||
| **Dry-run 预览** | `--dry-run` 展示调用参数但不执行,防止误操作生产数据 |
|
||||
| **凭证零落盘** | Client ID / Secret 仅在内存中使用,不写入配置文件或日志 |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>企业管理员安全机制</strong></summary>
|
||||
|
||||
| 机制 | 说明 |
|
||||
|------|------|
|
||||
| **OAuth 设备流认证** | 用户必须通过管理员授权的钉钉应用认证,未授权应用无法获取 Token |
|
||||
| **权限最小化** | CLI 仅能调用管理员授予该应用的 API 权限范围,无法越权 |
|
||||
| **白名单准入** | 共创阶段需管理员主动确认开通,后续支持自助审批 |
|
||||
| **操作全链路审计** | 每一次数据读写都经过钉钉开放平台 API,企业管理员可在管理后台实时追溯完整调用日志,任何异常操作无处隐藏 |
|
||||
|
||||
</details>
|
||||
|
||||
<details>
|
||||
<summary><strong>ISV / 企业服务商安全机制</strong></summary>
|
||||
|
||||
| 机制 | 说明 |
|
||||
|------|------|
|
||||
| **租户数据隔离** | 以已授权应用身份调用 API,不同租户数据严格隔离 |
|
||||
| **Skill 沙箱** | Agent Skills 是 Markdown 文档(`SKILL.md`),仅提供 prompt 描述,不执行任意代码 |
|
||||
| **集成链路零盲区** | ISV Skill 与 dws Skill 联调时,每一次 API 调用都强制经过钉钉开放平台鉴权,完整调用链路可追溯,不存在绕过审计的旁路 |
|
||||
|
||||
</details>
|
||||
|
||||
> 发现安全漏洞?请通过 [GitHub Security Advisories](https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/security/advisories/new) 报告,详见 [SECURITY.md](./SECURITY.md)。
|
||||
|
||||
## 参考与文档
|
||||
|
||||
- [参考手册](./docs/reference.md) — 环境变量、退出码、输出格式、Shell 补全
|
||||
- [架构设计](./docs/architecture.md) — 发现驱动管道、IR、Transport 层
|
||||
- [更新日志](./CHANGELOG.md) — 版本历史与迁移说明
|
||||
|
||||
## 贡献指南
|
||||
|
||||
参见 [CONTRIBUTING.md](./CONTRIBUTING.md) 了解构建、测试和开发工作流。
|
||||
|
||||
## 许可证
|
||||
|
||||
Apache-2.0
|
||||
@@ -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 = {
|
||||
|
||||
@@ -1,8 +1,24 @@
|
||||
{
|
||||
"name": "dingtalk-workspace-cli",
|
||||
"version": "__VERSION__",
|
||||
"description": "DingTalk Workspace CLI",
|
||||
"description": "DingTalk Workspace CLI - AI-powered productivity tools",
|
||||
"license": "Apache-2.0",
|
||||
"repository": {
|
||||
"type": "git",
|
||||
"url": "https://github.com/open-dingtalk/dingtalk-workspace-cli.git"
|
||||
},
|
||||
"homepage": "https://github.com/open-dingtalk/dingtalk-workspace-cli",
|
||||
"bugs": {
|
||||
"url": "https://github.com/open-dingtalk/dingtalk-workspace-cli/issues"
|
||||
},
|
||||
"keywords": [
|
||||
"dingtalk",
|
||||
"dws",
|
||||
"cli",
|
||||
"workspace",
|
||||
"ai",
|
||||
"productivity"
|
||||
],
|
||||
"bin": {
|
||||
"dws": "./bin/dws.js"
|
||||
},
|
||||
@@ -14,5 +30,8 @@
|
||||
"bin",
|
||||
"install.js",
|
||||
"README.md"
|
||||
]
|
||||
],
|
||||
"engines": {
|
||||
"node": ">=16"
|
||||
}
|
||||
}
|
||||
|
||||
Vendored
BIN
Binary file not shown.
@@ -0,0 +1,60 @@
|
||||
# Reference / 参考手册
|
||||
|
||||
## Environment Variables / 环境变量
|
||||
|
||||
| Variable | Purpose / 用途 |
|
||||
|---------|---------|
|
||||
| `DWS_CONFIG_DIR` | Override default config directory / 覆盖默认配置目录 |
|
||||
| `DWS_SERVERS_URL` | Point discovery at a custom server registry endpoint / 将服务发现指向自定义端点 |
|
||||
| `DWS_CLIENT_ID` | OAuth client ID (DingTalk AppKey) |
|
||||
| `DWS_CLIENT_SECRET` | OAuth client secret (DingTalk AppSecret) |
|
||||
| `DWS_TRUSTED_DOMAINS` | Comma-separated trusted domains for bearer token (default: `*.dingtalk.com`). `*` for dev only / Bearer token 允许发送的域名白名单,默认 `*.dingtalk.com`,仅开发环境可设为 `*` |
|
||||
| `DWS_ALLOW_HTTP_ENDPOINTS` | Set `1` to allow HTTP for loopback during dev / 设为 `1` 允许回环地址 HTTP,仅用于开发调试 |
|
||||
|
||||
## Exit Codes / 退出码
|
||||
|
||||
| Code | Category | Description / 描述 |
|
||||
|------|----------|-------------|
|
||||
| 0 | Success | Command completed successfully / 命令执行成功 |
|
||||
| 1 | API | MCP tool call or upstream API failure / MCP 工具调用或上游 API 失败 |
|
||||
| 2 | Auth | Authentication or authorization failure / 身份认证或授权失败 |
|
||||
| 3 | Validation | Invalid input, flags, or parameter schema mismatch / 输入参数校验失败 |
|
||||
| 4 | Discovery | Server discovery, cache, or protocol negotiation failure / 服务发现失败 |
|
||||
| 5 | Internal | Unexpected internal error / 未预期的内部错误 |
|
||||
|
||||
With `-f json`, error responses include structured payloads: `category`, `reason`, `hint`, `actions`.
|
||||
|
||||
使用 `-f json` 时,错误响应包含结构化字段:`category`、`reason`、`hint`、`actions`。
|
||||
|
||||
## Output Formats / 输出格式
|
||||
|
||||
```bash
|
||||
dws contact user search --keyword "Alice" -f table # Table (default, human-friendly / 表格,默认)
|
||||
dws contact user search --keyword "Alice" -f json # JSON (for agents and piping / 适合 agent)
|
||||
dws contact user search --keyword "Alice" -f raw # Raw API response / 原始响应
|
||||
```
|
||||
|
||||
## Dry Run / 试运行
|
||||
|
||||
```bash
|
||||
dws todo task list --dry-run # Preview MCP call without executing / 预览但不执行
|
||||
```
|
||||
|
||||
## Output to File / 输出到文件
|
||||
|
||||
```bash
|
||||
dws contact user search --keyword "Alice" -o result.json
|
||||
```
|
||||
|
||||
## Shell Completion / 自动补全
|
||||
|
||||
```bash
|
||||
# Bash
|
||||
dws completion bash > /etc/bash_completion.d/dws
|
||||
|
||||
# Zsh
|
||||
dws completion zsh > "${fpath[1]}/_dws"
|
||||
|
||||
# Fish
|
||||
dws completion fish > ~/.config/fish/completions/dws.fish
|
||||
```
|
||||
@@ -4,13 +4,19 @@ 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
|
||||
golang.org/x/sys v0.42.0
|
||||
golang.org/x/text v0.35.0
|
||||
)
|
||||
|
||||
require (
|
||||
github.com/danieljoos/wincred v1.2.3 // indirect
|
||||
github.com/godbus/dbus/v5 v5.2.2 // 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
|
||||
)
|
||||
|
||||
@@ -1,18 +1,38 @@
|
||||
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
|
||||
github.com/danieljoos/wincred v1.2.3 h1:v7dZC2x32Ut3nEfRH+vhoZGvN72+dQ/snVXo/vMFLdQ=
|
||||
github.com/danieljoos/wincred v1.2.3/go.mod h1:6qqX0WNrS4RzPZ1tnroDzq9kY3fu1KwE7MRLQK4X0bs=
|
||||
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
|
||||
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
|
||||
github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM=
|
||||
github.com/fatih/color v1.18.0/go.mod h1:4FelSpRwEGDpQ12mAdzqdOukCy4u8WUtOY6lkT/6HfU=
|
||||
github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ=
|
||||
github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c=
|
||||
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
|
||||
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
|
||||
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
|
||||
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
|
||||
github.com/itchyny/gojq v0.12.18 h1:gFGHyt/MLbG9n6dqnvlliiya2TaMMh6FFaR2b1H6Drc=
|
||||
github.com/itchyny/gojq v0.12.18/go.mod h1:4hPoZ/3lN9fDL1D+aK7DY1f39XZpY9+1Xpjz8atrEkg=
|
||||
github.com/itchyny/timefmt-go v0.1.7 h1:xyftit9Tbw+Dc/huSSPJaEmX1TVL8lw5vxjJLK4GMMA=
|
||||
github.com/itchyny/timefmt-go v0.1.7/go.mod h1:5E46Q+zj7vbTgWY8o5YkMeYb4I6GeWLFnetPy5oBrAI=
|
||||
github.com/mattn/go-colorable v0.1.13 h1:fFA4WZxdEF4tXPZVKMLwD8oUnCTTo08duU7wxecdEvA=
|
||||
github.com/mattn/go-colorable v0.1.13/go.mod h1:7S9/ev0klgBDR4GtXTXX8a3vIGJpMovkB8vQcUbaXHg=
|
||||
github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/yFXSvRLM=
|
||||
github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY=
|
||||
github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y=
|
||||
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
|
||||
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
|
||||
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
|
||||
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
|
||||
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
|
||||
github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY=
|
||||
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
|
||||
github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
|
||||
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
|
||||
github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
|
||||
github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
|
||||
github.com/zalando/go-keyring v0.2.8 h1:6sD/Ucpl7jNq10rM2pgqTs0sZ9V3qMrqfIIy5YPccHs=
|
||||
github.com/zalando/go-keyring v0.2.8/go.mod h1:tsMo+VpRq5NGyKfxoBVjCuMrG47yj8cmakZDO5QGii0=
|
||||
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
|
||||
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
|
||||
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
+173
-58
@@ -15,15 +15,18 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"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"
|
||||
)
|
||||
|
||||
@@ -46,11 +49,12 @@ func buildAuthCommand() *cobra.Command {
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(newAuthLoginCommand())
|
||||
if !edition.Get().HideAuthLogin {
|
||||
cmd.AddCommand(newAuthLoginCommand())
|
||||
}
|
||||
cmd.AddCommand(
|
||||
newAuthLogoutCommand(),
|
||||
newAuthStatusCommand(),
|
||||
newAuthImportCommand(),
|
||||
newAuthExchangeCommand(),
|
||||
newAuthResetCommand(),
|
||||
)
|
||||
@@ -59,8 +63,23 @@ func buildAuthCommand() *cobra.Command {
|
||||
|
||||
func newAuthLoginCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "login",
|
||||
Short: "登录钉钉(自动刷新 token,必要时扫码)",
|
||||
Use: "login",
|
||||
Short: "登录钉钉(自动刷新 token,必要时扫码)",
|
||||
Long: `登录钉钉并获取认证凭证。
|
||||
|
||||
支持的登录方式:
|
||||
- OAuth 设备流 (默认): 通过钉钉扫码授权登录
|
||||
- 直接提供 Token: 通过 --token 参数传入已有 token
|
||||
|
||||
不支持的登录方式:
|
||||
- 邮箱/密码登录
|
||||
- 手机号/验证码登录
|
||||
- 应用凭证 (AppKey/AppSecret) 直接登录
|
||||
|
||||
示例:
|
||||
dws auth login # 扫码登录
|
||||
dws auth login --force # 强制重新登录 (忽略缓存 token)
|
||||
dws auth login --token xxx # 使用指定 token`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cfg, err := resolveAuthLoginConfig(cmd)
|
||||
@@ -102,9 +121,18 @@ func newAuthLoginCommand() *cobra.Command {
|
||||
}
|
||||
}
|
||||
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
|
||||
// Check if JSON output is requested
|
||||
format, _ := cmd.Root().PersistentFlags().GetString("format")
|
||||
if strings.EqualFold(strings.TrimSpace(format), "json") {
|
||||
return writeAuthLoginJSON(w, tokenData, cfg.Force)
|
||||
}
|
||||
|
||||
// Default table output
|
||||
fmt.Fprintln(w)
|
||||
if !cfg.Device && tokenData != nil && tokenData.IsAccessTokenValid() && !cfg.Force {
|
||||
fmt.Fprintf(w, "[OK] Token 有效,无需重新登录\n")
|
||||
@@ -130,15 +158,23 @@ func newAuthLoginCommand() *cobra.Command {
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("token", "", "Access token")
|
||||
cmd.Flags().Bool("device", false, "Use device authorization flow (compatibility flag)")
|
||||
cmd.Flags().Bool("force", false, "Force interactive login flow (compatibility flag)")
|
||||
cmd.Flags().String("redirect-url", "", "Loopback redirect URL compatibility flag")
|
||||
cmd.Flags().Bool("device", false, "Use device authorization flow")
|
||||
cmd.Flags().Bool("force", false, "Force interactive login (ignore cached token)")
|
||||
// Hidden compatibility flags
|
||||
cmd.Flags().String("redirect-url", "", "Loopback redirect URL")
|
||||
cmd.Flags().String("scopes", "", "Space-separated DingTalk OAuth scopes")
|
||||
cmd.Flags().String("authorize-url", "", "Override DingTalk authorization URL")
|
||||
cmd.Flags().String("token-url", "", "Override DingTalk token exchange URL")
|
||||
cmd.Flags().String("refresh-url", "", "Override DingTalk refresh token URL")
|
||||
cmd.Flags().Int("login-timeout", 0, "Compatibility flag for login timeout seconds")
|
||||
cmd.Flags().Bool("no-browser", false, "Compatibility flag for browser launch suppression")
|
||||
cmd.Flags().Int("login-timeout", 0, "Login timeout seconds")
|
||||
cmd.Flags().Bool("no-browser", false, "Suppress browser launch")
|
||||
_ = cmd.Flags().MarkHidden("redirect-url")
|
||||
_ = cmd.Flags().MarkHidden("scopes")
|
||||
_ = cmd.Flags().MarkHidden("authorize-url")
|
||||
_ = cmd.Flags().MarkHidden("token-url")
|
||||
_ = cmd.Flags().MarkHidden("refresh-url")
|
||||
_ = cmd.Flags().MarkHidden("login-timeout")
|
||||
_ = cmd.Flags().MarkHidden("no-browser")
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -153,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
|
||||
},
|
||||
}
|
||||
@@ -176,28 +228,37 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
configDir := defaultConfigDir()
|
||||
|
||||
authenticated := false
|
||||
updatedAt := ""
|
||||
refreshed := false
|
||||
var tokenData *authpkg.TokenData
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
if data, err := provider.Status(); err == nil {
|
||||
tokenData = data
|
||||
if !data.IsAccessTokenValid() && data.IsRefreshTokenValid() {
|
||||
refreshCtx, cancel := context.WithTimeout(cmd.Context(), 15*time.Second)
|
||||
_, refreshErr := provider.GetAccessToken(refreshCtx)
|
||||
cancel()
|
||||
if refreshErr == nil {
|
||||
if updatedData, statusErr := provider.Status(); statusErr == nil {
|
||||
data = updatedData
|
||||
tokenData = updatedData
|
||||
refreshed = true
|
||||
}
|
||||
} else if edition.Get().AutoPurgeToken {
|
||||
_ = authpkg.DeleteTokenData(configDir)
|
||||
}
|
||||
}
|
||||
if authStatusAuthenticated(data) {
|
||||
if authStatusAuthenticated(tokenData) {
|
||||
authenticated = true
|
||||
updatedAt = authStatusUpdatedAt(data)
|
||||
}
|
||||
}
|
||||
|
||||
// Check if JSON output is requested
|
||||
format, _ := cmd.Root().PersistentFlags().GetString("format")
|
||||
if strings.EqualFold(strings.TrimSpace(format), "json") {
|
||||
return writeAuthStatusJSON(cmd.OutOrStdout(), authenticated, refreshed, tokenData)
|
||||
}
|
||||
|
||||
// Default table output
|
||||
w := cmd.OutOrStdout()
|
||||
if authenticated {
|
||||
if refreshed {
|
||||
@@ -206,58 +267,20 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
} else {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "状态:", "已登录 ✅")
|
||||
}
|
||||
if updatedAt != "" {
|
||||
if updatedAt := authStatusUpdatedAt(tokenData); updatedAt != "" {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "有效期:", updatedAt)
|
||||
}
|
||||
} 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
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newAuthImportCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "import <file>",
|
||||
Short: "导入认证信息",
|
||||
Hidden: true,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
configDir := defaultConfigDir()
|
||||
if err := validateOptionalPath("credentials file", args[0]); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := authpkg.LoadExportedCredentials(cmd.Context(), args[0], configDir); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("failed to import credentials: %v", err))
|
||||
}
|
||||
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
refreshCtx, cancel := context.WithTimeout(cmd.Context(), 30*time.Second)
|
||||
defer cancel()
|
||||
token, refreshErr := provider.GetAccessToken(refreshCtx)
|
||||
tokenData, statusErr := provider.Status()
|
||||
if statusErr != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to load imported token data: %v", statusErr))
|
||||
}
|
||||
if refreshErr == nil {
|
||||
tokenData.AccessToken = token
|
||||
}
|
||||
clearCompatCache()
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintln(w, "[OK] 认证信息导入成功")
|
||||
if refreshErr != nil {
|
||||
fmt.Fprintf(w, "[WARN] 凭证暂时无法刷新: %v\n", refreshErr)
|
||||
}
|
||||
fmt.Fprintln(w, "Token 将自动刷新,无需重复登录")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newAuthExchangeCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "exchange",
|
||||
@@ -287,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()
|
||||
@@ -326,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
|
||||
},
|
||||
}
|
||||
@@ -415,3 +442,91 @@ func authStatusUpdatedAt(data *authpkg.TokenData) string {
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// authStatusResponse is the JSON response for auth status command.
|
||||
type authStatusResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Authenticated bool `json:"authenticated"`
|
||||
Message string `json:"message,omitempty"`
|
||||
Refreshed bool `json:"refreshed,omitempty"`
|
||||
TokenValid bool `json:"token_valid,omitempty"`
|
||||
RefreshTokenValid bool `json:"refresh_token_valid,omitempty"`
|
||||
ExpiresAt string `json:"expires_at,omitempty"`
|
||||
RefreshExpiresAt string `json:"refresh_expires_at,omitempty"`
|
||||
CorpID string `json:"corp_id,omitempty"`
|
||||
CorpName string `json:"corp_name,omitempty"`
|
||||
UserID string `json:"user_id,omitempty"`
|
||||
UserName string `json:"user_name,omitempty"`
|
||||
}
|
||||
|
||||
func writeAuthStatusJSON(w io.Writer, authenticated, refreshed bool, data *authpkg.TokenData) error {
|
||||
resp := authStatusResponse{
|
||||
Success: true,
|
||||
Authenticated: authenticated,
|
||||
}
|
||||
|
||||
if !authenticated {
|
||||
resp.Message = "未登录"
|
||||
} else if data != nil {
|
||||
resp.Refreshed = refreshed
|
||||
resp.TokenValid = data.IsAccessTokenValid()
|
||||
resp.RefreshTokenValid = data.IsRefreshTokenValid()
|
||||
if !data.ExpiresAt.IsZero() {
|
||||
resp.ExpiresAt = data.ExpiresAt.Format(time.RFC3339Nano)
|
||||
}
|
||||
if !data.RefreshExpAt.IsZero() {
|
||||
resp.RefreshExpiresAt = data.RefreshExpAt.Format(time.RFC3339Nano)
|
||||
}
|
||||
resp.CorpID = data.CorpID
|
||||
resp.CorpName = data.CorpName
|
||||
resp.UserID = data.UserID
|
||||
resp.UserName = data.UserName
|
||||
}
|
||||
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
return enc.Encode(resp)
|
||||
}
|
||||
|
||||
// authLoginResponse is the JSON response for auth login command.
|
||||
type authLoginResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Message string `json:"message"`
|
||||
TokenValid bool `json:"token_valid,omitempty"`
|
||||
RefreshTokenValid bool `json:"refresh_token_valid,omitempty"`
|
||||
ExpiresAt string `json:"expires_at,omitempty"`
|
||||
RefreshExpiresAt string `json:"refresh_expires_at,omitempty"`
|
||||
CorpID string `json:"corp_id,omitempty"`
|
||||
CorpName string `json:"corp_name,omitempty"`
|
||||
UserID string `json:"user_id,omitempty"`
|
||||
UserName string `json:"user_name,omitempty"`
|
||||
}
|
||||
|
||||
func writeAuthLoginJSON(w io.Writer, data *authpkg.TokenData, forced bool) error {
|
||||
resp := authLoginResponse{
|
||||
Success: true,
|
||||
Message: "登录成功",
|
||||
}
|
||||
|
||||
if data != nil {
|
||||
if data.IsAccessTokenValid() && !forced {
|
||||
resp.Message = "Token 有效,无需重新登录"
|
||||
}
|
||||
resp.TokenValid = data.IsAccessTokenValid()
|
||||
resp.RefreshTokenValid = data.IsRefreshTokenValid()
|
||||
if !data.ExpiresAt.IsZero() {
|
||||
resp.ExpiresAt = data.ExpiresAt.Format(time.RFC3339Nano)
|
||||
}
|
||||
if !data.RefreshExpAt.IsZero() {
|
||||
resp.RefreshExpiresAt = data.RefreshExpAt.Format(time.RFC3339Nano)
|
||||
}
|
||||
resp.CorpID = data.CorpID
|
||||
resp.CorpName = data.CorpName
|
||||
resp.UserID = data.UserID
|
||||
resp.UserName = data.UserName
|
||||
}
|
||||
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
return enc.Encode(resp)
|
||||
}
|
||||
|
||||
@@ -17,15 +17,20 @@ import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
// Cleanup keychain after test
|
||||
t.Cleanup(func() {
|
||||
_ = keychain.Remove(keychain.Service, keychain.AccountToken)
|
||||
})
|
||||
|
||||
root := t.TempDir()
|
||||
configDir := filepath.Join(root, "config")
|
||||
|
||||
@@ -39,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
|
||||
@@ -60,11 +65,12 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
t.Fatalf("Execute() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
|
||||
if _, err := os.Stat(filepath.Join(configDir, ".data")); err != nil {
|
||||
t.Fatalf("secure token data should remain after refresh failure: %v", err)
|
||||
// Verify token data still exists in keychain after refresh failure
|
||||
if !authpkg.TokenDataExistsKeychain() {
|
||||
t.Fatal("secure token data should remain in keychain after refresh failure")
|
||||
}
|
||||
|
||||
if !bytes.Contains(out.Bytes(), []byte("已登录")) {
|
||||
if !bytes.Contains(out.Bytes(), []byte("\"authenticated\"")) {
|
||||
t.Fatalf("output should still report authenticated status:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -24,10 +24,11 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
dynamicMu sync.RWMutex
|
||||
dynamicEndpoints map[string]string
|
||||
dynamicProducts map[string]bool
|
||||
dynamicAliases map[string]string
|
||||
dynamicMu sync.RWMutex
|
||||
dynamicEndpoints map[string]string
|
||||
dynamicProducts map[string]bool
|
||||
dynamicAliases map[string]string
|
||||
dynamicToolEndpoints map[string]string // tool name → endpoint
|
||||
)
|
||||
|
||||
var legacyDirectRuntimeAliases = map[string]string{
|
||||
@@ -46,6 +47,7 @@ func SetDynamicServers(servers []market.ServerDescriptor) {
|
||||
endpoints := make(map[string]string)
|
||||
products := make(map[string]bool)
|
||||
aliases := make(map[string]string)
|
||||
toolEndpoints := make(map[string]string)
|
||||
for _, server := range servers {
|
||||
if server.CLI.Skip {
|
||||
continue
|
||||
@@ -70,10 +72,26 @@ func SetDynamicServers(servers []market.ServerDescriptor) {
|
||||
aliases[alias] = id
|
||||
}
|
||||
}
|
||||
// Build tool → endpoint mapping from CLI tools and overrides.
|
||||
if endpoint != "" {
|
||||
for _, tool := range server.CLI.Tools {
|
||||
toolName := strings.TrimSpace(tool.Name)
|
||||
if toolName != "" {
|
||||
toolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
for toolName := range server.CLI.ToolOverrides {
|
||||
toolName = strings.TrimSpace(toolName)
|
||||
if toolName != "" {
|
||||
toolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
dynamicEndpoints = endpoints
|
||||
dynamicProducts = products
|
||||
dynamicAliases = aliases
|
||||
dynamicToolEndpoints = toolEndpoints
|
||||
}
|
||||
|
||||
func shouldUseDirectRuntime(invocation executor.Invocation) bool {
|
||||
@@ -88,11 +106,9 @@ func shouldUseDirectRuntime(invocation executor.Invocation) bool {
|
||||
}
|
||||
}
|
||||
|
||||
func directRuntimeEndpoint(productID string) (string, bool) {
|
||||
func directRuntimeEndpoint(productID, toolName string) (string, bool) {
|
||||
// Priority 0: env-var override always wins (DINGTALK_<PRODUCT>_MCP_URL).
|
||||
normalized := normalizeDirectRuntimeProductID(productID)
|
||||
dynamicMu.RLock()
|
||||
de := dynamicEndpoints
|
||||
dynamicMu.RUnlock()
|
||||
for _, candidate := range []string{strings.TrimSpace(productID), normalized} {
|
||||
if candidate == "" {
|
||||
continue
|
||||
@@ -100,6 +116,25 @@ func directRuntimeEndpoint(productID string) (string, bool) {
|
||||
if override, ok := productEndpointOverride(candidate); ok {
|
||||
return override, true
|
||||
}
|
||||
}
|
||||
|
||||
dynamicMu.RLock()
|
||||
de := dynamicEndpoints
|
||||
te := dynamicToolEndpoints
|
||||
dynamicMu.RUnlock()
|
||||
|
||||
// Priority 1: tool-level endpoint (resolves multi-endpoint products).
|
||||
if tool := strings.TrimSpace(toolName); tool != "" && te != nil {
|
||||
if endpoint, ok := te[tool]; ok {
|
||||
return endpoint, true
|
||||
}
|
||||
}
|
||||
|
||||
// Priority 2: product-level endpoint.
|
||||
for _, candidate := range []string{strings.TrimSpace(productID), normalized} {
|
||||
if candidate == "" {
|
||||
continue
|
||||
}
|
||||
if de != nil {
|
||||
if endpoint, ok := de[candidate]; ok {
|
||||
return endpoint, true
|
||||
@@ -122,6 +157,66 @@ func DirectRuntimeProductIDs() map[string]bool {
|
||||
return ids
|
||||
}
|
||||
|
||||
// AppendDynamicServer adds a single server descriptor to the existing
|
||||
// dynamic server registry without replacing the current entries. This
|
||||
// is used by the plugin loader to inject plugin servers alongside
|
||||
// Market-discovered servers.
|
||||
func AppendDynamicServer(server market.ServerDescriptor) {
|
||||
dynamicMu.Lock()
|
||||
defer dynamicMu.Unlock()
|
||||
|
||||
if dynamicEndpoints == nil {
|
||||
dynamicEndpoints = make(map[string]string)
|
||||
}
|
||||
if dynamicProducts == nil {
|
||||
dynamicProducts = make(map[string]bool)
|
||||
}
|
||||
if dynamicAliases == nil {
|
||||
dynamicAliases = make(map[string]string)
|
||||
}
|
||||
if dynamicToolEndpoints == nil {
|
||||
dynamicToolEndpoints = make(map[string]string)
|
||||
}
|
||||
|
||||
if server.CLI.Skip {
|
||||
return
|
||||
}
|
||||
|
||||
id := strings.TrimSpace(server.CLI.ID)
|
||||
endpoint := strings.TrimSpace(server.Endpoint)
|
||||
if id != "" && endpoint != "" {
|
||||
dynamicEndpoints[id] = endpoint
|
||||
dynamicProducts[id] = true
|
||||
}
|
||||
cmd := strings.TrimSpace(server.CLI.Command)
|
||||
if cmd != "" && cmd != id && endpoint != "" {
|
||||
dynamicEndpoints[cmd] = endpoint
|
||||
dynamicProducts[cmd] = true
|
||||
}
|
||||
for _, alias := range server.CLI.Aliases {
|
||||
alias = strings.TrimSpace(alias)
|
||||
if alias != "" && endpoint != "" {
|
||||
dynamicEndpoints[alias] = endpoint
|
||||
dynamicProducts[alias] = true
|
||||
dynamicAliases[alias] = id
|
||||
}
|
||||
}
|
||||
if endpoint != "" {
|
||||
for _, tool := range server.CLI.Tools {
|
||||
toolName := strings.TrimSpace(tool.Name)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
for toolName := range server.CLI.ToolOverrides {
|
||||
toolName = strings.TrimSpace(toolName)
|
||||
if toolName != "" {
|
||||
dynamicToolEndpoints[toolName] = endpoint
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func normalizeDirectRuntimeProductID(productID string) string {
|
||||
dynamicMu.RLock()
|
||||
da := dynamicAliases
|
||||
|
||||
@@ -0,0 +1,438 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/upgrade"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// checkStatus represents the outcome of a single doctor check.
|
||||
type checkStatus string
|
||||
|
||||
const (
|
||||
statusPass checkStatus = "pass"
|
||||
statusWarn checkStatus = "warn"
|
||||
statusFail checkStatus = "fail"
|
||||
)
|
||||
|
||||
// checkResult holds the outcome of a single doctor check.
|
||||
type checkResult struct {
|
||||
Name string `json:"name"`
|
||||
Status checkStatus `json:"status"`
|
||||
Message string `json:"message"`
|
||||
Hint string `json:"hint,omitempty"`
|
||||
Detail any `json:"detail,omitempty"`
|
||||
}
|
||||
|
||||
func newDoctorCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "doctor",
|
||||
Short: "环境健康检查",
|
||||
Long: "一键检查登录态、网络连通性、缓存状态和版本更新,快速定位常见问题。",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runDoctor,
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "以 JSON 格式输出")
|
||||
cmd.Flags().Int("timeout", 10, "网络检查超时时间 (秒)")
|
||||
cmd.Flags().Bool("perf", false, "额外展示最近一次性能报告")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runDoctor(cmd *cobra.Command, _ []string) error {
|
||||
jsonOut, _ := cmd.Flags().GetBool("json")
|
||||
timeout, _ := cmd.Flags().GetInt("timeout")
|
||||
if timeout <= 0 {
|
||||
timeout = 10
|
||||
}
|
||||
networkTimeout := time.Duration(timeout) * time.Second
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
checks := make([]checkResult, 0, 4)
|
||||
|
||||
authResult := doctorCheckAuth(cmd.Context(), w, jsonOut)
|
||||
checks = append(checks, authResult)
|
||||
|
||||
networkResult := doctorCheckNetwork(cmd.Context(), w, jsonOut, networkTimeout)
|
||||
checks = append(checks, networkResult)
|
||||
|
||||
cacheResult := doctorCheckCache(w, jsonOut)
|
||||
checks = append(checks, cacheResult)
|
||||
|
||||
versionResult := doctorCheckVersion(w, jsonOut, networkTimeout)
|
||||
checks = append(checks, versionResult)
|
||||
|
||||
showPerf, _ := cmd.Flags().GetBool("perf")
|
||||
if showPerf {
|
||||
perfResult := doctorCheckPerf(w, jsonOut)
|
||||
checks = append(checks, perfResult)
|
||||
}
|
||||
|
||||
pass, warn, fail := countResults(checks)
|
||||
|
||||
if jsonOut {
|
||||
result := map[string]any{
|
||||
"kind": "doctor",
|
||||
"checks": checks,
|
||||
"summary": map[string]int{
|
||||
"pass": pass,
|
||||
"warn": warn,
|
||||
"fail": fail,
|
||||
},
|
||||
}
|
||||
if showPerf {
|
||||
if report, err := LoadLatestReport(); err == nil {
|
||||
result["perf_report"] = report
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(w, result)
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "\n诊断完成: %d 项通过, %d 项警告, %d 项失败\n", pass, warn, fail)
|
||||
if fail > 0 {
|
||||
return fmt.Errorf("诊断发现 %d 项失败", fail)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── Auth check ──────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckAuth(ctx context.Context, w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查登录状态... ")
|
||||
}
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
|
||||
data, err := provider.Status()
|
||||
if err != nil || data == nil {
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "未登录"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 进行登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if data.IsAccessTokenValid() || data.IsRefreshTokenValid() {
|
||||
if !data.IsAccessTokenValid() {
|
||||
refreshCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
|
||||
_, refreshErr := provider.GetAccessToken(refreshCtx)
|
||||
cancel()
|
||||
if refreshErr != nil {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusWarn,
|
||||
Message: "Refresh Token 有效, 但自动刷新 Access Token 失败",
|
||||
Hint: "运行 dws auth login 重新登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{Name: "auth", Status: statusFail, Message: "登录已过期"}
|
||||
if !edition.Get().IsEmbedded {
|
||||
r.Hint = "运行 dws auth login 重新登录"
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Network check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckNetwork(ctx context.Context, w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查网络连通性... ")
|
||||
}
|
||||
|
||||
baseURL := cli.DefaultMarketBaseURL
|
||||
httpClient := &http.Client{Timeout: timeout}
|
||||
client := market.NewClient(baseURL, httpClient)
|
||||
|
||||
start := time.Now()
|
||||
reqCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
_, err := client.FetchServers(reqCtx, 1)
|
||||
latency := time.Since(start)
|
||||
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 不可达: %v", err),
|
||||
Hint: "请检查网络连接或代理设置",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "network",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("mcp.dingtalk.com 可达 (延迟 %dms)", latency.Milliseconds()),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Cache check ─────────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckCache(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查缓存状态... ")
|
||||
}
|
||||
|
||||
store := cacheStoreFromEnv()
|
||||
files, _, err := cacheDirectoryStats(store.Root)
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("缓存目录不可读: %v", err),
|
||||
Hint: "运行 dws cache clean 清理后重试",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
entries, _ := store.ListToolsCacheEntries(config.DefaultPartition)
|
||||
|
||||
if files == 0 && len(entries) == 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: "缓存为空 (首次使用)",
|
||||
Hint: "运行任意 dws 命令后将自动建立缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
staleCount := 0
|
||||
for _, e := range entries {
|
||||
if e.Freshness == cache.FreshnessStale {
|
||||
staleCount++
|
||||
}
|
||||
}
|
||||
|
||||
if staleCount > 0 {
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("%d 个文件, %d 个工具缓存, %d 个已过期", files, len(entries), staleCount),
|
||||
Hint: "运行 dws cache refresh 刷新缓存",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
msg := fmt.Sprintf("%d 个文件, %d 个工具缓存", files, len(entries))
|
||||
if len(entries) > 0 {
|
||||
msg += ", 全部新鲜"
|
||||
}
|
||||
r := checkResult{
|
||||
Name: "cache",
|
||||
Status: statusPass,
|
||||
Message: msg,
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Version check ───────────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckVersion(w io.Writer, jsonOut bool, timeout time.Duration) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查版本更新... ")
|
||||
}
|
||||
|
||||
currentVer := version
|
||||
|
||||
client := upgrade.NewClient()
|
||||
latest, err := client.FetchLatestRelease()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusFail,
|
||||
Message: fmt.Sprintf("无法获取最新版本: %v", err),
|
||||
Hint: "请检查网络连接",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
if upgrade.NeedsUpgrade(currentVer, latest.Version) {
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusWarn,
|
||||
Message: fmt.Sprintf("有新版本 (当前 %s, 最新 v%s)", ensureV(currentVer), latest.Version),
|
||||
Hint: "运行 dws upgrade 升级到最新版本",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "version",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("已是最新版本 %s", ensureV(currentVer)),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
// ── Output helpers ──────────────────────────────────────────────────────
|
||||
|
||||
func printCheckResult(w io.Writer, r checkResult) {
|
||||
icon := statusIcon(r.Status)
|
||||
fmt.Fprintf(w, "%s %s\n", icon, r.Message)
|
||||
if r.Hint != "" {
|
||||
fmt.Fprintf(w, " %s\n", r.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func statusIcon(s checkStatus) string {
|
||||
switch s {
|
||||
case statusPass:
|
||||
return "✅"
|
||||
case statusWarn:
|
||||
return "⚠️"
|
||||
case statusFail:
|
||||
return "❌"
|
||||
default:
|
||||
return "?"
|
||||
}
|
||||
}
|
||||
|
||||
func countResults(checks []checkResult) (pass, warn, fail int) {
|
||||
for _, c := range checks {
|
||||
switch c.Status {
|
||||
case statusPass:
|
||||
pass++
|
||||
case statusWarn:
|
||||
warn++
|
||||
case statusFail:
|
||||
fail++
|
||||
}
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// ── Perf report check ──────────────────────────────────────────────────
|
||||
|
||||
func doctorCheckPerf(w io.Writer, jsonOut bool) checkResult {
|
||||
if !jsonOut {
|
||||
fmt.Fprint(w, "检查性能报告... ")
|
||||
}
|
||||
|
||||
report, err := LoadLatestReport()
|
||||
if err != nil {
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusWarn,
|
||||
Message: "未找到性能报告",
|
||||
Hint: "设置 DWS_PERF_REPORT=auto 后运行任意命令生成报告",
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
r := checkResult{
|
||||
Name: "perf",
|
||||
Status: statusPass,
|
||||
Message: fmt.Sprintf("报告可用 (%s, %s)", report.Command, report.Timestamp.Local().Format("2006-01-02 15:04")),
|
||||
}
|
||||
if !jsonOut {
|
||||
printCheckResult(w, r)
|
||||
printPerfReportSummary(w, report)
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
func printPerfReportSummary(w io.Writer, report *PerfReport) {
|
||||
fmt.Fprintf(w, "\n最近一次性能报告 (%s, %s):\n",
|
||||
report.Command, report.Timestamp.Local().Format("2006-01-02 15:04"))
|
||||
|
||||
for _, p := range report.Phases {
|
||||
marker := ""
|
||||
if p.Name == report.Slowest {
|
||||
marker = " ← 最慢"
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s %dms%s\n", p.Name, p.DurationMs, marker)
|
||||
}
|
||||
fmt.Fprintf(w, " %-25s ─────────\n", "─────────────────────────")
|
||||
fmt.Fprintf(w, " %-25s %dms (框架开销 %dms)\n", "总耗时", report.TotalMs, report.OverheadMs)
|
||||
}
|
||||
|
||||
func formatLocalTime(t time.Time) string {
|
||||
if t.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return t.Local().Format("2006-01-02 15:04")
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCountResults(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
{Status: statusWarn},
|
||||
{Status: statusFail},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 1 || fail != 1 {
|
||||
t.Errorf("expected (2,1,1), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCountResultsAllPass(t *testing.T) {
|
||||
checks := []checkResult{
|
||||
{Status: statusPass},
|
||||
{Status: statusPass},
|
||||
}
|
||||
pass, warn, fail := countResults(checks)
|
||||
if pass != 2 || warn != 0 || fail != 0 {
|
||||
t.Errorf("expected (2,0,0), got (%d,%d,%d)", pass, warn, fail)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStatusIcon(t *testing.T) {
|
||||
tests := []struct {
|
||||
status checkStatus
|
||||
want string
|
||||
}{
|
||||
{statusPass, "✅"},
|
||||
{statusWarn, "⚠️"},
|
||||
{statusFail, "❌"},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
got := statusIcon(tc.status)
|
||||
if got != tc.want {
|
||||
t.Errorf("statusIcon(%q) = %q, want %q", tc.status, got, tc.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResult(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusFail,
|
||||
Message: "something broke",
|
||||
Hint: "try fixing it",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "❌") {
|
||||
t.Error("expected fail icon")
|
||||
}
|
||||
if !strings.Contains(out, "something broke") {
|
||||
t.Error("expected message")
|
||||
}
|
||||
if !strings.Contains(out, "try fixing it") {
|
||||
t.Error("expected hint")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintCheckResultNoHint(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
r := checkResult{
|
||||
Name: "test",
|
||||
Status: statusPass,
|
||||
Message: "all good",
|
||||
}
|
||||
printCheckResult(&buf, r)
|
||||
|
||||
out := buf.String()
|
||||
if !strings.Contains(out, "✅") {
|
||||
t.Error("expected pass icon")
|
||||
}
|
||||
lines := strings.Split(strings.TrimSpace(out), "\n")
|
||||
if len(lines) != 1 {
|
||||
t.Errorf("expected 1 line (no hint), got %d", len(lines))
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmpty(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, false)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if !strings.Contains(r.Message, "缓存为空") {
|
||||
t.Errorf("expected empty cache message, got %q", r.Message)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCheckCacheEmptyJSON(t *testing.T) {
|
||||
t.Setenv("DWS_CACHE_DIR", t.TempDir())
|
||||
|
||||
var buf bytes.Buffer
|
||||
r := doctorCheckCache(&buf, true)
|
||||
|
||||
if r.Status != statusWarn {
|
||||
t.Errorf("expected warn for empty cache, got %s", r.Status)
|
||||
}
|
||||
if buf.Len() != 0 {
|
||||
t.Error("expected no output in JSON mode")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDoctorCommandStructure(t *testing.T) {
|
||||
cmd := newDoctorCommand()
|
||||
if cmd.Use != "doctor" {
|
||||
t.Errorf("Use = %q, want doctor", cmd.Use)
|
||||
}
|
||||
|
||||
jsonFlag := cmd.Flags().Lookup("json")
|
||||
if jsonFlag == nil {
|
||||
t.Error("expected --json flag")
|
||||
}
|
||||
timeoutFlag := cmd.Flags().Lookup("timeout")
|
||||
if timeoutFlag == nil {
|
||||
t.Error("expected --timeout flag")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckResultJSONMarshal(t *testing.T) {
|
||||
r := checkResult{
|
||||
Name: "auth",
|
||||
Status: statusPass,
|
||||
Message: "已登录",
|
||||
}
|
||||
data, err := json.Marshal(r)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if parsed["name"] != "auth" {
|
||||
t.Errorf("expected name=auth, got %v", parsed["name"])
|
||||
}
|
||||
if parsed["status"] != "pass" {
|
||||
t.Errorf("expected status=pass, got %v", parsed["status"])
|
||||
}
|
||||
if _, hasHint := parsed["hint"]; hasHint {
|
||||
t.Error("empty hint should be omitted")
|
||||
}
|
||||
}
|
||||
@@ -14,6 +14,7 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -23,7 +24,9 @@ type GlobalFlags struct {
|
||||
ClientSecret string
|
||||
Debug bool
|
||||
DryRun bool
|
||||
Fields string
|
||||
Format string
|
||||
JQ string
|
||||
Mock bool
|
||||
Output string
|
||||
Timeout int
|
||||
@@ -33,11 +36,13 @@ type GlobalFlags struct {
|
||||
}
|
||||
|
||||
func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", "Override OAuth client ID (DingTalk AppKey)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", "Override OAuth client secret (DingTalk AppSecret)")
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientID, "client-id", "", i18n.T("覆盖 OAuth 客户端 ID (钉钉 AppKey)"))
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", i18n.T("覆盖 OAuth 客户端密钥 (钉钉 AppSecret)"))
|
||||
cmd.PersistentFlags().BoolVar(&flags.Debug, "debug", false, "显示调试日志")
|
||||
cmd.PersistentFlags().BoolVar(&flags.DryRun, "dry-run", false, "预览操作内容,不实际执行")
|
||||
cmd.PersistentFlags().StringVarP(&flags.Format, "format", "f", "table", "输出格式: json|table|raw")
|
||||
cmd.PersistentFlags().StringVar(&flags.Fields, "fields", "", "筛选输出字段 (逗号分隔, 如: name,id,status)")
|
||||
cmd.PersistentFlags().StringVarP(&flags.Format, "format", "f", "json", "输出格式: json|table|raw")
|
||||
cmd.PersistentFlags().StringVar(&flags.JQ, "jq", "", "jq 表达式过滤输出 (如: '.items[] | .name')")
|
||||
cmd.PersistentFlags().BoolVar(&flags.Mock, "mock", false, "使用 Mock 数据 (开发调试用)")
|
||||
cmd.PersistentFlags().StringVarP(&flags.Output, "output", "o", "", "Write command output to a file")
|
||||
_ = cmd.PersistentFlags().MarkHidden("output")
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
+45
-11
@@ -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
|
||||
}
|
||||
@@ -332,10 +368,8 @@ func asyncRevalidateRegistry(parent context.Context, store *cache.Store, partiti
|
||||
}
|
||||
}
|
||||
|
||||
func newLegacyHiddenCommands(runner executor.Runner) []*cobra.Command {
|
||||
var commands []*cobra.Command
|
||||
commands = append(commands, helpers.NewHiddenVendorCommands(runner)...)
|
||||
return commands
|
||||
func newLegacyHiddenCommands(_ executor.Runner) []*cobra.Command {
|
||||
return nil
|
||||
}
|
||||
|
||||
func mergeTopLevelCommands(commands []*cobra.Command) []*cobra.Command {
|
||||
|
||||
@@ -277,6 +277,10 @@ func TestLoadDynamicCommandsFallsBackToStaleCacheOnNetworkError(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestLoadDynamicCommandsRefreshesRegistryCacheInBackgroundAfterAgedStart(t *testing.T) {
|
||||
// Skip: async revalidation is disabled when discoveryBaseURLOverride is set.
|
||||
// This test requires background refresh which only runs in production mode.
|
||||
t.Skip("async revalidation disabled in test mode")
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, "")
|
||||
|
||||
var phase atomic.Int32
|
||||
|
||||
@@ -0,0 +1,577 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"os/exec"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/fatih/color"
|
||||
)
|
||||
|
||||
const (
|
||||
// PatAuthRetryTimeout is the maximum time to wait for user authorization
|
||||
// when a PAT scope error is detected.
|
||||
PatAuthRetryTimeout = 10 * time.Minute
|
||||
|
||||
// PatAuthPollInterval is how often we poll to check if the user has
|
||||
// completed authorization.
|
||||
PatAuthPollInterval = 5 * time.Second
|
||||
)
|
||||
|
||||
// PatScopeError holds information about a missing PAT scope.
|
||||
type PatScopeError struct {
|
||||
OriginalError string
|
||||
Identity string
|
||||
ErrorType string
|
||||
Message string
|
||||
Hint string
|
||||
MissingScope string
|
||||
}
|
||||
|
||||
func (e *PatScopeError) Error() string {
|
||||
return e.OriginalError
|
||||
}
|
||||
|
||||
// patScopeRegex matches PAT-protocol scope error patterns from the API.
|
||||
// Only matches explicit scope-related keywords; generic "permission denied" or
|
||||
// "forbidden" are intentionally excluded to avoid false positives on business
|
||||
// authorization errors (e.g. mailbox access denied, 403 Forbidden).
|
||||
var patScopeRegex = regexp.MustCompile(`(?i)(missing_scope|insufficient_scope|scope.*required)`)
|
||||
|
||||
// scopeValueRegex extracts a scope identifier (e.g. "calendar:read",
|
||||
// "mail:user_mailbox.message:send") from an error message.
|
||||
// Supports multi-segment scopes with multiple colons (resource:sub:action).
|
||||
var scopeValueRegex = regexp.MustCompile(`([a-zA-Z][a-zA-Z0-9_.]*(?::[a-zA-Z][a-zA-Z0-9_.]*)+)`)
|
||||
|
||||
// identityValueRegex extracts an identity label from an error message.
|
||||
var identityValueRegex = regexp.MustCompile(`(?i)identity["\s:]+([a-zA-Z_]+)`)
|
||||
|
||||
// isPatScopeError checks if an error looks like a PAT scope/permission error
|
||||
// that can be resolved by re-authorizing with additional scopes.
|
||||
func isPatScopeError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
|
||||
// Check for missing_scope pattern in error message or hint
|
||||
if patScopeRegex.MatchString(msg) {
|
||||
return true
|
||||
}
|
||||
|
||||
var typed *apperrors.Error
|
||||
if stderrors.As(err, &typed) {
|
||||
// Check message, reason, and hint for scope-related patterns
|
||||
fullText := strings.ToLower(typed.Message + " " + typed.Reason + " " + typed.Hint)
|
||||
if typed.Category == apperrors.CategoryAuth {
|
||||
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") ||
|
||||
(strings.Contains(fullText, "scope") && strings.Contains(fullText, "required")) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
// Any category with scope/permission hints
|
||||
if strings.Contains(fullText, "missing_scope") || strings.Contains(fullText, "insufficient_scope") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// extractPatScopeError parses an error to extract PAT scope details.
|
||||
func extractPatScopeError(err error) *PatScopeError {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
msg := err.Error()
|
||||
scope := ""
|
||||
|
||||
var typed *apperrors.Error
|
||||
if stderrors.As(err, &typed) {
|
||||
msg = typed.Message
|
||||
if typed.Reason != "" {
|
||||
msg += " (" + typed.Reason + ")"
|
||||
}
|
||||
}
|
||||
|
||||
// Try to extract scope value (e.g. "calendar:read") from error message.
|
||||
scopeMatch := scopeValueRegex.FindStringSubmatch(msg)
|
||||
if len(scopeMatch) > 1 {
|
||||
scope = scopeMatch[1]
|
||||
}
|
||||
|
||||
// Try to extract identity from error message.
|
||||
identity := "user"
|
||||
identityMatch := identityValueRegex.FindStringSubmatch(msg)
|
||||
if len(identityMatch) > 1 {
|
||||
identity = identityMatch[1]
|
||||
}
|
||||
|
||||
return &PatScopeError{
|
||||
OriginalError: err.Error(),
|
||||
Identity: identity,
|
||||
ErrorType: "missing_scope",
|
||||
Message: msg,
|
||||
Hint: fmt.Sprintf("run `dws auth login --scope %q` to authorize the missing scope", scope),
|
||||
MissingScope: scope,
|
||||
}
|
||||
}
|
||||
|
||||
// PrintPatAuthError prints a human-readable PAT authorization error.
|
||||
func PrintPatAuthError(w io.Writer, scopeErr *PatScopeError) {
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
cyan := color.New(color.FgCyan).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
green := color.New(color.FgGreen).SprintFunc()
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "{\n")
|
||||
fmt.Fprintf(w, " %s: %s,\n", bold("\"ok\""), "false")
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"identity\""), scopeErr.Identity)
|
||||
fmt.Fprintf(w, " %s: {\n", bold("\"error\""))
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"type\""), scopeErr.ErrorType)
|
||||
fmt.Fprintf(w, " %s: %q,\n", bold("\"message\""), scopeErr.Message)
|
||||
fmt.Fprintf(w, " %s: %q\n", bold("\"hint\""), scopeErr.Hint)
|
||||
fmt.Fprintf(w, " }\n")
|
||||
fmt.Fprintf(w, "}\n")
|
||||
fmt.Fprintln(w)
|
||||
|
||||
// Print authorization instructions
|
||||
fmt.Fprintf(w, "%s %s\n", green("▶"), bold("需要额外授权"))
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, " %s %s\n", dim("#"), dim("运行以下命令完成授权"))
|
||||
|
||||
if scopeErr.MissingScope != "" {
|
||||
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan(fmt.Sprintf("dws auth login --scope %q", scopeErr.MissingScope)))
|
||||
} else {
|
||||
fmt.Fprintf(w, " %s %s\n", cyan("$"), cyan("dws auth login"))
|
||||
}
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, " %s 在浏览器中打开授权链接,完成授权后重新执行命令\n", dim("ℹ"))
|
||||
fmt.Fprintln(w)
|
||||
}
|
||||
|
||||
// PrintPatAuthJSON prints a machine-readable PAT authorization error.
|
||||
func PrintPatAuthJSON(w io.Writer, scopeErr *PatScopeError) {
|
||||
payload := map[string]any{
|
||||
"ok": false,
|
||||
"identity": scopeErr.Identity,
|
||||
"error": map[string]any{
|
||||
"type": scopeErr.ErrorType,
|
||||
"message": scopeErr.Message,
|
||||
"hint": scopeErr.Hint,
|
||||
},
|
||||
}
|
||||
if scopeErr.MissingScope != "" {
|
||||
payload["missing_scope"] = scopeErr.MissingScope
|
||||
}
|
||||
|
||||
data, _ := json.MarshalIndent(payload, "", " ")
|
||||
fmt.Fprintln(w, string(data))
|
||||
}
|
||||
|
||||
// WaitForPatAuthorization polls until the user completes authorization or timeout.
|
||||
// It returns true if authorization was completed, false if timed out or cancelled.
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
yellow := color.New(color.FgYellow).SprintFunc()
|
||||
green := color.New(color.FgGreen).SprintFunc()
|
||||
red := color.New(color.FgRed).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
|
||||
timeout := PatAuthRetryTimeout
|
||||
deadline := time.Now().Add(timeout)
|
||||
pollTicker := time.NewTicker(PatAuthPollInterval)
|
||||
defer pollTicker.Stop()
|
||||
start := time.Now()
|
||||
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", yellow("⏳"), bold("等待用户授权..."))
|
||||
fmt.Fprintf(output, " %s 请在另一个终端完成 dws auth login 授权\n", dim("ℹ"))
|
||||
fmt.Fprintf(output, " %s 超时时间: %s\n", dim("⏱"), timeout)
|
||||
fmt.Fprintln(output)
|
||||
|
||||
pollCount := 0
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fmt.Fprintf(output, "%s 操作已取消\n", red("✗"))
|
||||
return false
|
||||
|
||||
case <-time.After(time.Until(deadline)):
|
||||
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", red("✗"), timeout)
|
||||
fmt.Fprintf(output, " %s 请重新执行命令\n", dim("ℹ"))
|
||||
return false
|
||||
|
||||
case <-pollTicker.C:
|
||||
pollCount++
|
||||
elapsed := time.Since(start).Truncate(time.Second)
|
||||
remaining := time.Until(deadline).Truncate(time.Second)
|
||||
|
||||
// Check if token is now valid
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err == nil && tokenData != nil {
|
||||
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
green("✓"), bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// Show polling status
|
||||
fmt.Fprintf(output, "\r%s [%d] 等待授权中... (%s 已用, %s 剩余) ",
|
||||
dim("⟳"), pollCount, elapsed, remaining)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// retryWithPatAuthRetry wraps an invocation that failed with a PAT scope error.
|
||||
// It waits for the user to complete authorization and then retries the invocation.
|
||||
func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocation executor.Invocation, scopeErr *PatScopeError, configDir string, output io.Writer) (executor.Result, error) {
|
||||
// Print the PAT error in human-readable format
|
||||
PrintPatAuthError(output, scopeErr)
|
||||
|
||||
// Wait for user to complete authorization
|
||||
authorized := WaitForPatAuthorization(ctx, configDir, output)
|
||||
if !authorized {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"等待用户授权超时",
|
||||
apperrors.WithReason("pat_auth_timeout"),
|
||||
apperrors.WithHint(fmt.Sprintf("授权超时 (%s),请重新执行命令", PatAuthRetryTimeout)),
|
||||
apperrors.WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
// Clear the token cache so the new token is loaded
|
||||
ResetRuntimeTokenCache()
|
||||
|
||||
// Retry the invocation
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", color.New(color.FgGreen).SprintFunc()("▶"),
|
||||
color.New(color.Bold).SprintFunc()("授权完成,正在重试..."))
|
||||
fmt.Fprintln(output)
|
||||
|
||||
return runner.Run(ctx, invocation)
|
||||
}
|
||||
|
||||
// ---- handlePatAuthCheck (runner.go entry point) -----------------------------
|
||||
|
||||
const (
|
||||
// patPollInterval is how often we poll the device flow status endpoint.
|
||||
patPollInterval = 2 * time.Second
|
||||
// patPollTimeout is the maximum time to wait for user authorization via device flow.
|
||||
patPollTimeout = 10 * time.Minute
|
||||
)
|
||||
|
||||
// patRetryingKey is a context key to prevent recursive PAT auth checks.
|
||||
// After APPROVED, the retry should not trigger another PAT flow.
|
||||
type patRetryingKeyType struct{}
|
||||
|
||||
var patRetryingKey = patRetryingKeyType{}
|
||||
|
||||
// IsPatRetrying returns true if the current context is already in a PAT retry.
|
||||
func IsPatRetrying(ctx context.Context) bool {
|
||||
v, _ := ctx.Value(patRetryingKey).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
// handlePatAuthCheck is called by runner.executeInvocation when a PAT
|
||||
// authorization error is detected. It injects the server-assigned clientId
|
||||
// as x-robot-uid header, prints authorization details, opens the browser,
|
||||
// polls the device flow endpoint until the user authorizes, and retries the
|
||||
// original invocation on success.
|
||||
func handlePatAuthCheck(
|
||||
ctx context.Context,
|
||||
r *runtimeRunner,
|
||||
invocation executor.Invocation,
|
||||
patErr *apperrors.PATError,
|
||||
configDir string,
|
||||
output io.Writer,
|
||||
) (executor.Result, error) {
|
||||
// Parse authorization details from PATError.RawJSON.
|
||||
var patData struct {
|
||||
Code string `json:"code"`
|
||||
Data struct {
|
||||
Desc string `json:"desc"`
|
||||
FlowID string `json:"flowId"`
|
||||
URI string `json:"uri"`
|
||||
ClientID string `json:"clientId"`
|
||||
ClientSecret string `json:"clientSecret"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(patErr.RawJSON), &patData); err != nil {
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
slog.Debug("PAT auth check",
|
||||
"clientId", patData.Data.ClientID,
|
||||
"flowId", patData.Data.FlowID,
|
||||
"hasSecret", patData.Data.ClientSecret != "",
|
||||
)
|
||||
|
||||
// Inject clientId/clientSecret from PAT response as runtime credentials
|
||||
// so that subsequent device flow auth uses the server-assigned app identity.
|
||||
if patData.Data.ClientID != "" {
|
||||
if patData.Data.ClientSecret != "" {
|
||||
// When both clientId and clientSecret are provided, use direct mode
|
||||
// (DingTalk API) rather than MCP proxy — the MCP proxy does not hold
|
||||
// the secret for this particular app.
|
||||
authpkg.SetClientID(patData.Data.ClientID)
|
||||
authpkg.SetClientSecret(patData.Data.ClientSecret)
|
||||
} else {
|
||||
// No clientSecret — rely on MCP proxy to manage the secret server-side.
|
||||
authpkg.SetClientIDFromMCP(patData.Data.ClientID)
|
||||
}
|
||||
|
||||
// Persist clientId (and optionally secret) to ~/.dws/app.json so that
|
||||
// future process invocations can load it at startup and populate
|
||||
// DWS_CLIENT_ID env before the first MCP request.
|
||||
appCfg := &authpkg.AppConfig{
|
||||
ClientID: patData.Data.ClientID,
|
||||
}
|
||||
if patData.Data.ClientSecret != "" {
|
||||
appCfg.ClientSecret = authpkg.PlainSecret(patData.Data.ClientSecret)
|
||||
}
|
||||
if err := authpkg.SaveAppConfig(configDir, appCfg); err != nil {
|
||||
slog.Warn("failed to persist app config from PAT", "error", err)
|
||||
fmt.Fprintf(output, " \u26a0 保存应用配置失败: %v (下次启动可能需要重新授权)\n", err)
|
||||
}
|
||||
}
|
||||
|
||||
bold := color.New(color.Bold).SprintFunc()
|
||||
cyan := color.New(color.FgCyan).SprintFunc()
|
||||
greenFn := color.New(color.FgGreen).SprintFunc()
|
||||
yellowFn := color.New(color.FgYellow).SprintFunc()
|
||||
redFn := color.New(color.FgRed).SprintFunc()
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
|
||||
fmt.Fprintln(output)
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("需要 PAT 授权"))
|
||||
if patData.Data.Desc != "" {
|
||||
fmt.Fprintf(output, " %s %s\n", dim("ℹ"), patData.Data.Desc)
|
||||
}
|
||||
if patData.Data.URI != "" {
|
||||
fmt.Fprintf(output, " %s %s\n\n", dim("🔗"), cyan(patData.Data.URI))
|
||||
// Best-effort browser open.
|
||||
_ = tryOpenBrowser(patData.Data.URI)
|
||||
}
|
||||
|
||||
// If no flowId, we can't poll — fall back to returning PATError for host-app.
|
||||
if patData.Data.FlowID == "" {
|
||||
fmt.Fprintln(output)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
// Poll the device flow status until user authorizes, rejects, or timeout.
|
||||
fmt.Fprintf(output, "%s %s\n", yellowFn("⏳"), bold("等待用户授权..."))
|
||||
fmt.Fprintf(output, " %s 请在浏览器中完成授权,超时时间: %s\n", dim("ℹ"), patPollTimeout)
|
||||
fmt.Fprintln(output)
|
||||
|
||||
pollCtx, cancel := context.WithTimeout(ctx, patPollTimeout)
|
||||
defer cancel()
|
||||
|
||||
status, authCode, err := pollPatDeviceFlow(pollCtx, patData.Data.FlowID, configDir, output)
|
||||
if err != nil {
|
||||
fmt.Fprintf(output, "%s 轮询授权状态失败: %v\n", redFn("✗"), err)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
|
||||
switch status {
|
||||
case authpkg.StatusApproved:
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("✓"), bold("授权成功!"))
|
||||
fmt.Fprintln(output)
|
||||
|
||||
// Exchange authCode for a fresh access token (mirrors device_flow loginOnce).
|
||||
if authCode != "" {
|
||||
slog.Debug("PAT retry: exchanging authCode for token", "hasCode", true)
|
||||
tokenData, exchErr := authpkg.ExchangeCodeForToken(ctx, configDir, authCode)
|
||||
if exchErr != nil {
|
||||
slog.Warn("PAT retry: exchangeCode failed, retrying with existing token", "error", exchErr)
|
||||
fmt.Fprintf(output, " %s 换取新 token 失败: %v (将使用现有凭证重试)\n", yellowFn("⚠"), exchErr)
|
||||
} else {
|
||||
if err := authpkg.SaveTokenData(configDir, tokenData); err != nil {
|
||||
slog.Warn("PAT retry: failed to save new token", "error", err)
|
||||
fmt.Fprintf(output, " %s 保存新 token 失败: %v\n", yellowFn("⚠"), err)
|
||||
} else {
|
||||
slog.Debug("PAT retry: token refreshed and saved")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Clear token cache so the new credentials take effect.
|
||||
ResetRuntimeTokenCache()
|
||||
|
||||
// Workaround: brief delay to let server-side authorization state propagate
|
||||
// before retrying. Without this the retry may use stale credentials.
|
||||
slog.Debug("PAT retry: waiting for server-side state propagation", "delay", "1s")
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// Retry the original invocation with pat-retrying flag to prevent recursion.
|
||||
fmt.Fprintf(output, "%s %s\n", greenFn("▶"), bold("授权完成,正在重试..."))
|
||||
fmt.Fprintln(output)
|
||||
slog.Debug("PAT retry: identity env check",
|
||||
"DWS_CLIENT_ID", os.Getenv("DWS_CLIENT_ID"),
|
||||
)
|
||||
retryCtx := context.WithValue(ctx, patRetryingKey, true)
|
||||
return r.Run(retryCtx, invocation)
|
||||
|
||||
case authpkg.StatusRejected:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("用户已拒绝授权"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"用户已拒绝授权",
|
||||
apperrors.WithReason("pat_auth_rejected"),
|
||||
apperrors.WithHint("用户在浏览器中拒绝了授权请求,请重新执行命令。"),
|
||||
)
|
||||
|
||||
case authpkg.StatusExpired:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("授权超时"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"授权超时",
|
||||
apperrors.WithReason("pat_auth_expired"),
|
||||
apperrors.WithHint("授权链接已过期,请重新执行命令。"),
|
||||
)
|
||||
|
||||
case authpkg.StatusCancelled:
|
||||
fmt.Fprintf(output, "%s %s\n", redFn("✗"), bold("操作已取消"))
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"操作已取消",
|
||||
apperrors.WithReason("pat_auth_cancelled"),
|
||||
apperrors.WithHint("用户取消了授权操作。"),
|
||||
)
|
||||
|
||||
default:
|
||||
fmt.Fprintf(output, "%s 未知授权状态: %s\n", redFn("✗"), status)
|
||||
return executor.Result{}, patErr
|
||||
}
|
||||
}
|
||||
|
||||
// pollPatDeviceFlow polls the PAT device flow status endpoint until a terminal
|
||||
// state (APPROVED/REJECTED/EXPIRED) is reached or the context is cancelled.
|
||||
// Returns the final status string and the authCode (non-empty only on APPROVED).
|
||||
func pollPatDeviceFlow(ctx context.Context, flowID string, configDir string, output io.Writer) (string, string, error) {
|
||||
pollURL := fmt.Sprintf("%s%s?flowId=%s",
|
||||
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
|
||||
|
||||
// Load user access token for the poll request header.
|
||||
var accessToken string
|
||||
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
|
||||
accessToken = tokenData.AccessToken
|
||||
}
|
||||
|
||||
// Use a client that does NOT follow redirects, so we can detect SSO 302.
|
||||
noRedirectClient := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(patPollInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
dim := color.New(color.Faint).SprintFunc()
|
||||
pollCount := 0
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
if ctx.Err() == context.Canceled {
|
||||
return authpkg.StatusCancelled, "", nil
|
||||
}
|
||||
return authpkg.StatusExpired, "", nil
|
||||
case <-ticker.C:
|
||||
pollCount++
|
||||
fmt.Fprintf(output, "\r%s [%d] 等待授权中... ", dim("⟳"), pollCount)
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, pollURL, nil)
|
||||
if err != nil {
|
||||
slog.Debug("PAT poll: failed to create request", "error", err)
|
||||
continue
|
||||
}
|
||||
if accessToken != "" {
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
}
|
||||
resp, err := noRedirectClient.Do(req)
|
||||
if err != nil {
|
||||
slog.Debug("PAT poll: request failed", "error", err)
|
||||
continue // transient network error, keep polling
|
||||
}
|
||||
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
resp.Body.Close()
|
||||
|
||||
// If we got a redirect (302/301), SSO gateway intercepted — skip JSON parse.
|
||||
if resp.StatusCode == http.StatusFound || resp.StatusCode == http.StatusMovedPermanently {
|
||||
continue
|
||||
}
|
||||
|
||||
var pollResp authpkg.DevicePollResponse
|
||||
if err := json.Unmarshal(bodyBytes, &pollResp); err != nil {
|
||||
slog.Debug("PAT poll: failed to parse response", "error", err, "body", string(bodyBytes))
|
||||
continue
|
||||
}
|
||||
|
||||
status := authpkg.ParseDeviceFlowStatus(pollResp.Data.Status, pollResp.Success)
|
||||
switch status {
|
||||
case authpkg.StatusApproved:
|
||||
fmt.Fprintln(output) // clear the polling line
|
||||
return status, pollResp.Data.AuthCode, nil
|
||||
case authpkg.StatusRejected, authpkg.StatusExpired:
|
||||
fmt.Fprintln(output) // clear the polling line
|
||||
return status, "", nil
|
||||
case authpkg.StatusPending:
|
||||
// keep polling
|
||||
default:
|
||||
// ParseDeviceFlowStatus normalizes empty+!success to EXPIRED,
|
||||
// so this branch handles truly unknown statuses.
|
||||
fmt.Fprintln(output)
|
||||
return status, "", nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tryOpenBrowser opens url in the default browser; errors are silently ignored.
|
||||
func tryOpenBrowser(url string) error {
|
||||
var cmd *exec.Cmd
|
||||
switch runtime.GOOS {
|
||||
case "darwin":
|
||||
cmd = exec.Command("open", url)
|
||||
case "linux":
|
||||
cmd = exec.Command("xdg-open", url)
|
||||
case "windows":
|
||||
cmd = exec.Command("cmd", "/c", "start", url)
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
return cmd.Start()
|
||||
}
|
||||
@@ -0,0 +1,606 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
|
||||
func TestIsPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected missing_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_PlainString(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user lacks required scope",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user lacks required scope",
|
||||
}
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected plain string with missing_scope to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_NotScopeError(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewValidation("invalid parameter")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("expected validation error NOT to be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if isPatScopeError(nil) {
|
||||
t.Fatal("nil error should not be detected as scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_WithReason(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("API error",
|
||||
apperrors.WithReason("missing_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected error with missing_scope reason to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_InsufficientScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("insufficient_scope for resource",
|
||||
apperrors.WithReason("insufficient_scope"),
|
||||
)
|
||||
if !isPatScopeError(err) {
|
||||
t.Fatal("expected insufficient_scope error to be detected")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_MissingScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.ErrorType != "missing_scope" {
|
||||
t.Errorf("expected error type 'missing_scope', got %q", scopeErr.ErrorType)
|
||||
}
|
||||
if !strings.Contains(scopeErr.Hint, "dws auth login") {
|
||||
t.Errorf("expected hint to contain 'dws auth login', got %q", scopeErr.Hint)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ExtractsScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "missing_scope: user needs calendar:read",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "user needs calendar:read",
|
||||
}
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:read" {
|
||||
t.Errorf("expected MissingScope 'calendar:read', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthError_HumanReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:user_mailbox.message:send",
|
||||
Hint: "run `dws auth login --scope \"mail:user_mailbox.message:send\"` to authorize",
|
||||
MissingScope: "mail:user_mailbox.message:send",
|
||||
}
|
||||
PrintPatAuthError(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, "missing_scope") {
|
||||
t.Errorf("expected output to contain 'missing_scope', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "dws auth login") {
|
||||
t.Errorf("expected output to contain 'dws auth login', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "需要额外授权") {
|
||||
t.Errorf("expected output to contain Chinese auth prompt, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintPatAuthJSON_MachineReadable(t *testing.T) {
|
||||
t.Parallel()
|
||||
var buf strings.Builder
|
||||
scopeErr := &PatScopeError{
|
||||
Identity: "user",
|
||||
ErrorType: "missing_scope",
|
||||
Message: "missing required scope(s): mail:send",
|
||||
Hint: "run dws auth login --scope mail:send",
|
||||
MissingScope: "mail:send",
|
||||
}
|
||||
PrintPatAuthJSON(&buf, scopeErr)
|
||||
|
||||
output := buf.String()
|
||||
if !strings.Contains(output, `"ok": false`) {
|
||||
t.Errorf("expected JSON to contain ok: false, got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, `"missing_scope": "mail:send"`) {
|
||||
t.Errorf("expected JSON to contain missing_scope, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_BusinessPermissionDenied(t *testing.T) {
|
||||
t.Parallel()
|
||||
// Generic business "permission denied" should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("User has no permission to access this mailbox, permission denied")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("generic 'permission denied' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatScopeError_GenericForbidden(t *testing.T) {
|
||||
t.Parallel()
|
||||
// HTTP 403 Forbidden should NOT trigger PAT re-auth.
|
||||
err := apperrors.NewAuth("403 Forbidden")
|
||||
if isPatScopeError(err) {
|
||||
t.Fatal("'403 Forbidden' should not be detected as PAT scope error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_ComplexScope(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth("missing required scope(s): mail:user_mailbox.message:send")
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.MissingScope != "mail:user_mailbox.message:send" {
|
||||
t.Errorf("expected MissingScope 'mail:user_mailbox.message:send', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPatScopeError_Error(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := &PatScopeError{
|
||||
OriginalError: "test error message",
|
||||
}
|
||||
if err.Error() != "test error message" {
|
||||
t.Errorf("expected Error() to return OriginalError, got %q", err.Error())
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow integration tests — httptest mock covering four terminal
|
||||
// states: APPROVED, REJECTED, EXPIRED, CANCELLED (ctx cancel).
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// setupPollServer creates an httptest server that responds to
|
||||
// /cli/oauth/device/poll?flowId=<fid> with the given status sequence.
|
||||
// It also writes the server URL into a temp DWS_CONFIG_DIR/mcp_url so that
|
||||
// GetMCPBaseURL() returns the test server address.
|
||||
func setupPollServer(t *testing.T, statuses []authpkg.DevicePollResponse) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var callCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
idx := int(callCount.Add(1)) - 1
|
||||
if idx >= len(statuses) {
|
||||
idx = len(statuses) - 1
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(statuses[idx])
|
||||
}))
|
||||
|
||||
// Write mcp_url so GetMCPBaseURL picks up the test server.
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Approved(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "APPROVED", AuthCode: "code123"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-1", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED, got %q", status)
|
||||
}
|
||||
if authCode != "code123" {
|
||||
t.Errorf("expected authCode 'code123', got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Rejected(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "REJECTED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-2", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "REJECTED" {
|
||||
t.Errorf("expected REJECTED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for REJECTED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Expired(t *testing.T) {
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: "EXPIRED"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-3", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for EXPIRED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_Cancelled(t *testing.T) {
|
||||
// Server always returns PENDING so context cancellation is the only exit.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
// Cancel immediately after first poll tick.
|
||||
go func() {
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
cancel()
|
||||
}()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-4", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "CANCELLED" {
|
||||
t.Errorf("expected CANCELLED, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for CANCELLED, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// IsPatRetrying tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestIsPatRetrying_Default(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.Background()
|
||||
if IsPatRetrying(ctx) {
|
||||
t.Fatal("expected false for plain context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPatRetrying_WithValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
ctx := context.WithValue(context.Background(), patRetryingKey, true)
|
||||
if !IsPatRetrying(ctx) {
|
||||
t.Fatal("expected true when pat retry key is set")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// pollPatDeviceFlow edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestPollPatDeviceFlow_ServerErrorFallback(t *testing.T) {
|
||||
// When server returns success=false with empty status, should treat as EXPIRED.
|
||||
server, configDir := setupPollServer(t, []authpkg.DevicePollResponse{
|
||||
{Success: false, Data: authpkg.DevicePollData{Status: ""}},
|
||||
})
|
||||
defer server.Close()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, authCode, err := pollPatDeviceFlow(ctx, "flow-err", configDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "EXPIRED" {
|
||||
t.Errorf("expected EXPIRED for server error fallback, got %q", status)
|
||||
}
|
||||
if authCode != "" {
|
||||
t.Errorf("expected empty authCode for server error, got %q", authCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPollPatDeviceFlow_RedirectSkipped(t *testing.T) {
|
||||
// When server returns 302 (SSO redirect), poll should continue until real response.
|
||||
var callCount int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
callCount++
|
||||
if callCount <= 1 {
|
||||
// First call: simulate SSO redirect
|
||||
w.Header().Set("Location", "https://sso.example.com")
|
||||
w.WriteHeader(http.StatusFound)
|
||||
return
|
||||
}
|
||||
// Second call: return APPROVED
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
resp := authpkg.DevicePollResponse{
|
||||
Success: true,
|
||||
Data: authpkg.DevicePollData{Status: "APPROVED"},
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
var buf bytes.Buffer
|
||||
status, _, err := pollPatDeviceFlow(ctx, "flow-redirect", tmpDir, &buf)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status != "APPROVED" {
|
||||
t.Errorf("expected APPROVED after redirect, got %q", status)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// extractPatScopeError edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestExtractPatScopeError_Nil(t *testing.T) {
|
||||
t.Parallel()
|
||||
if got := extractPatScopeError(nil); got != nil {
|
||||
t.Fatalf("expected nil for nil error, got %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPatScopeError_WithIdentity(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := apperrors.NewAuth(`insufficient_scope: identity "app_user" needs calendar:write`)
|
||||
scopeErr := extractPatScopeError(err)
|
||||
if scopeErr == nil {
|
||||
t.Fatal("expected non-nil PatScopeError")
|
||||
}
|
||||
if scopeErr.Identity != "app_user" {
|
||||
t.Errorf("expected Identity 'app_user', got %q", scopeErr.Identity)
|
||||
}
|
||||
if scopeErr.MissingScope != "calendar:write" {
|
||||
t.Errorf("expected MissingScope 'calendar:write', got %q", scopeErr.MissingScope)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// handlePatAuthCheck integration tests — cover the main orchestrator with
|
||||
// mock runner + httptest poll server for APPROVED, REJECTED, EmptyFlowID.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// mockRunner is a simple executor.Runner for testing handlePatAuthCheck.
|
||||
type mockRunner struct {
|
||||
runFunc func(ctx context.Context, inv executor.Invocation) (executor.Result, error)
|
||||
}
|
||||
|
||||
func (m *mockRunner) Run(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
return m.runFunc(ctx, inv)
|
||||
}
|
||||
|
||||
// setupHandlePATServer creates an httptest server for handlePatAuthCheck tests.
|
||||
// It responds to device poll requests with the given status after the first poll.
|
||||
func setupHandlePATServer(t *testing.T, terminalStatus string, authCode string) (*httptest.Server, string) {
|
||||
t.Helper()
|
||||
var pollCount atomic.Int32
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if strings.Contains(r.URL.Path, "/cli/oauth/device/poll") {
|
||||
idx := int(pollCount.Add(1)) - 1
|
||||
var resp authpkg.DevicePollResponse
|
||||
if idx == 0 {
|
||||
resp = authpkg.DevicePollResponse{Success: true, Data: authpkg.DevicePollData{Status: "PENDING"}}
|
||||
} else {
|
||||
resp = authpkg.DevicePollResponse{
|
||||
Success: terminalStatus == "APPROVED",
|
||||
Data: authpkg.DevicePollData{Status: terminalStatus, AuthCode: authCode},
|
||||
}
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_ = json.NewEncoder(w).Encode(resp)
|
||||
return
|
||||
}
|
||||
http.NotFound(w, r)
|
||||
}))
|
||||
|
||||
tmpDir := t.TempDir()
|
||||
_ = os.WriteFile(filepath.Join(tmpDir, "mcp_url"), []byte(server.URL), 0644)
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
return server, tmpDir
|
||||
}
|
||||
|
||||
func makePATErrorJSON(flowID, clientID string) string {
|
||||
type patData struct {
|
||||
Desc string `json:"desc"`
|
||||
FlowID string `json:"flowId"`
|
||||
URI string `json:"uri"`
|
||||
ClientID string `json:"clientId"`
|
||||
}
|
||||
payload := struct {
|
||||
Code string `json:"code"`
|
||||
Data patData `json:"data"`
|
||||
}{
|
||||
Code: "AGENT_CODE_NOT_EXISTS",
|
||||
Data: patData{
|
||||
Desc: "test auth",
|
||||
FlowID: flowID,
|
||||
URI: "", // empty to avoid opening browser in test
|
||||
ClientID: clientID,
|
||||
},
|
||||
}
|
||||
data, _ := json.Marshal(payload)
|
||||
return string(data)
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Approved(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "APPROVED", "test-auth-code")
|
||||
defer server.Close()
|
||||
|
||||
var retryCalled bool
|
||||
var retryHasKey bool
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
retryCalled = true
|
||||
retryHasKey = IsPatRetrying(ctx)
|
||||
return executor.Result{Response: map[string]any{"ok": true}}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-approved", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !retryCalled {
|
||||
t.Fatal("expected mock runner to be called for retry")
|
||||
}
|
||||
if !retryHasKey {
|
||||
t.Fatal("expected retry context to have patRetryingKey")
|
||||
}
|
||||
// Verify SetClientIDFromMCP was called with the PAT response clientId.
|
||||
if cid := authpkg.ClientID(); cid != "test-client-id" {
|
||||
t.Errorf("expected ClientID 'test-client-id', got %q", cid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_Rejected(t *testing.T) {
|
||||
server, configDir := setupHandlePATServer(t, "REJECTED", "")
|
||||
defer server.Close()
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called on REJECTED")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("flow-rejected", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, configDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for REJECTED")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "用户已拒绝授权") {
|
||||
t.Errorf("expected rejection error, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHandlePatAuthCheck_EmptyFlowID_FallsBackToPATError(t *testing.T) {
|
||||
// No poll server needed — empty flowId means no polling, return PATError directly.
|
||||
tmpDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", tmpDir)
|
||||
|
||||
mock := &mockRunner{
|
||||
runFunc: func(ctx context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
t.Fatal("runner should not be called when flowId is empty")
|
||||
return executor.Result{}, nil
|
||||
},
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{fallback: mock}
|
||||
patErr := &apperrors.PATError{RawJSON: makePATErrorJSON("", "test-client-id")}
|
||||
|
||||
ctx := context.Background()
|
||||
var buf bytes.Buffer
|
||||
_, err := handlePatAuthCheck(ctx, runner, executor.Invocation{
|
||||
CanonicalProduct: "test",
|
||||
Tool: "test_tool",
|
||||
}, patErr, tmpDir, &buf)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected PATError when flowId is empty")
|
||||
}
|
||||
// Should return the original PATError.
|
||||
if _, ok := err.(*apperrors.PATError); !ok {
|
||||
t.Errorf("expected *PATError, got %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,664 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newPluginCommand() *cobra.Command {
|
||||
pluginCmd := newPlaceholderParent("plugin", i18n.T("插件管理"))
|
||||
|
||||
pluginCmd.AddCommand(
|
||||
newPluginListCommand(),
|
||||
newPluginInstallCommand(),
|
||||
newPluginInfoCommand(),
|
||||
newPluginEnableCommand(),
|
||||
newPluginDisableCommand(),
|
||||
newPluginRemoveCommand(),
|
||||
newPluginValidateCommand(),
|
||||
newPluginCreateCommand(),
|
||||
newPluginDevCommand(),
|
||||
newPluginConfigCommand(),
|
||||
newPluginBuildCommand(),
|
||||
)
|
||||
|
||||
return pluginCmd
|
||||
}
|
||||
|
||||
func newPluginListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("列出已安装的插件"),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
if wantJSON {
|
||||
return output.WriteJSON(cmd.OutOrStdout(), plugins)
|
||||
}
|
||||
|
||||
if len(plugins) == 0 {
|
||||
fmt.Fprintln(cmd.OutOrStdout(), "No plugins installed.")
|
||||
return nil
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
"NAME", "VERSION", "TYPE", "STATUS", "DESCRIPTION")
|
||||
fmt.Fprintln(w, strings.Repeat("-", 85))
|
||||
for _, p := range plugins {
|
||||
fmt.Fprintf(w, "%-35s %-12s %-10s %-10s %s\n",
|
||||
p.Name, p.Version, p.Type, statusStr(p.Enabled), p.Description)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install",
|
||||
Short: i18n.T("安装插件"),
|
||||
Example: ` dws plugin install --dir ./conference
|
||||
dws plugin install --git https://github.com/DingTalk-Real-AI/conference.git`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dirPath, _ := cmd.Flags().GetString("dir")
|
||||
gitURL, _ := cmd.Flags().GetString("git")
|
||||
|
||||
if dirPath == "" && gitURL == "" {
|
||||
return apperrors.NewValidation("specify install source: --dir <path> or --git <url>")
|
||||
}
|
||||
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if gitURL != "" {
|
||||
p, err := loader.InstallFromGit(gitURL)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
}
|
||||
|
||||
p, err := loader.InstallFromDir(dirPath)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("install failed: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Installed %s (%s)\n", p.Manifest.Name, p.Manifest.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("dir", "", "Install from a local directory")
|
||||
cmd.Flags().String("git", "", "Install from a Git repository")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginInfoCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "info <name>",
|
||||
Short: i18n.T("查看插件详情"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
plugins := loader.ListInstalled()
|
||||
|
||||
for _, p := range plugins {
|
||||
if p.Name == name {
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Name: %s\n", p.Name)
|
||||
fmt.Fprintf(w, "Version: %s\n", p.Version)
|
||||
fmt.Fprintf(w, "Type: %s\n", p.Type)
|
||||
fmt.Fprintf(w, "Status: %s\n", statusStr(p.Enabled))
|
||||
fmt.Fprintf(w, "Path: %s\n", p.Path)
|
||||
if p.Description != "" {
|
||||
fmt.Fprintf(w, "Description: %s\n", p.Description)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found", name))
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginEnableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "enable <name>",
|
||||
Short: i18n.T("启用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], true); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s enabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginDisableCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "disable <name>",
|
||||
Short: i18n.T("禁用插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.SetEnabled(args[0], false); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s disabled.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginRemoveCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "remove <name>",
|
||||
Short: i18n.T("卸载已安装的插件"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
// Stop stdio clients before removing to release file locks
|
||||
StopStdioClientsByPlugin(args[0])
|
||||
keepData, _ := cmd.Flags().GetBool("keep-data")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
if err := loader.RemovePlugin(args[0], keepData); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Plugin %s removed.\n", args[0])
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("keep-data", false, "Keep plugin data directory")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginValidateCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "validate <dir>",
|
||||
Short: i18n.T("校验 plugin.json"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
m, err := plugin.ParseManifest(dir + "/plugin.json")
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("parse failed: %v", err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Valid: %s (%s)\n", m.Name, m.Version)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginCreateCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create <name>",
|
||||
Short: i18n.T("脚手架生成新插件目录"),
|
||||
Example: ` dws plugin create my-tool
|
||||
dws plugin create my-tool --description "My awesome tool"`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name := args[0]
|
||||
desc, _ := cmd.Flags().GetString("description")
|
||||
pluginType := "user"
|
||||
|
||||
// Validate name format
|
||||
m := &plugin.Manifest{Name: name, Version: "0.1.0", Type: pluginType}
|
||||
if err := m.Validate(""); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin name: %v", err))
|
||||
}
|
||||
|
||||
dir := filepath.Join(".", name)
|
||||
if _, err := os.Stat(dir); err == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("directory %q already exists", dir))
|
||||
}
|
||||
|
||||
// Create directory structure
|
||||
dirs := []string{
|
||||
dir,
|
||||
filepath.Join(dir, "skills", name),
|
||||
filepath.Join(dir, "hooks"),
|
||||
}
|
||||
for _, d := range dirs {
|
||||
if err := os.MkdirAll(d, 0o755); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create directory: %v", err))
|
||||
}
|
||||
}
|
||||
|
||||
// Write plugin.json
|
||||
pluginJSON := fmt.Sprintf(`{
|
||||
"name": %q,
|
||||
"version": "0.1.0",
|
||||
"description": %q,
|
||||
"type": %q,
|
||||
"minCLIVersion": %q,
|
||||
"mcpServers": {
|
||||
%q: {
|
||||
"type": "stdio",
|
||||
"command": "${DWS_PLUGIN_ROOT}/bin/server",
|
||||
"args": []
|
||||
}
|
||||
},
|
||||
"build": {
|
||||
"command": "echo 'TODO: replace with your build command, e.g.: bun build --compile src/server.ts --outfile bin/server'",
|
||||
"output": "bin/server"
|
||||
},
|
||||
"skills": "./skills/",
|
||||
"hooks": "./hooks/hooks.json"
|
||||
}
|
||||
`, name, desc, pluginType, RawVersion(), name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "plugin.json"), []byte(pluginJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write plugin.json: %v", err))
|
||||
}
|
||||
|
||||
// Write SKILL.md template
|
||||
skillMD := fmt.Sprintf(`---
|
||||
name: %s
|
||||
description: %s
|
||||
cli_version: ">=%s"
|
||||
---
|
||||
|
||||
# %s
|
||||
|
||||
## Intent Recognition
|
||||
|
||||
Use this skill when the user mentions:
|
||||
- TODO: add your intent keywords here
|
||||
|
||||
## Command Decision Tree
|
||||
|
||||
| User Intent | Command | Required Parameters |
|
||||
|-------------|---------|---------------------|
|
||||
| TODO | `+"`dws %s <sub-command>`"+` | `+"`--param`"+` |
|
||||
|
||||
## Parameter Rules
|
||||
|
||||
### TODO: parameter type
|
||||
- Format description
|
||||
- Conversion rules
|
||||
`, name, desc, RawVersion(), name, name)
|
||||
|
||||
if err := os.WriteFile(filepath.Join(dir, "skills", name, "SKILL.md"), []byte(skillMD), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write SKILL.md: %v", err))
|
||||
}
|
||||
|
||||
// Write hooks.json template
|
||||
hooksJSON := `{
|
||||
"hooks": []
|
||||
}
|
||||
`
|
||||
if err := os.WriteFile(filepath.Join(dir, "hooks", "hooks.json"), []byte(hooksJSON), 0o644); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to write hooks.json: %v", err))
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "Created plugin scaffold at ./%s/\n", name)
|
||||
fmt.Fprintf(w, " %s/\n", name)
|
||||
fmt.Fprintf(w, " ├── plugin.json\n")
|
||||
fmt.Fprintf(w, " ├── skills/%s/SKILL.md\n", name)
|
||||
fmt.Fprintf(w, " └── hooks/hooks.json\n")
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintf(w, "Next steps:\n")
|
||||
fmt.Fprintf(w, " 1. Edit plugin.json to configure your MCP servers\n")
|
||||
fmt.Fprintf(w, " 2. Edit skills/%s/SKILL.md to describe your commands\n", name)
|
||||
fmt.Fprintf(w, " 3. Run: dws plugin validate ./%s\n", name)
|
||||
fmt.Fprintf(w, " 4. Run: dws plugin dev ./%s\n", name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("description", "", "Plugin description")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginDevCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "dev <dir>",
|
||||
Short: i18n.T("将本地目录注册为开发态插件"),
|
||||
Long: `Registers a plugin from a local source directory for development.
|
||||
The plugin is loaded directly from the source directory on next CLI invocation,
|
||||
without copying files to ~/.dws/plugins/. Use 'dws plugin dev --off <name>'
|
||||
to unregister.`,
|
||||
Example: ` dws plugin dev ./my-tool
|
||||
dws plugin dev --off my-tool`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
off, _ := cmd.Flags().GetBool("off")
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if off {
|
||||
// Unregister dev plugin
|
||||
name := args[0]
|
||||
if err := loader.UnregisterDevPlugin(name); err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q unregistered.\n", name)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Register dev plugin
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
// Validate the plugin first
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
if err := m.Validate(RawVersion()); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("validation failed: %v", err))
|
||||
}
|
||||
|
||||
if err := loader.RegisterDevPlugin(m.Name, absDir); err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to register: %v", err))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Dev plugin %q registered from %s\n", m.Name, absDir)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "It will be loaded on next dws invocation.\n")
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "To unregister: dws plugin dev --off %s\n", m.Name)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("off", false, "Unregister a dev plugin")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigCommand() *cobra.Command {
|
||||
configCmd := newPlaceholderParent("config", i18n.T("管理插件配置"))
|
||||
configCmd.AddCommand(
|
||||
newPluginConfigSetCommand(),
|
||||
newPluginConfigGetCommand(),
|
||||
newPluginConfigListCommand(),
|
||||
newPluginConfigUnsetCommand(),
|
||||
)
|
||||
return configCmd
|
||||
}
|
||||
|
||||
func newPluginConfigSetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "set <plugin-name> <key> <value>",
|
||||
Short: i18n.T("设置插件配置项"),
|
||||
Long: `Persistently set a configuration value for a plugin.
|
||||
The value is stored in ~/.dws/settings.json and automatically injected
|
||||
as an environment variable when the plugin is loaded.
|
||||
|
||||
Environment variables set by the user (e.g. via export) take precedence
|
||||
over values stored in settings.json.`,
|
||||
Example: ` dws plugin config set demo-devtool DASHSCOPE_API_KEY sk-xxx
|
||||
dws plugin config set my-plugin API_ENDPOINT https://api.example.com`,
|
||||
Args: cobra.ExactArgs(3),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key, value := args[0], args[1], args[2]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// Validate that the plugin exists.
|
||||
plugins := loader.ListInstalled()
|
||||
found := false
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return apperrors.NewValidation(fmt.Sprintf("plugin %q not found; use 'dws plugin list' to see installed plugins", pluginName))
|
||||
}
|
||||
|
||||
loader.SetPluginConfig(pluginName, key, value)
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config saved: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigGetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "get <plugin-name> <key>",
|
||||
Short: i18n.T("读取插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
val, ok := loader.GetPluginConfig(pluginName, key)
|
||||
if !ok {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintln(cmd.OutOrStdout(), val)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newPluginConfigListCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list <plugin-name>",
|
||||
Short: i18n.T("列出插件所有配置项"),
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName := args[0]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
wantJSON, _ := cmd.Flags().GetBool("json")
|
||||
configs := loader.ListPluginConfig(pluginName)
|
||||
|
||||
// Also load the plugin manifest to show declared userConfig keys.
|
||||
declaredKeys := loadDeclaredUserConfig(loader, pluginName)
|
||||
|
||||
if wantJSON {
|
||||
result := make(map[string]any)
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
if sensitive {
|
||||
result[k] = maskSensitiveValue(v)
|
||||
} else {
|
||||
result[k] = v
|
||||
}
|
||||
}
|
||||
// Include declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
entry := map[string]any{
|
||||
"value": nil,
|
||||
"description": ci.Description,
|
||||
"required": ci.Default == "",
|
||||
}
|
||||
result[k] = entry
|
||||
}
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"kind": "plugin_config",
|
||||
"plugin": pluginName,
|
||||
"config": result,
|
||||
})
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
if len(configs) == 0 && len(declaredKeys) == 0 {
|
||||
fmt.Fprintf(w, "No configuration for plugin %q.\n", pluginName)
|
||||
return nil
|
||||
}
|
||||
|
||||
fmt.Fprintf(w, "Configuration for %s:\n\n", pluginName)
|
||||
|
||||
// Show set values.
|
||||
for k, v := range configs {
|
||||
sensitive := false
|
||||
if ci, ok := declaredKeys[k]; ok {
|
||||
sensitive = ci.Sensitive
|
||||
}
|
||||
displayVal := v
|
||||
if sensitive {
|
||||
displayVal = maskSensitiveValue(v)
|
||||
}
|
||||
fmt.Fprintf(w, " %s = %s\n", k, displayVal)
|
||||
}
|
||||
|
||||
// Show declared but unset keys.
|
||||
for k, ci := range declaredKeys {
|
||||
if _, set := configs[k]; !set {
|
||||
desc := ""
|
||||
if ci.Description != "" {
|
||||
desc = " # " + ci.Description
|
||||
}
|
||||
fmt.Fprintf(w, " %s = (not set)%s\n", k, desc)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Output in JSON format")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newPluginConfigUnsetCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "unset <plugin-name> <key>",
|
||||
Short: i18n.T("删除插件配置项"),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
pluginName, key := args[0], args[1]
|
||||
loader := plugin.NewLoader(RawVersion())
|
||||
|
||||
if !loader.UnsetPluginConfig(pluginName, key) {
|
||||
return apperrors.NewValidation(fmt.Sprintf("config key %q not set for plugin %q", key, pluginName))
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Config removed: %s.%s\n", pluginName, key)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// loadDeclaredUserConfig loads the userConfig section from a plugin's manifest.
|
||||
func loadDeclaredUserConfig(loader *plugin.Loader, pluginName string) map[string]plugin.ConfigItem {
|
||||
plugins := loader.ListInstalled()
|
||||
for _, p := range plugins {
|
||||
if p.Name == pluginName {
|
||||
m, err := plugin.ParseManifest(filepath.Join(p.Path, "plugin.json"))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return m.UserConfig
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// maskSensitiveValue masks a sensitive value, showing only the first 4
|
||||
// and last 2 characters for values longer than 8 characters.
|
||||
func maskSensitiveValue(value string) string {
|
||||
if len(value) <= 8 {
|
||||
return strings.Repeat("*", len(value))
|
||||
}
|
||||
return value[:4] + strings.Repeat("*", len(value)-6) + value[len(value)-2:]
|
||||
}
|
||||
|
||||
func newPluginBuildCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "build <dir>",
|
||||
Short: i18n.T("将插件 stdio server 编译为原生二进制"),
|
||||
Long: `Runs the build command declared in plugin.json to compile the
|
||||
plugin's server into a single executable. This ensures plugin users
|
||||
don't need any language runtime (Node.js, Python, etc.) installed.
|
||||
|
||||
The build configuration is read from the "build" field in plugin.json:
|
||||
|
||||
{
|
||||
"build": {
|
||||
"command": "bun build --compile src/server.ts --outfile bin/server",
|
||||
"output": "bin/server"
|
||||
}
|
||||
}`,
|
||||
Example: ` dws plugin build ./my-plugin
|
||||
dws plugin build .`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
dir := args[0]
|
||||
absDir, err := filepath.Abs(dir)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid path: %v", err))
|
||||
}
|
||||
|
||||
m, err := plugin.ParseManifest(filepath.Join(absDir, "plugin.json"))
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid plugin at %s: %v", dir, err))
|
||||
}
|
||||
|
||||
if m.Build == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf(
|
||||
"plugin %q has no \"build\" field in plugin.json.\n"+
|
||||
"Add a build config, e.g.:\n\n"+
|
||||
" \"build\": {\n"+
|
||||
" \"command\": \"bun build --compile src/server.js --outfile bin/server\",\n"+
|
||||
" \"output\": \"bin/server\"\n"+
|
||||
" }", m.Name))
|
||||
}
|
||||
|
||||
if err := plugin.BuildPlugin(absDir); err != nil {
|
||||
return apperrors.NewInternal(err.Error())
|
||||
}
|
||||
|
||||
fmt.Fprintf(cmd.OutOrStdout(), "Build succeeded: %s\n", m.Build.Output)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func statusStr(enabled bool) string {
|
||||
if enabled {
|
||||
return "enabled"
|
||||
}
|
||||
return "disabled"
|
||||
}
|
||||
@@ -0,0 +1,198 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
// TestSharedCacheStoreConcurrentSaveTools verifies that a single *cache.Store
|
||||
// instance is safe for goroutines saving tool snapshots concurrently, as long
|
||||
// as each goroutine targets a distinct (partition, serverKey). This mirrors
|
||||
// the real plugin discovery path where each goroutine owns one plugin/server.
|
||||
//
|
||||
// Each call serializes to its own "<key>.json.tmp" file followed by a
|
||||
// rename(2) to the final path, so concurrent writers targeting distinct keys
|
||||
// never collide. The invariant asserted here: after N parallel writes, the
|
||||
// Store returns each written snapshot intact under LoadTools.
|
||||
func TestSharedCacheStoreConcurrentSaveTools(t *testing.T) {
|
||||
const (
|
||||
partition = "default/default"
|
||||
writers = 16
|
||||
)
|
||||
|
||||
store := cache.NewStore(t.TempDir())
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < writers; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", idx)
|
||||
if err := store.SaveTools(partition, key, cache.ToolsSnapshot{
|
||||
ServerKey: key,
|
||||
}); err != nil {
|
||||
t.Errorf("SaveTools(%s): %v", key, err)
|
||||
}
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < writers; i++ {
|
||||
key := fmt.Sprintf("plugin:concurrent:%d", i)
|
||||
snapshot, _, err := store.LoadTools(partition, key)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTools(%s): %v", key, err)
|
||||
}
|
||||
if snapshot.ServerKey != key {
|
||||
t.Errorf("LoadTools(%s) returned ServerKey %q", key, snapshot.ServerKey)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestAppendDynamicServerConcurrent exercises the dynamicMu mutex on the
|
||||
// write path by spraying distinct server descriptors in parallel. Afterwards
|
||||
// every injected product ID must be resolvable — a missing entry would
|
||||
// indicate a lost write through an un-synchronized map update.
|
||||
func TestAppendDynamicServerConcurrent(t *testing.T) {
|
||||
dynamicMu.Lock()
|
||||
prev := struct {
|
||||
endpoints map[string]string
|
||||
products map[string]bool
|
||||
aliases map[string]string
|
||||
toolEndpoints map[string]string
|
||||
}{dynamicEndpoints, dynamicProducts, dynamicAliases, dynamicToolEndpoints}
|
||||
dynamicEndpoints = nil
|
||||
dynamicProducts = nil
|
||||
dynamicAliases = nil
|
||||
dynamicToolEndpoints = nil
|
||||
dynamicMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
dynamicMu.Lock()
|
||||
dynamicEndpoints = prev.endpoints
|
||||
dynamicProducts = prev.products
|
||||
dynamicAliases = prev.aliases
|
||||
dynamicToolEndpoints = prev.toolEndpoints
|
||||
dynamicMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
id := fmt.Sprintf("plugin-id-%d", idx)
|
||||
endpoint := fmt.Sprintf("https://example.test/%d", idx)
|
||||
AppendDynamicServer(market.ServerDescriptor{
|
||||
Endpoint: endpoint,
|
||||
CLI: market.CLIOverlay{
|
||||
ID: id,
|
||||
Command: id,
|
||||
},
|
||||
})
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
id := fmt.Sprintf("plugin-id-%d", i)
|
||||
if endpoint, ok := directRuntimeEndpoint(id, ""); !ok || endpoint == "" {
|
||||
t.Errorf("directRuntimeEndpoint(%q) = (%q, %v), want non-empty", id, endpoint, ok)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegisterStdioClientConcurrent verifies the stdioMu-protected registry
|
||||
// survives concurrent writers — every registered client must be looked up
|
||||
// afterwards. Uses nil client pointers since LookupStdioClient only compares
|
||||
// keys, not values.
|
||||
func TestRegisterStdioClientConcurrent(t *testing.T) {
|
||||
stdioMu.Lock()
|
||||
prev := stdioClients
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
stdioMu.Unlock()
|
||||
t.Cleanup(func() {
|
||||
stdioMu.Lock()
|
||||
stdioClients = prev
|
||||
stdioMu.Unlock()
|
||||
})
|
||||
|
||||
const n = 32
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < n; i++ {
|
||||
wg.Add(1)
|
||||
go func(idx int) {
|
||||
defer wg.Done()
|
||||
RegisterStdioClient(fmt.Sprintf("plugin/%d", idx), nil)
|
||||
}(i)
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
for i := 0; i < n; i++ {
|
||||
key := fmt.Sprintf("plugin/%d", i)
|
||||
if _, ok := LookupStdioClient(key); !ok {
|
||||
t.Errorf("LookupStdioClient(%q) missing after concurrent registration", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestResolvePluginColdTimeouts covers the three code paths of the env
|
||||
// parser: unset (defaults), valid duration (applied to all three slots),
|
||||
// and invalid duration (logged and ignored, defaults returned).
|
||||
func TestResolvePluginColdTimeouts(t *testing.T) {
|
||||
t.Run("defaults when env unset", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("httpNoAuth = %v, want 1s", got.httpNoAuth)
|
||||
}
|
||||
if got.httpAuth != 1500*time.Millisecond {
|
||||
t.Errorf("httpAuth = %v, want 1.5s", got.httpAuth)
|
||||
}
|
||||
if got.stdio != 2*time.Second {
|
||||
t.Errorf("stdio = %v, want 2s", got.stdio)
|
||||
}
|
||||
})
|
||||
t.Run("env override applies to all slots", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "3500ms")
|
||||
got := resolvePluginColdTimeouts()
|
||||
want := 3500 * time.Millisecond
|
||||
if got.httpNoAuth != want || got.httpAuth != want || got.stdio != want {
|
||||
t.Errorf("override not propagated: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("invalid env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "not-a-duration")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second || got.stdio != 2*time.Second {
|
||||
t.Errorf("invalid env should not override defaults: %+v", got)
|
||||
}
|
||||
})
|
||||
t.Run("non-positive env falls back to defaults", func(t *testing.T) {
|
||||
t.Setenv(cli.PluginColdTimeoutEnv, "0")
|
||||
got := resolvePluginColdTimeouts()
|
||||
if got.httpNoAuth != 1*time.Second {
|
||||
t.Errorf("zero duration should not override defaults, got %v", got.httpNoAuth)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,501 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"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/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newRecoveryCommand(_ context.Context, loader cli.CatalogLoader, flags *GlobalFlags) *cobra.Command {
|
||||
var (
|
||||
planUseLast bool
|
||||
planEventID string
|
||||
executeUseLast bool
|
||||
executeEventID string
|
||||
finalEventID string
|
||||
finalOutcome string
|
||||
executionFile string
|
||||
)
|
||||
|
||||
runtime := newRecoveryRuntime(loader, flags)
|
||||
|
||||
cmd := &cobra.Command{
|
||||
Use: "recovery",
|
||||
Short: "错误恢复辅助命令",
|
||||
Long: "读取失败快照,生成恢复分析,并回写恢复结果。",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
planCmd := &cobra.Command{
|
||||
Use: "plan",
|
||||
Short: "基于失败快照生成恢复计划",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
store := recovery.NewStore(defaultConfigDir())
|
||||
last, err := loadRecoverySnapshot(store, planUseLast, planEventID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
planner := recovery.NewPlanner(runtime)
|
||||
plan := planner.PlanWithOptions(cmd.Context(), last.Context, recovery.PlanOptions{
|
||||
EventID: last.EventID,
|
||||
EnableDocSearch: true,
|
||||
})
|
||||
recovery.HydratePlanForEvent(last.EventID, last.Context, last.Replay, &plan)
|
||||
if err := store.SavePlan(last.EventID, plan); err != nil {
|
||||
return fmt.Errorf("保存恢复计划失败: %w", err)
|
||||
}
|
||||
|
||||
payload := map[string]any{
|
||||
"event_id": last.EventID,
|
||||
"context": last.Context,
|
||||
"plan": plan,
|
||||
}
|
||||
return output.WriteCommandPayload(cmd, payload, output.FormatJSON)
|
||||
},
|
||||
}
|
||||
planCmd.Flags().BoolVar(&planUseLast, "last", false, "读取最近一次失败快照")
|
||||
planCmd.Flags().StringVar(&planEventID, "event-id", "", "按 event_id 读取失败快照")
|
||||
|
||||
executeCmd := &cobra.Command{
|
||||
Use: "execute",
|
||||
Short: "生成面向 Agent 的恢复分析包",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
store := recovery.NewStore(defaultConfigDir())
|
||||
last, err := loadRecoverySnapshot(store, executeUseLast, executeEventID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
planner := recovery.NewPlanner(runtime)
|
||||
executor := recovery.NewExecutor(planner, runtime)
|
||||
bundle := executor.Execute(cmd.Context(), *last)
|
||||
if err := store.SaveAnalysis(last.EventID, bundle.Plan, bundle); err != nil {
|
||||
return fmt.Errorf("保存恢复分析失败: %w", err)
|
||||
}
|
||||
|
||||
return output.WriteCommandPayload(cmd, bundle, output.FormatJSON)
|
||||
},
|
||||
}
|
||||
executeCmd.Flags().BoolVar(&executeUseLast, "last", false, "读取最近一次失败快照")
|
||||
executeCmd.Flags().StringVar(&executeEventID, "event-id", "", "按 event_id 读取失败快照")
|
||||
|
||||
finalizeCmd := &cobra.Command{
|
||||
Use: "finalize",
|
||||
Short: "回写恢复闭环结果",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
if strings.TrimSpace(finalEventID) == "" {
|
||||
return fmt.Errorf("必须提供 --event-id")
|
||||
}
|
||||
if strings.TrimSpace(finalOutcome) == "" {
|
||||
return fmt.Errorf("必须提供 --outcome")
|
||||
}
|
||||
switch finalOutcome {
|
||||
case "recovered", "failed", "handoff":
|
||||
default:
|
||||
return fmt.Errorf("--outcome 仅支持 recovered|failed|handoff")
|
||||
}
|
||||
|
||||
store := recovery.NewStore(defaultConfigDir())
|
||||
var execution *recovery.RecoveryExecution
|
||||
if strings.TrimSpace(executionFile) != "" {
|
||||
loaded, err := loadRecoveryExecution(executionFile)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
execution = &loaded
|
||||
}
|
||||
if err := store.Finalize(finalEventID, finalOutcome, execution); err != nil {
|
||||
return fmt.Errorf("回写恢复结果失败: %w", err)
|
||||
}
|
||||
|
||||
payload := map[string]any{
|
||||
"event_id": finalEventID,
|
||||
"outcome": finalOutcome,
|
||||
"success": true,
|
||||
}
|
||||
if execution != nil {
|
||||
payload["execution_recorded"] = true
|
||||
}
|
||||
return output.WriteCommandPayload(cmd, payload, output.FormatJSON)
|
||||
},
|
||||
}
|
||||
finalizeCmd.Flags().StringVar(&finalEventID, "event-id", "", "恢复事件 ID")
|
||||
finalizeCmd.Flags().StringVar(&finalOutcome, "outcome", "", "恢复结果: recovered|failed|handoff")
|
||||
finalizeCmd.Flags().StringVar(&executionFile, "execution-file", "", "Agent 执行详情 JSON 文件")
|
||||
|
||||
cmd.AddCommand(planCmd, executeCmd, finalizeCmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func loadRecoverySnapshot(store *recovery.Store, useLast bool, eventID string) (*recovery.LastError, error) {
|
||||
if useLast && strings.TrimSpace(eventID) != "" {
|
||||
return nil, fmt.Errorf("--last 和 --event-id 不能同时使用")
|
||||
}
|
||||
switch {
|
||||
case useLast:
|
||||
last, err := store.LoadLastError()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取失败快照失败: %w", err)
|
||||
}
|
||||
return last, nil
|
||||
case strings.TrimSpace(eventID) != "":
|
||||
last, err := store.LoadErrorByEvent(strings.TrimSpace(eventID))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取失败快照失败: %w", err)
|
||||
}
|
||||
return last, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("必须通过 --last 或 --event-id 指定失败快照")
|
||||
}
|
||||
}
|
||||
|
||||
func loadRecoveryExecution(path string) (recovery.RecoveryExecution, error) {
|
||||
var execution recovery.RecoveryExecution
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return execution, fmt.Errorf("读取恢复执行详情失败: %w", err)
|
||||
}
|
||||
var payload recoveryExecutionPayload
|
||||
if err := json.Unmarshal(data, &payload); err != nil {
|
||||
return execution, fmt.Errorf("解析恢复执行详情失败: %w", err)
|
||||
}
|
||||
execution.Actions = append([]string(nil), payload.Actions...)
|
||||
if len(execution.Actions) == 0 && strings.TrimSpace(payload.Action) != "" {
|
||||
execution.Actions = []string{strings.TrimSpace(payload.Action)}
|
||||
}
|
||||
execution.Result = strings.TrimSpace(payload.Result)
|
||||
execution.ErrorSummary = strings.TrimSpace(payload.ErrorSummary)
|
||||
if execution.ErrorSummary == "" {
|
||||
execution.ErrorSummary = strings.TrimSpace(payload.Error)
|
||||
}
|
||||
|
||||
attempts, err := decodeRecoveryAttempts(payload.Attempts, execution.Actions, execution.Result, execution.ErrorSummary)
|
||||
if err != nil {
|
||||
return execution, fmt.Errorf("解析恢复执行详情失败: %w", err)
|
||||
}
|
||||
if len(attempts) == 0 && payload.Attempt > 0 {
|
||||
attempts = legacyRecoveryAttempts(payload.Attempt, execution.Actions, execution.Result, execution.ErrorSummary)
|
||||
}
|
||||
execution.Attempts = attempts
|
||||
return execution, nil
|
||||
}
|
||||
|
||||
type recoveryExecutionPayload struct {
|
||||
Action string `json:"action,omitempty"`
|
||||
Actions []string `json:"actions,omitempty"`
|
||||
Attempt int `json:"attempt,omitempty"`
|
||||
Attempts json.RawMessage `json:"attempts,omitempty"`
|
||||
Result string `json:"result,omitempty"`
|
||||
Error string `json:"error,omitempty"`
|
||||
ErrorSummary string `json:"error_summary,omitempty"`
|
||||
}
|
||||
|
||||
func decodeRecoveryAttempts(raw json.RawMessage, actions []string, result, errorSummary string) ([]recovery.RecoveryAttempt, error) {
|
||||
trimmed := strings.TrimSpace(string(raw))
|
||||
if trimmed == "" || trimmed == "null" {
|
||||
return nil, nil
|
||||
}
|
||||
if strings.HasPrefix(trimmed, "[") {
|
||||
var attempts []recovery.RecoveryAttempt
|
||||
if err := json.Unmarshal(raw, &attempts); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return attempts, nil
|
||||
}
|
||||
|
||||
var count int
|
||||
if err := json.Unmarshal(raw, &count); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return legacyRecoveryAttempts(count, actions, result, errorSummary), nil
|
||||
}
|
||||
|
||||
func legacyRecoveryAttempts(count int, actions []string, result, errorSummary string) []recovery.RecoveryAttempt {
|
||||
if count <= 0 {
|
||||
return nil
|
||||
}
|
||||
summary := strings.TrimSpace(strings.Join(actions, ", "))
|
||||
if summary == "" {
|
||||
summary = "legacy execution attempt"
|
||||
}
|
||||
attempts := make([]recovery.RecoveryAttempt, 0, count)
|
||||
for i := 0; i < count; i++ {
|
||||
attempts = append(attempts, recovery.RecoveryAttempt{
|
||||
CommandSummary: summary,
|
||||
Result: result,
|
||||
ErrorSummary: errorSummary,
|
||||
Source: "legacy_execution_file",
|
||||
})
|
||||
}
|
||||
return attempts
|
||||
}
|
||||
|
||||
type recoveryRuntime struct {
|
||||
loader cli.CatalogLoader
|
||||
transport *transport.Client
|
||||
flags *GlobalFlags
|
||||
}
|
||||
|
||||
func newRecoveryRuntime(loader cli.CatalogLoader, flags *GlobalFlags) *recoveryRuntime {
|
||||
var httpClient *http.Client
|
||||
if flags != nil && flags.Timeout > 0 {
|
||||
httpClient = &http.Client{Timeout: time.Duration(flags.Timeout) * time.Second}
|
||||
}
|
||||
client := transport.NewClient(httpClient)
|
||||
client.ExtraHeaders = resolveIdentityHeaders()
|
||||
return &recoveryRuntime{
|
||||
loader: loader,
|
||||
transport: client,
|
||||
flags: flags,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *recoveryRuntime) Search(ctx context.Context, query string, rc recovery.RecoveryContext) (recovery.KnowledgeRetrieval, error) {
|
||||
const (
|
||||
searchPage = 1
|
||||
searchSize = 5
|
||||
)
|
||||
requestArgs := map[string]any{
|
||||
"keyword": query,
|
||||
"page": searchPage,
|
||||
"size": searchSize,
|
||||
}
|
||||
|
||||
retrieval := recovery.KnowledgeRetrieval{
|
||||
DocSearch: recovery.DocSearch{
|
||||
Provider: "open_platform_docs",
|
||||
Query: query,
|
||||
Page: searchPage,
|
||||
Size: searchSize,
|
||||
Status: "empty",
|
||||
Request: &recovery.ToolCallRecord{
|
||||
ServerID: "devdoc",
|
||||
ToolName: "search_open_platform_docs",
|
||||
Arguments: cloneRecoveryArgs(requestArgs),
|
||||
},
|
||||
},
|
||||
}
|
||||
if r == nil || strings.TrimSpace(query) == "" {
|
||||
retrieval.DocSearch.Status = "skipped"
|
||||
return retrieval, nil
|
||||
}
|
||||
result, err := r.CallToolDirect(ctx, "devdoc", "search_open_platform_docs", requestArgs)
|
||||
if result != nil {
|
||||
retrieval.DocSearch.Response = toRecoveryToolResponse(result)
|
||||
}
|
||||
if err != nil {
|
||||
retrieval.DocSearch.Status = "error"
|
||||
retrieval.DocSearch.Error = err.Error()
|
||||
return retrieval, err
|
||||
}
|
||||
|
||||
retrieval.DocSearch.Items = parseDocSearchItems(result)
|
||||
if len(retrieval.DocSearch.Items) > 0 {
|
||||
retrieval.DocSearch.Status = "success"
|
||||
retrieval.KBHits = rerankDocSearchHits(query, rc, retrieval.DocSearch.Items)
|
||||
}
|
||||
return retrieval, nil
|
||||
}
|
||||
|
||||
func (r *recoveryRuntime) CallToolDirect(ctx context.Context, serverID, toolName string, args map[string]any) (*transport.ToolCallResult, error) {
|
||||
if r == nil || r.transport == nil {
|
||||
return nil, fmt.Errorf("recovery runtime not initialized")
|
||||
}
|
||||
endpoint, err := r.resolveEndpoint(ctx, serverID, toolName)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tc := r.transport.WithAuth(resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags)), resolveIdentityHeaders())
|
||||
result, err := tc.CallTool(ctx, endpoint, toolName, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if result.IsError {
|
||||
return &result, apperrors.NewAPI(
|
||||
extractMCPErrorMessage(result),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("mcp_tool_error"),
|
||||
apperrors.WithServerKey(serverID),
|
||||
)
|
||||
}
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
func (r *recoveryRuntime) resolveEndpoint(ctx context.Context, productID, toolName string) (string, error) {
|
||||
if endpoint, ok := directRuntimeEndpoint(productID, toolName); ok {
|
||||
return endpoint, nil
|
||||
}
|
||||
if r == nil || r.loader == nil {
|
||||
return "", fmt.Errorf("未找到服务 %s 的 endpoint", productID)
|
||||
}
|
||||
catalog, err := r.loader.Load(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
product, ok := catalog.FindProduct(productID)
|
||||
if !ok || strings.TrimSpace(product.Endpoint) == "" {
|
||||
return "", fmt.Errorf("未找到服务 %s 的 endpoint", productID)
|
||||
}
|
||||
return strings.TrimSpace(product.Endpoint), nil
|
||||
}
|
||||
|
||||
func recoveryRuntimeToken(flags *GlobalFlags) string {
|
||||
if flags == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(flags.Token)
|
||||
}
|
||||
|
||||
func toRecoveryToolResponse(result *transport.ToolCallResult) *recovery.ToolResponse {
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
response := &recovery.ToolResponse{IsError: result.IsError}
|
||||
if len(result.Blocks) > 0 {
|
||||
response.Content = make([]recovery.ToolResponseBlock, 0, len(result.Blocks))
|
||||
for _, block := range result.Blocks {
|
||||
response.Content = append(response.Content, recovery.ToolResponseBlock{
|
||||
Type: block.Type,
|
||||
Text: block.Text,
|
||||
})
|
||||
}
|
||||
}
|
||||
return response
|
||||
}
|
||||
|
||||
func parseDocSearchItems(result *transport.ToolCallResult) []recovery.DocSearchItem {
|
||||
if result == nil {
|
||||
return nil
|
||||
}
|
||||
if items := parseDocSearchItemsFromMap(result.Content); len(items) > 0 {
|
||||
return items
|
||||
}
|
||||
for _, block := range result.Blocks {
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(block.Text), &payload); err == nil {
|
||||
if items := parseDocSearchItemsFromMap(payload); len(items) > 0 {
|
||||
return items
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func parseDocSearchItemsFromMap(payload map[string]any) []recovery.DocSearchItem {
|
||||
if len(payload) == 0 {
|
||||
return nil
|
||||
}
|
||||
if items := toDocSearchItems(payload["items"]); len(items) > 0 {
|
||||
return items
|
||||
}
|
||||
if data, ok := payload["data"].(map[string]any); ok {
|
||||
if items := toDocSearchItems(data["items"]); len(items) > 0 {
|
||||
return items
|
||||
}
|
||||
}
|
||||
if result, ok := payload["result"].(map[string]any); ok {
|
||||
if items := toDocSearchItems(result["items"]); len(items) > 0 {
|
||||
return items
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func toDocSearchItems(raw any) []recovery.DocSearchItem {
|
||||
list, ok := raw.([]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
items := make([]recovery.DocSearchItem, 0, len(list))
|
||||
for _, entry := range list {
|
||||
object, ok := entry.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
item := recovery.DocSearchItem{}
|
||||
if title, ok := object["title"].(string); ok {
|
||||
item.Title = title
|
||||
}
|
||||
if url, ok := object["url"].(string); ok {
|
||||
item.URL = url
|
||||
}
|
||||
if desc, ok := object["desc"].(string); ok {
|
||||
item.Desc = desc
|
||||
}
|
||||
if item.Title != "" || item.URL != "" || item.Desc != "" {
|
||||
items = append(items, item)
|
||||
}
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
func rerankDocSearchHits(query string, rc recovery.RecoveryContext, items []recovery.DocSearchItem) []recovery.KBHit {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
keywords := strings.Fields(strings.ToLower(strings.TrimSpace(query)))
|
||||
type scoredHit struct {
|
||||
hit recovery.KBHit
|
||||
score float64
|
||||
}
|
||||
scored := make([]scoredHit, 0, len(items))
|
||||
for _, item := range items {
|
||||
text := strings.ToLower(strings.Join(append([]string{
|
||||
item.Title,
|
||||
item.URL,
|
||||
item.Desc,
|
||||
rc.ToolName,
|
||||
}, rc.CommandPath...), " "))
|
||||
score := 0.0
|
||||
for _, keyword := range keywords {
|
||||
if strings.Contains(text, keyword) {
|
||||
score += 1
|
||||
}
|
||||
}
|
||||
scored = append(scored, scoredHit{
|
||||
hit: recovery.KBHit{
|
||||
Source: "open_platform_docs",
|
||||
Title: item.Title,
|
||||
URL: item.URL,
|
||||
Snippet: item.Desc,
|
||||
Score: score,
|
||||
},
|
||||
score: score,
|
||||
})
|
||||
}
|
||||
sort.SliceStable(scored, func(i, j int) bool {
|
||||
return scored[i].score > scored[j].score
|
||||
})
|
||||
limit := len(scored)
|
||||
if limit > 3 {
|
||||
limit = 3
|
||||
}
|
||||
hits := make([]recovery.KBHit, 0, limit)
|
||||
for _, item := range scored[:limit] {
|
||||
hits = append(hits, item.hit)
|
||||
}
|
||||
return hits
|
||||
}
|
||||
@@ -0,0 +1,324 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
)
|
||||
|
||||
func TestRecoveryPlanReadsLastSnapshotAndPrintsJSON(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
writeRecoverySnapshot(t, configDir, recovery.LastError{
|
||||
EventID: "evt_test",
|
||||
RecordedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
Context: recovery.RecoveryContext{
|
||||
CommandPath: []string{"approval", "instance", "get"},
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationRead,
|
||||
CLIErrorCode: "RESOURCE_NOT_FOUND",
|
||||
RawError: "resource_not_found",
|
||||
Fingerprint: "fp-1",
|
||||
},
|
||||
Replay: recovery.Replay{
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationRead,
|
||||
ToolArgs: map[string]any{"instanceId": "ins_1"},
|
||||
RedactedCommand: "dws approval instance get --instance-id ins_1 --format json",
|
||||
},
|
||||
})
|
||||
|
||||
root := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
root.SetOut(&out)
|
||||
root.SetErr(&out)
|
||||
root.SetArgs([]string{"recovery", "plan", "--last", "-f", "json"})
|
||||
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("Execute(recovery plan) error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), `"event_id": "evt_test"`) {
|
||||
t.Fatalf("output missing event id:\n%s", out.String())
|
||||
}
|
||||
if !strings.Contains(out.String(), `"category": "resource"`) {
|
||||
t.Fatalf("output missing resource category:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoveryExecuteReadsLastSnapshotAndPrintsJSON(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
writeRecoverySnapshot(t, configDir, recovery.LastError{
|
||||
EventID: "evt_exec",
|
||||
RecordedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
Context: recovery.RecoveryContext{
|
||||
CommandPath: []string{"approval", "instance", "get"},
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationRead,
|
||||
CLIErrorCode: "RESOURCE_NOT_FOUND",
|
||||
RawError: "resource_not_found",
|
||||
Fingerprint: "fp-2",
|
||||
},
|
||||
Replay: recovery.Replay{
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationRead,
|
||||
ToolArgs: map[string]any{"instanceId": "ins_1"},
|
||||
RedactedCommand: "dws approval instance get --instance-id ins_1 --format json",
|
||||
},
|
||||
})
|
||||
|
||||
root := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
root.SetOut(&out)
|
||||
root.SetErr(&out)
|
||||
root.SetArgs([]string{"recovery", "execute", "--last", "-f", "json"})
|
||||
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("Execute(recovery execute) error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), `"event_id": "evt_exec"`) {
|
||||
t.Fatalf("output missing event id:\n%s", out.String())
|
||||
}
|
||||
if !strings.Contains(out.String(), `"status": "needs_agent_action"`) {
|
||||
t.Fatalf("output missing bundle status:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoveryFinalizeRequiresEventIDAndOutcome(t *testing.T) {
|
||||
root := NewRootCommand()
|
||||
root.SetOut(&bytes.Buffer{})
|
||||
root.SetErr(&bytes.Buffer{})
|
||||
root.SetArgs([]string{"recovery", "finalize"})
|
||||
|
||||
err := root.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute(recovery finalize) error = nil, want validation")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--event-id") {
|
||||
t.Fatalf("error = %v, want event-id requirement", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoveryPlanRejectsLastAndEventIDTogether(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
writeRecoverySnapshot(t, configDir, recovery.LastError{
|
||||
EventID: "evt_conflict",
|
||||
RecordedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
Context: recovery.RecoveryContext{
|
||||
CommandPath: []string{"approval", "instance", "get"},
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationRead,
|
||||
CLIErrorCode: "RESOURCE_NOT_FOUND",
|
||||
RawError: "resource_not_found",
|
||||
Fingerprint: "fp-conflict",
|
||||
},
|
||||
})
|
||||
|
||||
root := NewRootCommand()
|
||||
root.SetOut(&bytes.Buffer{})
|
||||
root.SetErr(&bytes.Buffer{})
|
||||
root.SetArgs([]string{"recovery", "plan", "--last", "--event-id", "evt_conflict"})
|
||||
|
||||
err := root.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute(recovery plan) error = nil, want conflict validation")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--last") || !strings.Contains(err.Error(), "--event-id") {
|
||||
t.Fatalf("error = %v, want mutually exclusive flags", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoveryFinalizeAcceptsLegacyExecutionFile(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
writeRecoverySnapshot(t, configDir, recovery.LastError{
|
||||
EventID: "evt_legacy_finalize",
|
||||
RecordedAt: time.Now().UTC().Format(time.RFC3339Nano),
|
||||
Context: recovery.RecoveryContext{
|
||||
CommandPath: []string{"approval", "instance", "get"},
|
||||
ServerID: "approval",
|
||||
ToolName: "get_approval_instance",
|
||||
OperationKind: recovery.OperationUnknown,
|
||||
RawError: "unexpected upstream failure",
|
||||
Fingerprint: "fp-legacy-finalize",
|
||||
},
|
||||
})
|
||||
|
||||
executionPath := filepath.Join(configDir, "legacy_execution.json")
|
||||
if err := os.WriteFile(executionPath, []byte(`{"action":"verify_resource_exists","attempts":2,"result":"failed","error":"resource still missing"}`), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile(legacy execution) error = %v", err)
|
||||
}
|
||||
|
||||
root := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
root.SetOut(&out)
|
||||
root.SetErr(&out)
|
||||
root.SetArgs([]string{
|
||||
"recovery", "finalize",
|
||||
"--event-id", "evt_legacy_finalize",
|
||||
"--outcome", "failed",
|
||||
"--execution-file", executionPath,
|
||||
"-f", "json",
|
||||
})
|
||||
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("Execute(recovery finalize) error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), `"execution_recorded": true`) {
|
||||
t.Fatalf("output missing execution_recorded flag:\n%s", out.String())
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(filepath.Join(configDir, "recovery", "recovery_events.jsonl"))
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(recovery_events.jsonl) error = %v", err)
|
||||
}
|
||||
lines := strings.Split(strings.TrimSpace(string(data)), "\n")
|
||||
lastLine := lines[len(lines)-1]
|
||||
if !strings.Contains(lastLine, `"phase":"finalized"`) {
|
||||
t.Fatalf("expected finalized event, got %s", lastLine)
|
||||
}
|
||||
if !strings.Contains(lastLine, `"legacy_execution_file"`) {
|
||||
t.Fatalf("expected legacy execution attempts to be normalized, got %s", lastLine)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteWritesRecoveryEventIDToStderrOnCapturedFailure(t *testing.T) {
|
||||
setupRuntimeCommandTest(t)
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
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
|
||||
}
|
||||
switch req["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{
|
||||
{
|
||||
"type": "text",
|
||||
"text": "baseId is required",
|
||||
},
|
||||
},
|
||||
"isError": true,
|
||||
},
|
||||
})
|
||||
}
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
t.Setenv(cli.CatalogFixtureEnv, writeDocCatalogFixture(t, server.URL, false))
|
||||
|
||||
oldArgs := os.Args
|
||||
defer func() { os.Args = oldArgs }()
|
||||
os.Args = []string{"dws", "mcp", "doc", "search_documents", "--json", `{"keyword":"design"}`, "--token", "test-token"}
|
||||
|
||||
stdoutR, stdoutW, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatalf("os.Pipe(stdout) error = %v", err)
|
||||
}
|
||||
stderrR, stderrW, err := os.Pipe()
|
||||
if err != nil {
|
||||
t.Fatalf("os.Pipe(stderr) error = %v", err)
|
||||
}
|
||||
oldStdout := os.Stdout
|
||||
oldStderr := os.Stderr
|
||||
defer func() {
|
||||
os.Stdout = oldStdout
|
||||
os.Stderr = oldStderr
|
||||
}()
|
||||
os.Stdout = stdoutW
|
||||
os.Stderr = stderrW
|
||||
|
||||
exitCode := Execute()
|
||||
|
||||
_ = stdoutW.Close()
|
||||
_ = stderrW.Close()
|
||||
stdoutData, _ := io.ReadAll(stdoutR)
|
||||
stderrData, _ := io.ReadAll(stderrR)
|
||||
|
||||
if exitCode == 0 {
|
||||
t.Fatalf("Execute() exitCode = 0, want failure\nstdout:\n%s\nstderr:\n%s", stdoutData, stderrData)
|
||||
}
|
||||
if !strings.Contains(string(stderrData), "RECOVERY_EVENT_ID=evt_") {
|
||||
t.Fatalf("stderr missing recovery event id:\n%s", stderrData)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(filepath.Join(configDir, "recovery", "last_error.json"))
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(last_error.json) error = %v", err)
|
||||
}
|
||||
var last recovery.LastError
|
||||
if err := json.Unmarshal(data, &last); err != nil {
|
||||
t.Fatalf("json.Unmarshal(last_error) error = %v", err)
|
||||
}
|
||||
if last.EventID == "" || last.Context.ToolName != "search_documents" {
|
||||
t.Fatalf("unexpected recovery snapshot %#v", last)
|
||||
}
|
||||
}
|
||||
|
||||
func writeRecoverySnapshot(t *testing.T, configDir string, last recovery.LastError) {
|
||||
t.Helper()
|
||||
|
||||
recoveryDir := filepath.Join(configDir, "recovery")
|
||||
if err := os.MkdirAll(recoveryDir, 0o700); err != nil {
|
||||
t.Fatalf("MkdirAll(recovery) error = %v", err)
|
||||
}
|
||||
data, err := json.MarshalIndent(last, "", " ")
|
||||
if err != nil {
|
||||
t.Fatalf("json.MarshalIndent() error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(recoveryDir, "last_error.json"), append(data, '\n'), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile(last_error.json) error = %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
)
|
||||
|
||||
func captureRuntimeFailure(invocation executor.Invocation, rawErr, wrappedErr error) {
|
||||
if rawErr == nil && wrappedErr == nil {
|
||||
return
|
||||
}
|
||||
store := recovery.NewStore(defaultConfigDir())
|
||||
if store == nil || !store.Enabled() {
|
||||
return
|
||||
}
|
||||
input := recovery.CaptureInput{
|
||||
CommandPath: runtimeCommandPath(invocation),
|
||||
ServerID: strings.TrimSpace(invocation.CanonicalProduct),
|
||||
ToolName: strings.TrimSpace(invocation.Tool),
|
||||
Args: cloneRecoveryArgs(invocation.Params),
|
||||
Argv: append([]string(nil), os.Args[1:]...),
|
||||
RawErr: rawErr,
|
||||
WrappedErr: wrappedErr,
|
||||
}
|
||||
_, _ = store.Capture(recovery.BuildContext(input), recovery.BuildReplay(input))
|
||||
}
|
||||
|
||||
func runtimeCommandPath(invocation executor.Invocation) []string {
|
||||
if path := currentCommandPath(); len(path) > 0 {
|
||||
return path
|
||||
}
|
||||
if legacy := strings.Fields(strings.TrimSpace(invocation.LegacyPath)); len(legacy) > 0 {
|
||||
return legacy
|
||||
}
|
||||
if product := strings.TrimSpace(invocation.CanonicalProduct); product != "" {
|
||||
if tool := strings.TrimSpace(invocation.Tool); tool != "" {
|
||||
return []string{product, tool}
|
||||
}
|
||||
return []string{product}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func currentCommandPath() []string {
|
||||
boolFlags := map[string]struct{}{
|
||||
"--verbose": {},
|
||||
"-v": {},
|
||||
"--debug": {},
|
||||
"--mock": {},
|
||||
"--dry-run": {},
|
||||
"--yes": {},
|
||||
"-y": {},
|
||||
"--help": {},
|
||||
"-h": {},
|
||||
"--json": {},
|
||||
}
|
||||
path := make([]string, 0, len(os.Args))
|
||||
skipNext := false
|
||||
for _, arg := range os.Args[1:] {
|
||||
if skipNext {
|
||||
skipNext = false
|
||||
continue
|
||||
}
|
||||
if arg == "--" {
|
||||
break
|
||||
}
|
||||
if strings.HasPrefix(arg, "-") {
|
||||
if strings.Contains(arg, "=") {
|
||||
continue
|
||||
}
|
||||
if _, ok := boolFlags[arg]; ok {
|
||||
continue
|
||||
}
|
||||
skipNext = true
|
||||
continue
|
||||
}
|
||||
path = append(path, arg)
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
func cloneRecoveryArgs(args map[string]any) map[string]any {
|
||||
if len(args) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]any, len(args))
|
||||
for key, value := range args {
|
||||
out[key] = value
|
||||
}
|
||||
return out
|
||||
}
|
||||
+895
-33
File diff suppressed because it is too large
Load Diff
@@ -29,7 +29,15 @@ import (
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
func TestPrintExecutionErrorDefaultsToHumanReadable(t *testing.T) {
|
||||
// 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()
|
||||
|
||||
root := NewRootCommand()
|
||||
@@ -44,13 +52,10 @@ func TestPrintExecutionErrorDefaultsToHumanReadable(t *testing.T) {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty for human-readable error output", stdout.String())
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stderr.String(), "Error: [VALIDATION] bad flag") {
|
||||
t.Fatalf("stderr = %q, want human-readable header", stderr.String())
|
||||
}
|
||||
if !strings.Contains(stderr.String(), "Hint: Pass the required flag and retry.") {
|
||||
t.Fatalf("stderr = %q, want hint line", stderr.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -68,11 +73,11 @@ func TestPrintExecutionErrorUsesJSONWhenFormatIsJSON(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -100,11 +105,11 @@ func TestPrintExecutionErrorUsesJSONWhenCommandSetsJSONFlag(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.String())
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty when errors go to stderr", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
if !strings.Contains(stderr.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stderr = %q, want JSON error payload", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -175,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(), "版本:") {
|
||||
t.Fatalf("version output missing version header:\n%s", out.String())
|
||||
if !strings.Contains(out.String(), "Version:") {
|
||||
t.Fatalf("version output missing Version line:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -216,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(), "版本:") {
|
||||
t.Fatalf("version output missing version header:\n%s", out.String())
|
||||
if !strings.Contains(out.String(), "Version:") {
|
||||
t.Fatalf("version output missing Version line:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -258,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) {
|
||||
@@ -345,3 +355,86 @@ func TestNestedShortHelpDoesNotRequirePINOrLogin(t *testing.T) {
|
||||
t.Fatalf("nested short help output missing command title:\n%s", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_writes_raw_JSON_to_stderr(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_LOW_RISK_NO_PERMISSION","data":{}}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty for RawStderrError", stdout.String())
|
||||
}
|
||||
got := strings.TrimSpace(stderr.String())
|
||||
if got != rawJSON {
|
||||
t.Fatalf("stderr = %q, want raw JSON %q", got, rawJSON)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_exit_code_is_4(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := &patLikeError{raw: `{"code":"PAT_MEDIUM_RISK_NO_PERMISSION"}`}
|
||||
exitCode := apperrors.ExitCode(err)
|
||||
if exitCode != 4 {
|
||||
t.Fatalf("apperrors.ExitCode(patLikeError) = %d, want 4", exitCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrintExecutionError_RawStderrError_takes_precedence_over_JSON_mode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
rawJSON := `{"success":false,"code":"PAT_HIGH_RISK_NO_PERMISSION"}`
|
||||
err := &patLikeError{raw: rawJSON}
|
||||
|
||||
root := NewRootCommand()
|
||||
_ = root.PersistentFlags().Set("format", "json")
|
||||
|
||||
var stdout, stderr bytes.Buffer
|
||||
writeErr := printExecutionError(root, &stdout, &stderr, err)
|
||||
if writeErr != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", writeErr)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty — RawStderrError should bypass JSON mode", stdout.String())
|
||||
}
|
||||
if !strings.Contains(stderr.String(), "PAT_HIGH_RISK_NO_PERMISSION") {
|
||||
t.Fatalf("stderr = %q, want raw PAT JSON", stderr.String())
|
||||
}
|
||||
}
|
||||
|
||||
// simulateExecuteWithPanic mirrors the recovery pattern in Execute():
|
||||
// named return + defer recover → exitCode = 5 on panic.
|
||||
func simulateExecuteWithPanic(doPanic bool) (exitCode int) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
exitCode = 5
|
||||
}
|
||||
}()
|
||||
if doPanic {
|
||||
panic("test panic")
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func TestExecute_panic_recovery_returns_exit_5(t *testing.T) {
|
||||
t.Parallel()
|
||||
code := simulateExecuteWithPanic(true)
|
||||
if code != 5 {
|
||||
t.Fatalf("panic recovery exitCode = %d, want 5", code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecute_no_panic_returns_0(t *testing.T) {
|
||||
t.Parallel()
|
||||
code := simulateExecuteWithPanic(false)
|
||||
if code != 0 {
|
||||
t.Fatalf("no-panic exitCode = %d, want 0", code)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -13,6 +15,26 @@ func configureRootHelp(root *cobra.Command) {
|
||||
return
|
||||
}
|
||||
|
||||
// Replace the cobra-default English help command with a localized one so
|
||||
// that both its listing short (shown in `dws --help`) and its own
|
||||
// `dws help --help` long text follow the active locale.
|
||||
root.SetHelpCommand(&cobra.Command{
|
||||
Use: "help [command]",
|
||||
Short: i18n.T("查看任意命令的帮助信息"),
|
||||
Long: i18n.T("显示任意命令的帮助文案。\n" +
|
||||
"用法:dws help [命令路径] 查看完整说明。"),
|
||||
DisableAutoGenTag: true,
|
||||
Run: func(c *cobra.Command, args []string) {
|
||||
target, _, err := c.Root().Find(args)
|
||||
if target == nil || err != nil {
|
||||
c.Root().HelpFunc()(c.Root(), args)
|
||||
return
|
||||
}
|
||||
target.InitDefaultHelpFlag()
|
||||
_ = target.Help()
|
||||
},
|
||||
})
|
||||
|
||||
defaultHelpFunc := root.HelpFunc()
|
||||
root.SetHelpFunc(func(cmd *cobra.Command, args []string) {
|
||||
if cmd != root {
|
||||
@@ -25,6 +47,7 @@ func configureRootHelp(root *cobra.Command) {
|
||||
|
||||
func renderRootHelp(root *cobra.Command) {
|
||||
services := visibleMCPRootCommands(root)
|
||||
utilities := visibleUtilityRootCommands(root)
|
||||
w := root.OutOrStdout()
|
||||
|
||||
if len(services) == 0 {
|
||||
@@ -44,8 +67,21 @@ func renderRootHelp(root *cobra.Command) {
|
||||
|
||||
_, _ = fmt.Fprintln(w, "Usage:")
|
||||
_, _ = fmt.Fprintln(w, " dws <service> [command] [flags]")
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, " dws <command> [flags]")
|
||||
}
|
||||
_, _ = fmt.Fprintln(w)
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service.`)
|
||||
if len(utilities) > 0 {
|
||||
_, _ = fmt.Fprintln(w, "Utility Commands:")
|
||||
_, _ = fmt.Fprintln(w)
|
||||
tw := tabwriter.NewWriter(w, 0, 0, 2, ' ', 0)
|
||||
for _, utility := range utilities {
|
||||
_, _ = fmt.Fprintf(tw, " %s\t%s\n", utility.Name(), strings.TrimSpace(utility.Short))
|
||||
}
|
||||
_ = tw.Flush()
|
||||
_, _ = fmt.Fprintln(w)
|
||||
}
|
||||
_, _ = fmt.Fprintln(w, `Use "dws <service> --help" for more information about a discovered MCP service or "dws <command> --help" for utility commands.`)
|
||||
}
|
||||
|
||||
func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
@@ -53,7 +89,16 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
return nil
|
||||
}
|
||||
|
||||
allowed := DirectRuntimeProductIDs()
|
||||
var allowed map[string]bool
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
products := fn()
|
||||
allowed = make(map[string]bool, len(products))
|
||||
for _, p := range products {
|
||||
allowed[p] = true
|
||||
}
|
||||
} else {
|
||||
allowed = DirectRuntimeProductIDs()
|
||||
}
|
||||
if len(allowed) == 0 {
|
||||
return nil
|
||||
}
|
||||
@@ -70,3 +115,29 @@ func visibleMCPRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
func visibleUtilityRootCommands(root *cobra.Command) []*cobra.Command {
|
||||
if root == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
productCommands := DirectRuntimeProductIDs()
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
productCommands = make(map[string]bool, len(fn()))
|
||||
for _, product := range fn() {
|
||||
productCommands[product] = true
|
||||
}
|
||||
}
|
||||
|
||||
commands := make([]*cobra.Command, 0)
|
||||
for _, cmd := range root.Commands() {
|
||||
if cmd == nil || cmd.Hidden {
|
||||
continue
|
||||
}
|
||||
if productCommands[cmd.Name()] {
|
||||
continue
|
||||
}
|
||||
commands = append(commands, cmd)
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
+407
-27
@@ -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"
|
||||
@@ -42,15 +88,28 @@ const (
|
||||
envDingtalkTraceID = "DINGTALK_TRACE_ID"
|
||||
envDingtalkSessionID = "DINGTALK_SESSION_ID"
|
||||
envDingtalkMessageID = "DINGTALK_MESSAGE_ID"
|
||||
|
||||
// Environment variables for third-party channel integration
|
||||
envDWSChannel = "DWS_CHANNEL"
|
||||
)
|
||||
|
||||
func newCommandRunnerWithFlags(loader cli.CatalogLoader, flags *GlobalFlags) executor.Runner {
|
||||
// Ensure DWS_CLIENT_ID env is populated from persisted config before
|
||||
// resolveIdentityHeaders reads it. This covers fresh-process cold starts
|
||||
// where no env var has been inherited from a parent process.
|
||||
if os.Getenv("DWS_CLIENT_ID") == "" {
|
||||
if cid := authpkg.ClientID(); cid != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", cid)
|
||||
}
|
||||
}
|
||||
|
||||
var httpClient *http.Client
|
||||
if flags != nil && flags.Timeout > 0 {
|
||||
httpClient = &http.Client{Timeout: time.Duration(flags.Timeout) * time.Second}
|
||||
}
|
||||
transportClient := transport.NewClient(httpClient)
|
||||
transportClient.ExtraHeaders = resolveIdentityHeaders()
|
||||
transportClient.FileLogger = FileLoggerInstance()
|
||||
return &runtimeRunner{
|
||||
loader: loader,
|
||||
transport: transportClient,
|
||||
@@ -86,15 +145,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); ok {
|
||||
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)
|
||||
@@ -115,8 +184,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{
|
||||
@@ -147,19 +267,105 @@ 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
|
||||
}
|
||||
}
|
||||
}
|
||||
// PAT scope error: offer human-readable output and retry after authorization
|
||||
if isPatScopeError(err) {
|
||||
scopeErr := extractPatScopeError(err)
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return executor.Result{}, err
|
||||
}
|
||||
|
||||
// ---- Edition hook gets first dibs (preserves overlay PATError passthrough) ----
|
||||
if fn := edition.Get().ClassifyToolResult; fn != nil {
|
||||
if editionErr := fn(callResult.Content); editionErr != nil {
|
||||
if patCheck := apperrors.AsPatAuthCheckError(editionErr); patCheck != nil {
|
||||
if IsPatRetrying(ctx) {
|
||||
return executor.Result{}, patCheck // already retried once, don't loop
|
||||
}
|
||||
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
return executor.Result{}, editionErr
|
||||
}
|
||||
}
|
||||
|
||||
// ---- Structured PAT auth check (open-source fallback) ----
|
||||
if patCheck := apperrors.ClassifyPatAuthCheck(callResult.Content); patCheck != nil {
|
||||
if IsPatRetrying(ctx) {
|
||||
return executor.Result{}, patCheck // already retried once, don't loop
|
||||
}
|
||||
return handlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
|
||||
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
|
||||
|
||||
// ClassifyToolResult hook: let the overlay intercept known error
|
||||
// patterns (PAT permission, gateway-auth) before generic handling.
|
||||
if classify := edition.Get().ClassifyToolResult; classify != nil {
|
||||
if hookErr := classify(callResult.Content); hookErr != nil {
|
||||
captureRuntimeFailure(invocation, hookErr, hookErr)
|
||||
return executor.Result{}, hookErr
|
||||
}
|
||||
}
|
||||
|
||||
mcpErr := apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
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),
|
||||
)
|
||||
// PAT scope error in business response: offer human-readable output and retry
|
||||
if isPatScopeError(mcpErr) {
|
||||
scopeErr := extractPatScopeError(mcpErr)
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return retryWithPatAuthRetry(ctx, r, invocation, scopeErr, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return executor.Result{}, mcpErr
|
||||
}
|
||||
|
||||
@@ -168,6 +374,18 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
return executor.Result{}, err
|
||||
}
|
||||
|
||||
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),
|
||||
)
|
||||
}
|
||||
|
||||
invocation.Implemented = true
|
||||
response := map[string]any{
|
||||
"endpoint": transport.RedactURL(endpoint),
|
||||
@@ -179,37 +397,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 {
|
||||
@@ -243,6 +552,14 @@ func runtimeFlagEnabled(raw string, defaultValue bool) bool {
|
||||
}
|
||||
}
|
||||
|
||||
func isAuthError(err error) bool {
|
||||
var appErr *apperrors.Error
|
||||
if errors.As(err, &appErr) {
|
||||
return appErr.Category == apperrors.CategoryAuth
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func productEndpointOverride(productID string) (string, bool) {
|
||||
key := "DINGTALK_" + strings.ToUpper(strings.ReplaceAll(strings.TrimSpace(productID), "-", "_")) + "_MCP_URL"
|
||||
value := strings.TrimSpace(os.Getenv(key))
|
||||
@@ -261,7 +578,7 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers = make(map[string]string)
|
||||
}
|
||||
|
||||
// Inject environment variable based headers for MCP gateway tracking
|
||||
// Inject environment variable based headers for MCP gateway tracking.
|
||||
envHeaders := map[string]string{
|
||||
"x-dingtalk-agent": os.Getenv(envDingtalkAgent),
|
||||
"x-dingtalk-trace-id": os.Getenv(envDingtalkTraceID),
|
||||
@@ -273,9 +590,39 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers[k] = v
|
||||
}
|
||||
}
|
||||
|
||||
// Inject third-party channel headers
|
||||
if v := os.Getenv(envDWSChannel); v != "" {
|
||||
headers["x-dws-channel"] = v
|
||||
}
|
||||
|
||||
if fn := edition.Get().MergeHeaders; fn != nil {
|
||||
headers = fn(headers)
|
||||
}
|
||||
return headers
|
||||
}
|
||||
|
||||
// detectBusinessError checks the MCP response content for DingTalk business
|
||||
// errors (success=false + errorCode/errorMsg) that are not flagged at the MCP
|
||||
// protocol level. Returns the error message, or "" if the response is OK.
|
||||
func detectBusinessError(content map[string]any) string {
|
||||
success, ok := content["success"]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
b, ok := success.(bool)
|
||||
if !ok || b {
|
||||
return ""
|
||||
}
|
||||
if msg, ok := content["errorMsg"].(string); ok && strings.TrimSpace(msg) != "" {
|
||||
return strings.TrimSpace(msg)
|
||||
}
|
||||
if code, ok := content["errorCode"].(string); ok && strings.TrimSpace(code) != "" {
|
||||
return "business error: code " + strings.TrimSpace(code)
|
||||
}
|
||||
return "business error: success=false"
|
||||
}
|
||||
|
||||
// extractMCPErrorMessage builds an error message from a ToolCallResult with
|
||||
// isError=true. It extracts text from content blocks when available.
|
||||
func extractMCPErrorMessage(result transport.ToolCallResult) string {
|
||||
@@ -295,3 +642,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,632 @@
|
||||
// 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 search`.
|
||||
type CliSkillDTO struct {
|
||||
SkillID string `json:"skillId"`
|
||||
Name string `json:"name"`
|
||||
Desc string `json:"desc"`
|
||||
Icon string `json:"icon"`
|
||||
}
|
||||
|
||||
// agentSkillPaths maps target names to their relative skill installation paths.
|
||||
// These paths are relative to the user's home directory.
|
||||
var agentSkillPaths = map[string]string{
|
||||
"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(
|
||||
newSkillInstallCommand(),
|
||||
newSkillGetCommand(),
|
||||
newSkillSearchCommand(),
|
||||
newSkillFindHintCommand(),
|
||||
newSkillAddHintCommand(),
|
||||
)
|
||||
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 newSkillSearchCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "search",
|
||||
Short: "从钉钉技能市场搜索技能",
|
||||
Long: "从钉钉技能市场搜索技能,根据关键词返回匹配的技能列表。",
|
||||
Example: " dws skill search --query 关键词",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillFind,
|
||||
}
|
||||
cmd.Flags().String("query", "", "搜索关键词(必填)")
|
||||
_ = cmd.MarkFlagRequired("query")
|
||||
cmd.Flags().String("scopes", "", "查询范围,空格分隔。备选值:DingtalkMarket(钉钉市场)、OrgInternal(企业内部)。为空默认查市场技能")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillFindHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "find",
|
||||
Short: "兼容旧用法,提示使用 skill search",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill search --query <关键词>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func newSkillInstallCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "install <skillId> <target>",
|
||||
Short: "下载并安装技能到指定目录",
|
||||
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
|
||||
|
||||
参数:
|
||||
skillId 技能 ID(必填),可从钉钉技能市场获取
|
||||
target 安装目标(必填),支持: %s
|
||||
|
||||
安装路径:
|
||||
qoder -> ~/.qoder/skills/
|
||||
claude -> ~/.claude/skills/
|
||||
cursor -> ~/.cursor/skills/
|
||||
codex -> ~/.codex/skills/
|
||||
opencode -> ~/.config/opencode/skills/
|
||||
. -> 当前目录
|
||||
|
||||
示例:
|
||||
dws skill install skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill install skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill install skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillAdd,
|
||||
}
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddHintCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "add",
|
||||
Short: "兼容旧用法,提示使用 skill install",
|
||||
Hidden: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "use: dws skill install <skillId> <target>")
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/install?skillId=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(skillID)))
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "⬇️ 下载技能包...")
|
||||
|
||||
tmpDir, err := downloadSkillToTmpDir(cmd.Context(), apiURL, accessToken)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), tmpDir)
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
keyword, _ := cmd.Flags().GetString("query")
|
||||
scopes, _ := cmd.Flags().GetString("scopes")
|
||||
accessToken, err := loadSkillAccessToken()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
apiURL := fmt.Sprintf("%s/cli/find-skills?keyword=%s", skillAPIHost(), url.QueryEscape(strings.TrimSpace(keyword)))
|
||||
if scopes != "" {
|
||||
apiURL += "&scopes=" + url.QueryEscape(scopes)
|
||||
}
|
||||
req, err := http.NewRequestWithContext(cmd.Context(), http.MethodGet, apiURL, nil)
|
||||
if err != nil {
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to create request: %v", err))
|
||||
}
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
client := &http.Client{Timeout: 30 * time.Second}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %v", err), apperrors.WithRetryable(true))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return parseLegacySkillAPIError(resp)
|
||||
}
|
||||
|
||||
var result findSkillsResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to parse search response: %v", err))
|
||||
}
|
||||
if !result.Success {
|
||||
errMsg := strings.TrimSpace(result.ErrorMsg)
|
||||
if errMsg == "" {
|
||||
errMsg = strings.TrimSpace(result.ErrorCode)
|
||||
}
|
||||
if errMsg == "" {
|
||||
errMsg = "unknown error"
|
||||
}
|
||||
return apperrors.NewAPI(fmt.Sprintf("failed to search skills: %s", errMsg))
|
||||
}
|
||||
|
||||
if len(result.Result) == 0 {
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "未找到匹配的技能")
|
||||
return nil
|
||||
}
|
||||
|
||||
for _, skill := range result.Result {
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "SkillID: %s\n", skill.SkillID)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Name: %s\n", skill.Name)
|
||||
_, _ = fmt.Fprintf(cmd.OutOrStdout(), "Desc: %s\n", skill.Desc)
|
||||
_, _ = fmt.Fprintln(cmd.OutOrStdout(), "---")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
skillID := strings.TrimSpace(args[0])
|
||||
target := strings.TrimSpace(args[1])
|
||||
|
||||
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 TestSkillInstallCommandValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
wantErr bool
|
||||
errMsg string
|
||||
}{
|
||||
{
|
||||
name: "missing arguments",
|
||||
args: []string{"skill", "install"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "missing target",
|
||||
args: []string{"skill", "install", "skill-123"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "too many arguments",
|
||||
args: []string{"skill", "install", "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 TestSkillInstallInvalidTarget(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", "install", "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 TestSkillInstallRequiresAuth(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", "install", "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{"install", "search", "get"} {
|
||||
if !strings.Contains(output, subcmd) {
|
||||
t.Errorf("help should mention %q subcommand, got: %s", subcmd, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillInstallCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "install", "--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 TestSkillSearchCommandValidation(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "search"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() error = nil, want missing required flag error")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "required flag") {
|
||||
t.Fatalf("error = %v, want required flag message", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillFindHintCommand(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "find"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if !strings.Contains(out.String(), "dws skill search --query") {
|
||||
t.Fatalf("output = %q, want legacy hint", out.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestDownloadSkillFileSuccess(t *testing.T) {
|
||||
// Create a mock server that returns a zip file
|
||||
expectedContent := []byte("fake zip content")
|
||||
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,119 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
const stdioEndpointScheme = "stdio://"
|
||||
|
||||
var (
|
||||
stdioMu sync.RWMutex
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
)
|
||||
|
||||
// RegisterStdioClient stores a StdioClient keyed by its canonical product ID
|
||||
// (the CLI.ID used in the server descriptor). The runner looks up this client
|
||||
// when a stdio:// endpoint is resolved at execution time.
|
||||
func RegisterStdioClient(productID string, client *transport.StdioClient) {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
stdioClients[productID] = client
|
||||
}
|
||||
|
||||
// LookupStdioClient returns the StdioClient registered for the given product ID.
|
||||
// The productID can be either the full key (pluginName/serverKey) or just the serverKey.
|
||||
// This supports backward compatibility with existing CanonicalProduct values.
|
||||
func LookupStdioClient(productID string) (*transport.StdioClient, bool) {
|
||||
stdioMu.RLock()
|
||||
defer stdioMu.RUnlock()
|
||||
// Try exact match first
|
||||
if c, ok := stdioClients[productID]; ok {
|
||||
return c, true
|
||||
}
|
||||
// If not found, try matching by serverKey suffix (for backward compatibility)
|
||||
for id, c := range stdioClients {
|
||||
if idx := strings.LastIndex(id, "/"); idx >= 0 {
|
||||
if id[idx+1:] == productID {
|
||||
return c, true
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
// StdioEndpoint returns a virtual endpoint URL for a stdio-based MCP server.
|
||||
// Format: stdio://{pluginName}/{serverKey}
|
||||
func StdioEndpoint(pluginName, serverKey string) string {
|
||||
return stdioEndpointScheme + pluginName + "/" + serverKey
|
||||
}
|
||||
|
||||
// IsStdioEndpoint returns true if the endpoint uses the stdio:// scheme.
|
||||
func IsStdioEndpoint(endpoint string) bool {
|
||||
return strings.HasPrefix(endpoint, stdioEndpointScheme)
|
||||
}
|
||||
|
||||
// StopAllStdioClients stops all registered stdio clients.
|
||||
// This should be called on program exit to terminate child processes.
|
||||
func StopAllStdioClients() {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
for id, client := range stdioClients {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
}
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
}
|
||||
|
||||
// StopStdioClient stops a specific stdio client by product ID.
|
||||
// Returns true if the client was found and stopped, false otherwise.
|
||||
func StopStdioClient(productID string) bool {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
client, ok := stdioClients[productID]
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", productID, "error", err)
|
||||
}
|
||||
delete(stdioClients, productID)
|
||||
return true
|
||||
}
|
||||
|
||||
// StopStdioClientsByPlugin stops all stdio clients belonging to a plugin.
|
||||
// The productID format is "pluginName/serverKey". This function stops all
|
||||
// clients whose productID has the given pluginName prefix.
|
||||
func StopStdioClientsByPlugin(pluginName string) int {
|
||||
stdioMu.Lock()
|
||||
defer stdioMu.Unlock()
|
||||
prefix := pluginName + "/"
|
||||
count := 0
|
||||
for id, client := range stdioClients {
|
||||
if len(id) > len(prefix) && id[:len(prefix)] == prefix {
|
||||
if err := client.Stop(); err != nil {
|
||||
slog.Warn("failed to stop stdio client", "id", id, "error", err)
|
||||
}
|
||||
delete(stdioClients, id)
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestStdioEndpoint(t *testing.T) {
|
||||
endpoint := StdioEndpoint("hello-plugin", "hello")
|
||||
want := "stdio://hello-plugin/hello"
|
||||
if endpoint != want {
|
||||
t.Errorf("StdioEndpoint() = %q, want %q", endpoint, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsStdioEndpoint(t *testing.T) {
|
||||
tests := []struct {
|
||||
endpoint string
|
||||
want bool
|
||||
}{
|
||||
{"stdio://hello-plugin/hello", true},
|
||||
{"stdio://conference/local", true},
|
||||
{"https://mcp.dingtalk.com", false},
|
||||
{"", false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := IsStdioEndpoint(tt.endpoint); got != tt.want {
|
||||
t.Errorf("IsStdioEndpoint(%q) = %v, want %v", tt.endpoint, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioClientRegistry(t *testing.T) {
|
||||
// Clean up after test
|
||||
defer func() {
|
||||
stdioMu.Lock()
|
||||
delete(stdioClients, "test-product")
|
||||
stdioMu.Unlock()
|
||||
}()
|
||||
|
||||
// Initially not found
|
||||
if _, ok := LookupStdioClient("test-product"); ok {
|
||||
t.Error("expected LookupStdioClient to return false for unregistered product")
|
||||
}
|
||||
|
||||
// Register a client
|
||||
client := transport.NewStdioClient("echo", nil, nil)
|
||||
RegisterStdioClient("test-product", client)
|
||||
|
||||
// Now should be found
|
||||
got, ok := LookupStdioClient("test-product")
|
||||
if !ok {
|
||||
t.Fatal("expected LookupStdioClient to return true after registration")
|
||||
}
|
||||
if got != client {
|
||||
t.Error("LookupStdioClient returned different client instance")
|
||||
}
|
||||
}
|
||||
@@ -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.0"
|
||||
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
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
@@ -27,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)
|
||||
@@ -35,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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,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)
|
||||
@@ -60,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)
|
||||
@@ -296,73 +306,6 @@ func TestRevokeTokenRemote(t *testing.T) {
|
||||
// Can't easily test since LogoutURL is a const. Just test that it doesn't panic with real URL.
|
||||
}
|
||||
|
||||
// ─── export.go ─────────────────────────────────────────────────────────
|
||||
|
||||
func TestLoadExportedCredentials_ValidPersistentCode(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
creds := ExportedCredentials{
|
||||
PersistentCode: "pcode-123",
|
||||
CorpID: "corp1",
|
||||
ExportedAt: time.Now().Format(time.RFC3339),
|
||||
}
|
||||
data, _ := json.Marshal(creds)
|
||||
path := filepath.Join(dir, "creds.json")
|
||||
_ = os.WriteFile(path, data, 0o600)
|
||||
|
||||
_, err := LoadExportedCredentials(context.Background(), path, dir)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExportedCredentials_ValidRefreshToken(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
creds := ExportedCredentials{
|
||||
RefreshToken: "refresh-123",
|
||||
CorpID: "corp1",
|
||||
ExportedAt: time.Now().Format(time.RFC3339),
|
||||
}
|
||||
data, _ := json.Marshal(creds)
|
||||
path := filepath.Join(dir, "creds.json")
|
||||
_ = os.WriteFile(path, data, 0o600)
|
||||
|
||||
_, err := LoadExportedCredentials(context.Background(), path, dir)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExportedCredentials_NoCredential(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
creds := ExportedCredentials{CorpID: "corp1"}
|
||||
data, _ := json.Marshal(creds)
|
||||
path := filepath.Join(dir, "creds.json")
|
||||
os.WriteFile(path, data, 0o600)
|
||||
|
||||
_, err := LoadExportedCredentials(context.Background(), path, dir)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing credentials")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExportedCredentials_InvalidJSON(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "creds.json")
|
||||
os.WriteFile(path, []byte("not json"), 0o600)
|
||||
|
||||
_, err := LoadExportedCredentials(context.Background(), path, dir)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for invalid JSON")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadExportedCredentials_MissingFile(t *testing.T) {
|
||||
_, err := LoadExportedCredentials(context.Background(), "/nonexistent/path", t.TempDir())
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file")
|
||||
}
|
||||
}
|
||||
|
||||
// ─── oauth_helpers.go ──────────────────────────────────────────────────
|
||||
|
||||
type tokenResponse struct {
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
// 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 "testing"
|
||||
|
||||
func TestClassifyDenialReason(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
status *CLIAuthStatus
|
||||
currentChannel string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "error CHANNEL_REQUIRED",
|
||||
status: &CLIAuthStatus{
|
||||
ErrorCode: "CHANNEL_REQUIRED",
|
||||
},
|
||||
want: "channel_required",
|
||||
},
|
||||
{
|
||||
name: "error NO_AUTH",
|
||||
status: &CLIAuthStatus{
|
||||
ErrorCode: "NO_AUTH",
|
||||
},
|
||||
want: "no_auth",
|
||||
},
|
||||
{
|
||||
name: "success false or nil result → unknown",
|
||||
status: &CLIAuthStatus{
|
||||
Success: false,
|
||||
},
|
||||
want: "unknown",
|
||||
},
|
||||
{
|
||||
name: "cliAuthEnabled true → no denial",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "userScope forbidden wins over channel",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "forbidden",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "channel-b",
|
||||
want: "user_forbidden",
|
||||
},
|
||||
{
|
||||
// Real-world case reported: user is in allowedUsers but the current
|
||||
// DWS_CHANNEL is not in allowedChannels. Reason must be channel,
|
||||
// NOT user.
|
||||
name: "user allowed but channel not in allowedChannels → channel_not_allowed",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
AllowedUsers: []string{"014566033934857460"},
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"2a4a658e467998befb7fa333c19ba2b3a3bacfa4"},
|
||||
},
|
||||
},
|
||||
currentChannel: "different-channel",
|
||||
want: "channel_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "channelScope specified but current channel empty → channel_required",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "",
|
||||
want: "channel_required",
|
||||
},
|
||||
{
|
||||
name: "channel matches allowedChannels → fall back to user denial",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
ChannelScope: "specified",
|
||||
AllowedChannels: []string{"channel-a"},
|
||||
},
|
||||
},
|
||||
currentChannel: "channel-a",
|
||||
want: "user_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "only userScope=specified, no channel restriction → user_not_allowed",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{
|
||||
CLIAuthEnabled: false,
|
||||
UserScope: "specified",
|
||||
},
|
||||
},
|
||||
currentChannel: "",
|
||||
want: "user_not_allowed",
|
||||
},
|
||||
{
|
||||
name: "no user or channel restriction → cli_not_enabled",
|
||||
status: &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
},
|
||||
want: "cli_not_enabled",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
got := classifyDenialReason(tc.status, tc.currentChannel)
|
||||
if got != tc.want {
|
||||
t.Fatalf("classifyDenialReason() = %q, want %q", got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,847 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// setupMCPConfigDir creates a temp config dir with mcp_url pointing to the
|
||||
// given test server and sets DWS_CONFIG_DIR via t.Setenv.
|
||||
// NOTE: tests calling this must NOT use t.Parallel().
|
||||
func setupMCPConfigDir(t *testing.T, srvURL string) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
os.WriteFile(filepath.Join(dir, "mcp_url"), []byte(srvURL), 0o600)
|
||||
t.Setenv("DWS_CONFIG_DIR", dir)
|
||||
return dir
|
||||
}
|
||||
|
||||
// resetClientIDFromMCP clears the MCP-sourced flag (test helper).
|
||||
func resetClientIDFromMCP() {
|
||||
clientMu.Lock()
|
||||
defer clientMu.Unlock()
|
||||
clientIDFromMCP = false
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 1. CheckCLIAuthEnabled: interface error → fail-closed with retry
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error from CheckCLIAuthEnabled when server returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Server 500 → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
configDir := setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 2 * time.Second},
|
||||
}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ Connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `{this is not valid json}`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Timeout_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(5 * time.Second)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 200 * time.Millisecond},
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(ctx, "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error on request timeout, got nil")
|
||||
}
|
||||
t.Logf("✅ Timeout → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 2. CheckCLIAuthEnabled: transient error then recovery → succeeds
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !status.Success || status.Result == nil || !status.Result.CLIAuthEnabled {
|
||||
t.Fatalf("expected CLIAuthEnabled=true, got %+v", status)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ Transient error then success: attempts=%d, enabled=%v", calls.Load(), status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 3. CheckCLIAuthEnabled: normal responses (pass-through)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status.Result == nil || !status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=true")
|
||||
}
|
||||
t.Logf("✅ Normal enabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status.Result == nil || status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=false")
|
||||
}
|
||||
t.Logf("✅ Normal disabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 4. OAuth callback: simulates the fail-closed logic at the /callback level
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestOAuthCallback_CLIAuthError_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error = fmt.Errorf("simulated network error")
|
||||
var authStatus *CLIAuthStatus
|
||||
_ = authStatus
|
||||
|
||||
// This is the exact expression used in oauth_provider.go callback:
|
||||
// cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
cliAuthEnabled := statusErr == nil // false
|
||||
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when statusErr != nil")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=%v → cliAuthEnabled=%v → shows notEnabledHTML (fail-closed)", statusErr, cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthEnabled_ShowsSuccessPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && authStatus.Result.CLIAuthEnabled
|
||||
if !cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be true when API returns enabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=true → cliAuthEnabled=%v → shows successHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthDisabledByServer_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result != nil && authStatus.Result.CLIAuthEnabled
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when server says disabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=false → cliAuthEnabled=%v → shows notEnabledHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 5. Device Flow: loginOnce with broken cliAuthEnabled endpoint
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
|
||||
default:
|
||||
http.Error(w, "not found: "+r.URL.Path, http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth check fails, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "检查 CLI 授权状态失败") && !strings.Contains(err.Error(), "Failed to check CLI auth status") {
|
||||
t.Fatalf("unexpected error message: %s", err)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth check error → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: false},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, SuperAdminPath):
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "张三"}},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth is disabled, got nil")
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth disabled by server → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DevicePollPath):
|
||||
// New terminal API: return APPROVED status
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "test-auth-code",
|
||||
},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: &CLIAuthResult{CLIAuthEnabled: true},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
terminalBaseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
token, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected loginOnce to succeed, got error: %v", err)
|
||||
}
|
||||
if token.AccessToken != "test-access-token" {
|
||||
t.Fatalf("unexpected token: %s", token.AccessToken)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth enabled → login succeeded, token=%s", token.AccessToken)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 6. FetchClientIDFromMCP: /cli/clientId error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestFetchClientIDFromMCP_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/clientId returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId 500 → fail-closed with retry: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `not json at all`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_BusinessError_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: false,
|
||||
ErrorCode: "forbidden",
|
||||
ErrorMsg: "access denied",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when server returns success=false, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "access denied") {
|
||||
t.Fatalf("error should contain server error message, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId business error → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "recovered-client-id",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if id != "recovered-client-id" {
|
||||
t.Fatalf("expected client ID 'recovered-client-id', got %q", id)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId transient then success: attempts=%d, id=%s", calls.Load(), id)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != ClientIDPath {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "my-client-id-123",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if id != "my-client-id-123" {
|
||||
t.Fatalf("expected 'my-client-id-123', got %q", id)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId normal success: id=%s", id)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 7. GetSuperAdmins: /cli/superAdmin error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetSuperAdmins_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/superAdmin returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "a1", Name: "张三"}, {StaffID: "a2", Name: "李四"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !result.Success || len(result.Result) != 2 {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin transient then success: attempts=%d, admins=%v", calls.Load(), result.Result)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "王五"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(result.Result) != 1 || result.Result[0].Name != "王五" {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin normal success: admins=%v", result.Result)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 8. SendCliAuthApply: /cli/sendCliAuthApply error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendCliAuthApply_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/sendCliAuthApply returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 1 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failure, got error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 2 {
|
||||
t.Fatalf("expected 2 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply transient then success: attempts=%d", calls.Load())
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
if !strings.Contains(r.URL.RawQuery, "adminStaffId=admin123") {
|
||||
t.Errorf("missing or wrong adminStaffId param: %s", r.URL.RawQuery)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "good-token", "admin123")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply normal success: result=%+v", result)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_BusinessError(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{
|
||||
Success: false,
|
||||
ErrorCode: "invalid_admin",
|
||||
ErrorMsg: "admin not found",
|
||||
Result: false,
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "nonexistent")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected transport error: %v", err)
|
||||
}
|
||||
if result.Success {
|
||||
t.Fatal("expected success=false for business error")
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply business error: errorCode=%s, errorMsg=%s", result.ErrorCode, result.ErrorMsg)
|
||||
}
|
||||
+253
-17
@@ -26,39 +26,44 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/fatih/color"
|
||||
)
|
||||
|
||||
const (
|
||||
// defaultPollInterval is the default seconds between device token polls.
|
||||
defaultPollInterval = 5
|
||||
// The server-side Redis TTL is 10 minutes; a 2-second interval keeps the
|
||||
// user-perceived latency low while staying well within rate limits.
|
||||
defaultPollInterval = 2
|
||||
// maxPollInterval caps the polling interval to prevent DoS via slow_down.
|
||||
maxPollInterval = 30
|
||||
// maxPollTotalWait caps the total wait time for device authorization.
|
||||
maxPollTotalWait = 15 * time.Minute
|
||||
// Aligned with the server-side Redis TTL (10 minutes).
|
||||
maxPollTotalWait = 10 * time.Minute
|
||||
)
|
||||
|
||||
type DeviceFlowProvider struct {
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
terminalBaseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
}
|
||||
|
||||
func NewDeviceFlowProvider(configDir string, logger *slog.Logger) *DeviceFlowProvider {
|
||||
return &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: ClientID(),
|
||||
scope: DefaultScopes,
|
||||
baseURL: DefaultDeviceBaseURL,
|
||||
logger: logger,
|
||||
Output: os.Stderr,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
configDir: configDir,
|
||||
clientID: ClientID(),
|
||||
scope: DefaultScopes,
|
||||
baseURL: DefaultDeviceBaseURL,
|
||||
terminalBaseURL: GetMCPBaseURL(),
|
||||
logger: logger,
|
||||
Output: os.Stderr,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
}
|
||||
}
|
||||
|
||||
@@ -66,6 +71,18 @@ func (p *DeviceFlowProvider) SetBaseURL(baseURL string) {
|
||||
p.baseURL = strings.TrimRight(baseURL, "/")
|
||||
}
|
||||
|
||||
// SetTerminalBaseURL sets the terminal API base URL for device flow polling.
|
||||
func (p *DeviceFlowProvider) SetTerminalBaseURL(baseURL string) {
|
||||
p.terminalBaseURL = strings.TrimRight(baseURL, "/")
|
||||
}
|
||||
|
||||
// SetScope overrides the OAuth scope for the device flow.
|
||||
func (p *DeviceFlowProvider) SetScope(scope string) {
|
||||
if p != nil {
|
||||
p.scope = scope
|
||||
}
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) output() io.Writer {
|
||||
if p != nil && p.Output != nil {
|
||||
return p.Output
|
||||
@@ -80,6 +97,7 @@ type DeviceAuthResponse struct {
|
||||
VerificationURIComplete string `json:"verificationUriComplete"`
|
||||
ExpiresIn int `json:"expiresIn"`
|
||||
Interval int `json:"interval"`
|
||||
FlowID string `json:"flowId"`
|
||||
}
|
||||
|
||||
type DeviceTokenResponse struct {
|
||||
@@ -88,6 +106,20 @@ type DeviceTokenResponse struct {
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
// DevicePollResponse represents the response from the terminal API poll endpoint.
|
||||
type DevicePollResponse struct {
|
||||
Success bool `json:"success"`
|
||||
Code string `json:"code,omitempty"`
|
||||
Message string `json:"message,omitempty"`
|
||||
Data DevicePollData `json:"data"`
|
||||
}
|
||||
|
||||
type DevicePollData struct {
|
||||
Status string `json:"status"`
|
||||
AuthCode string `json:"authCode,omitempty"`
|
||||
FlowID string `json:"flowId,omitempty"`
|
||||
}
|
||||
|
||||
type serviceResult struct {
|
||||
Success bool `json:"success"`
|
||||
Result json.RawMessage `json:"result"`
|
||||
@@ -96,6 +128,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 +198,89 @@ 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)
|
||||
}
|
||||
denialReason := classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
|
||||
if denialReason != "" {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
switch denialReason {
|
||||
case "user_forbidden":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织已禁止所有成员使用 CLI")))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织已禁止所有成员使用 CLI"))
|
||||
case "user_not_allowed":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 您不在该组织的 CLI 授权人员范围内")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员将您加入 CLI 授权人员名单。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员"))
|
||||
case "channel_not_allowed":
|
||||
ch := os.Getenv("DWS_CHANNEL")
|
||||
_, _ = fmt.Fprintf(p.output(), dfRed(i18n.T("⚠️ 当前渠道 %s 未获得该组织授权"))+"\n", ch)
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织管理员开通该渠道的访问权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf(i18n.T("当前渠道 %s 未获得该组织授权,请联系组织管理员"), ch)
|
||||
case "channel_required":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 当前组织已开启渠道管控")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请升级到最新版本的 CLI 后重试。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("当前组织已开启渠道管控,请升级到最新版本的 CLI 后重试"))
|
||||
case "no_auth":
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 认证已失效")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请执行 dws auth 重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("认证已失效,请执行 dws auth 重新登录"))
|
||||
default:
|
||||
// cli_not_enabled or unknown — show existing admin-apply flow
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
|
||||
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
|
||||
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
|
||||
maxAdmins := 3
|
||||
if len(admins.Result) < maxAdmins {
|
||||
maxAdmins = len(admins.Result)
|
||||
}
|
||||
var adminNames []string
|
||||
for i := 0; i < maxAdmins; i++ {
|
||||
adminNames = append(adminNames, admins.Result[i].Name)
|
||||
}
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("管理员操作入口:"), config.GetDeveloperSettingsURL())
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
|
||||
}
|
||||
}
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := SaveTokenData(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Always persist clientId to app.json so future process startups
|
||||
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
|
||||
if p.clientID != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
|
||||
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
oauthProvider.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
@@ -217,7 +346,91 @@ func (p *DeviceFlowProvider) pollDeviceToken(ctx context.Context, deviceCode str
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
// pollDeviceStatus polls the terminal API for device authorization status.
|
||||
//
|
||||
// Note: The server returns success=false for REJECTED and EXPIRED terminal
|
||||
// states (with a valid data.Status value). These are normal business outcomes,
|
||||
// not transport errors, so we return the response to the caller and let the
|
||||
// status-switch handle them.
|
||||
func (p *DeviceFlowProvider) pollDeviceStatus(ctx context.Context, flowID string) (*DevicePollResponse, error) {
|
||||
endpoint := fmt.Sprintf("%s%s?flowId=%s", p.terminalBaseURL, DevicePollPath, url.QueryEscape(flowID))
|
||||
body, err := p.doGet(ctx, endpoint)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var resp DevicePollResponse
|
||||
if err := json.Unmarshal(body, &resp); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("解析响应失败"), err)
|
||||
}
|
||||
// REJECTED/EXPIRED carry success=false but have a valid data.Status;
|
||||
// only treat as a real server error when data.Status is empty.
|
||||
if !resp.Success && resp.Data.Status == "" {
|
||||
return nil, fmt.Errorf("%s: [%s] %s", i18n.T("服务端返回错误"), resp.Code, resp.Message)
|
||||
}
|
||||
return &resp, nil
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorization(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
if auth.FlowID == "" {
|
||||
// Keep the pre-flowId device-code polling path for regular device flow
|
||||
// login responses that do not include terminal polling metadata.
|
||||
return p.waitForAuthorizationByDeviceCode(ctx, auth)
|
||||
}
|
||||
return p.waitForAuthorizationByFlowID(ctx, auth)
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorizationByFlowID(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
startTime := time.Now()
|
||||
interval := time.Duration(auth.Interval) * time.Second
|
||||
deadline := time.Duration(auth.ExpiresIn) * time.Second
|
||||
pollCount := 0
|
||||
|
||||
for {
|
||||
elapsed := time.Since(startTime)
|
||||
if elapsed >= maxPollTotalWait || elapsed >= deadline {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf("%s", i18n.Tf("设备授权码已过期(%d 秒),请重试", auth.ExpiresIn))
|
||||
}
|
||||
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-time.After(interval):
|
||||
}
|
||||
|
||||
pollCount++
|
||||
elapsedSec := int(time.Since(startTime).Seconds())
|
||||
dfPrintPollStatus(p.output(), pollCount, elapsedSec)
|
||||
|
||||
pollResp, err := p.pollDeviceStatus(ctx, auth.FlowID)
|
||||
if err != nil {
|
||||
dfPrintPollResult(p.output(), "network_error", i18n.T("网络错误,继续重试..."))
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("poll error", "error", err)
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
switch pollResp.Data.Status {
|
||||
case StatusApproved:
|
||||
dfPrintPollResult(p.output(), "authorized", i18n.T("授权成功!"))
|
||||
return &DeviceTokenResponse{AuthCode: pollResp.Data.AuthCode}, nil
|
||||
case StatusPending:
|
||||
dfPrintPollResult(p.output(), "pending", i18n.T("等待用户授权..."))
|
||||
case StatusRejected:
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("用户拒绝了授权请求"))
|
||||
case StatusExpired:
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("设备授权码已过期"))
|
||||
default:
|
||||
dfPrintPollResult(p.output(), "unknown", fmt.Sprintf(i18n.T("未知状态: %s"), pollResp.Data.Status))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) waitForAuthorizationByDeviceCode(ctx context.Context, auth *DeviceAuthResponse) (*DeviceTokenResponse, error) {
|
||||
startTime := time.Now()
|
||||
interval := time.Duration(auth.Interval) * time.Second
|
||||
deadline := time.Duration(auth.ExpiresIn) * time.Second
|
||||
@@ -297,6 +510,29 @@ func (p *DeviceFlowProvider) postForm(ctx context.Context, endpoint string, para
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// doGet performs an HTTP GET request and returns the response body.
|
||||
func (p *DeviceFlowProvider) doGet(ctx context.Context, endpoint string) ([]byte, error) {
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("创建请求失败"), err)
|
||||
}
|
||||
|
||||
resp, err := p.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("发送请求失败"), err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("读取响应失败"), err)
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncateBody(body, 200))
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
// truncateBody returns a string of at most maxLen bytes from body, appending
|
||||
// "...(truncated)" when the content exceeds the limit. This prevents leaking
|
||||
// potentially sensitive response payloads in error messages.
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
// Device flow authorization status constants.
|
||||
// Shared across device_flow.go and pat_auth_retry.go to avoid maintaining
|
||||
// string literals in multiple places.
|
||||
const (
|
||||
StatusPending = "PENDING"
|
||||
StatusApproved = "APPROVED"
|
||||
StatusRejected = "REJECTED"
|
||||
StatusExpired = "EXPIRED"
|
||||
StatusCancelled = "CANCELLED"
|
||||
)
|
||||
|
||||
// ParseDeviceFlowStatus normalizes a raw status string from the device flow
|
||||
// poll response into a canonical status constant. When the server returns an
|
||||
// empty status with success=false, it falls back to StatusExpired (server
|
||||
// error / flow not found).
|
||||
func ParseDeviceFlowStatus(rawStatus string, success bool) string {
|
||||
switch rawStatus {
|
||||
case StatusApproved, StatusRejected, StatusExpired, StatusPending, StatusCancelled:
|
||||
return rawStatus
|
||||
default:
|
||||
if rawStatus == "" && !success {
|
||||
return StatusExpired
|
||||
}
|
||||
return rawStatus
|
||||
}
|
||||
}
|
||||
@@ -14,15 +14,21 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
)
|
||||
|
||||
func newDeviceFlowTestLogger() *slog.Logger {
|
||||
@@ -46,6 +52,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)
|
||||
@@ -85,22 +95,42 @@ func TestWaitForAuthorizationSucceedsAfterPending(t *testing.T) {
|
||||
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// New terminal API uses GET method
|
||||
if r.Method != http.MethodGet {
|
||||
t.Fatalf("method = %s, want GET", r.Method)
|
||||
}
|
||||
if !strings.Contains(r.URL.RawQuery, "flowId=") {
|
||||
t.Fatal("flowId query parameter should be present")
|
||||
}
|
||||
if calls.Add(1) <= 2 {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
// Return PENDING status
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{"status": "PENDING"},
|
||||
})
|
||||
return
|
||||
}
|
||||
writeServiceResult(w, true, DeviceTokenResponse{AuthCode: "final-auth-code"}, "", "")
|
||||
// Return APPROVED status with authCode
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{
|
||||
"status": "APPROVED",
|
||||
"authCode": "final-auth-code",
|
||||
},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
provider.Output = io.Discard
|
||||
provider.SetBaseURL(server.URL)
|
||||
provider.SetTerminalBaseURL(server.URL)
|
||||
|
||||
resp, err := provider.waitForAuthorization(context.Background(), &DeviceAuthResponse{
|
||||
DeviceCode: "dc-1",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
FlowID: "test-flow-id",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("waitForAuthorization() error = %v", err)
|
||||
@@ -113,9 +143,87 @@ func TestWaitForAuthorizationSucceedsAfterPending(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationFallsBackToDeviceCodeWhenFlowIDMissing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var calls atomic.Int32
|
||||
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)
|
||||
}
|
||||
if err := r.ParseForm(); err != nil {
|
||||
t.Fatalf("ParseForm() error = %v", err)
|
||||
}
|
||||
if got := r.FormValue("device_code"); got != "legacy-device-code" {
|
||||
t.Fatalf("device_code = %q, want legacy-device-code", got)
|
||||
}
|
||||
if calls.Add(1) <= 2 {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
return
|
||||
}
|
||||
writeServiceResult(w, true, DeviceTokenResponse{AuthCode: "legacy-auth-code"}, "", "")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
var output bytes.Buffer
|
||||
provider.Output = &output
|
||||
provider.SetBaseURL(server.URL)
|
||||
|
||||
resp, err := provider.waitForAuthorization(context.Background(), &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 10,
|
||||
Interval: 1,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("waitForAuthorization() error = %v", err)
|
||||
}
|
||||
if resp.AuthCode != "legacy-auth-code" {
|
||||
t.Fatalf("auth code = %q, want legacy-auth-code", resp.AuthCode)
|
||||
}
|
||||
if calls.Load() != 3 {
|
||||
t.Fatalf("poll calls = %d, want 3", calls.Load())
|
||||
}
|
||||
if !strings.Contains(output.String(), i18n.T("等待用户授权...")) {
|
||||
t.Fatalf("expected device-code path to emit pending output, got %q", output.String())
|
||||
}
|
||||
if !strings.Contains(output.String(), i18n.T("授权成功!")) {
|
||||
t.Fatalf("expected device-code path to emit success output, got %q", output.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationHonorsContextCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// New terminal API uses GET method
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"success": true,
|
||||
"data": map[string]string{"status": "PENDING"},
|
||||
})
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
provider.Output = io.Discard
|
||||
provider.SetTerminalBaseURL(server.URL)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
if _, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
FlowID: "test-flow-id-2",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}); err == nil {
|
||||
t.Fatal("waitForAuthorization() error = nil, want context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationByDeviceCodeHonorsContextCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
writeServiceResult(w, true, DeviceTokenResponse{Error: "authorization_pending"}, "", "")
|
||||
}))
|
||||
@@ -128,11 +236,89 @@ func TestWaitForAuthorizationHonorsContextCancellation(t *testing.T) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 1500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
if _, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "dc-2",
|
||||
_, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}); err == nil {
|
||||
t.Fatal("waitForAuthorization() error = nil, want context cancellation")
|
||||
})
|
||||
if !errors.Is(err, context.DeadlineExceeded) {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want context deadline exceeded", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWaitForAuthorizationByDeviceCodeErrorStates(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
timeout time.Duration
|
||||
responses []DeviceTokenResponse
|
||||
wantErr string
|
||||
wantErrIs error
|
||||
wantOutput string
|
||||
}{
|
||||
{
|
||||
name: "slow_down_then_context_cancelled",
|
||||
timeout: 1500 * time.Millisecond,
|
||||
responses: []DeviceTokenResponse{{Error: "slow_down"}},
|
||||
wantErrIs: context.DeadlineExceeded,
|
||||
wantOutput: fmt.Sprintf(i18n.T("轮询过快,间隔增加至 %ds"), 6),
|
||||
},
|
||||
{
|
||||
name: "access_denied",
|
||||
timeout: 5 * time.Second,
|
||||
responses: []DeviceTokenResponse{{Error: "access_denied"}},
|
||||
wantErr: i18n.T("用户拒绝了授权请求"),
|
||||
wantOutput: fmt.Sprintf(i18n.T("[%d] 轮询中... (%ds)"), 1, 1),
|
||||
},
|
||||
{
|
||||
name: "expired_token",
|
||||
timeout: 5 * time.Second,
|
||||
responses: []DeviceTokenResponse{{Error: "expired_token"}},
|
||||
wantErr: i18n.T("设备授权码已过期"),
|
||||
wantOutput: fmt.Sprintf(i18n.T("[%d] 轮询中... (%ds)"), 1, 1),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
idx := int(calls.Add(1)) - 1
|
||||
if idx >= len(tt.responses) {
|
||||
idx = len(tt.responses) - 1
|
||||
}
|
||||
writeServiceResult(w, true, tt.responses[idx], "", "")
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := NewDeviceFlowProvider(t.TempDir(), newDeviceFlowTestLogger())
|
||||
var output bytes.Buffer
|
||||
provider.Output = &output
|
||||
provider.SetBaseURL(server.URL)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), tt.timeout)
|
||||
defer cancel()
|
||||
|
||||
_, err := provider.waitForAuthorization(ctx, &DeviceAuthResponse{
|
||||
DeviceCode: "legacy-device-code",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
})
|
||||
|
||||
if tt.wantErrIs != nil {
|
||||
if !errors.Is(err, tt.wantErrIs) {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want %v", err, tt.wantErrIs)
|
||||
}
|
||||
} else if err == nil || err.Error() != tt.wantErr {
|
||||
t.Fatalf("waitForAuthorization() error = %v, want %q", err, tt.wantErr)
|
||||
}
|
||||
|
||||
if !strings.Contains(output.String(), tt.wantOutput) {
|
||||
t.Fatalf("expected output to contain %q, got %q", tt.wantOutput, output.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
+197
-5
@@ -15,9 +15,37 @@ package auth
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_ID",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppKey (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CLIENT_SECRET",
|
||||
Category: configmeta.CategoryAuth,
|
||||
Description: "OAuth AppSecret (DingTalk 应用凭证)",
|
||||
DefaultValue: "(内置)",
|
||||
Sensitive: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CHANNEL",
|
||||
Category: configmeta.CategoryExternal,
|
||||
Description: "第三方渠道编码 (channelCode),如 Qoderwork",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
// AuthorizeURL is the DingTalk OAuth authorization page.
|
||||
AuthorizeURL = "https://login.dingtalk.com/oauth2/auth"
|
||||
@@ -56,18 +84,130 @@ const (
|
||||
// DeviceGrantType is the grant_type value defined by RFC 8628.
|
||||
DeviceGrantType = "urn:ietf:params:oauth:grant-type:device_code"
|
||||
|
||||
// Terminal API base URL for developer settings page.
|
||||
DefaultTerminalBaseURL = "https://open-dev.dingtalk.com"
|
||||
// DevicePollPath is the device flow polling path (used with MCP base URL).
|
||||
DevicePollPath = "/cli/oauth/device/poll"
|
||||
|
||||
// DeveloperSettingsPath is the path to the organization developer settings page.
|
||||
DeveloperSettingsPath = "/fe/old#/developerSettings"
|
||||
|
||||
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
|
||||
LogoutContinueURL = "https://login.dingtalk.com"
|
||||
|
||||
// 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"
|
||||
)
|
||||
|
||||
// GetTerminalBaseURL returns the terminal base URL with priority:
|
||||
// 1. ~/.dws/terminal_url file content (for pre-release environment)
|
||||
// 2. Default value (https://open-dev.dingtalk.com)
|
||||
func GetTerminalBaseURL() string {
|
||||
return config.GetTerminalBaseURL()
|
||||
}
|
||||
|
||||
// GetDeveloperSettingsURL returns the full URL to the organization developer
|
||||
// settings page, derived from the terminal base URL.
|
||||
func GetDeveloperSettingsURL() string {
|
||||
return config.GetDeveloperSettingsURL()
|
||||
}
|
||||
|
||||
// GetMCPBaseURL returns the MCP base URL with priority:
|
||||
// 1. ~/.dws/mcp_url file content (for pre-release environment)
|
||||
// 2. Default value (https://mcp.dingtalk.com)
|
||||
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 +224,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 +236,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 +265,45 @@ func ClientSecret() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
// Try loading from persisted app config (secret is in keychain)
|
||||
if _, secret := ResolveAppCredentials(getDefaultConfigDir()); secret != "" {
|
||||
return secret
|
||||
}
|
||||
if v := os.Getenv("DWS_CLIENT_SECRET"); v != "" {
|
||||
return v
|
||||
}
|
||||
return DefaultClientSecret
|
||||
}
|
||||
|
||||
// HasValidClientSecret returns true if a valid client secret is available.
|
||||
// A valid secret is one that is not a placeholder (e.g., <YOUR_CLIENT_SECRET>).
|
||||
func HasValidClientSecret() bool {
|
||||
secret := ClientSecret()
|
||||
return secret != "" && !strings.HasPrefix(secret, "<")
|
||||
}
|
||||
|
||||
// getRuntimeCredentials returns the runtime-override credentials if set.
|
||||
// Returns empty strings if no runtime overrides were provided.
|
||||
func getRuntimeCredentials() (clientID, clientSecret string) {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return runtimeClientID, runtimeClientSecret
|
||||
}
|
||||
|
||||
// getEnvClientID returns the environment variable client ID if set.
|
||||
func getEnvClientID() string {
|
||||
return os.Getenv("DWS_CLIENT_ID")
|
||||
}
|
||||
|
||||
// getDefaultConfigDir returns the default configuration directory.
|
||||
// Priority: DWS_CONFIG_DIR env var > ~/.dws
|
||||
func getDefaultConfigDir() string {
|
||||
if envDir := os.Getenv("DWS_CONFIG_DIR"); envDir != "" {
|
||||
return envDir
|
||||
}
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
return ".dws"
|
||||
}
|
||||
return filepath.Join(homeDir, ".dws")
|
||||
}
|
||||
|
||||
@@ -1,64 +0,0 @@
|
||||
// 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"
|
||||
"os"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ExportedCredentials represents the JSON structure of an exported credentials file,
|
||||
// used by auth import to restore credentials on another machine.
|
||||
type ExportedCredentials struct {
|
||||
RefreshToken string `json:"refresh_token,omitempty"`
|
||||
PersistentCode string `json:"persistent_code,omitempty"`
|
||||
CorpID string `json:"corp_id"`
|
||||
UserID string `json:"user_id,omitempty"`
|
||||
UserName string `json:"user_name,omitempty"`
|
||||
CorpName string `json:"corp_name,omitempty"`
|
||||
ExportedAt string `json:"exported_at"`
|
||||
}
|
||||
|
||||
func LoadExportedCredentials(ctx context.Context, path, configDir string) (string, error) {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("reading credentials file: %w", err)
|
||||
}
|
||||
|
||||
var creds ExportedCredentials
|
||||
if err := json.Unmarshal(b, &creds); err != nil {
|
||||
return "", fmt.Errorf("parsing credentials file: %w", err)
|
||||
}
|
||||
// Accept either persistent_code or refresh_token as a valid credential.
|
||||
if creds.PersistentCode == "" && creds.RefreshToken == "" {
|
||||
return "", fmt.Errorf("credentials file has no usable credential (need persistent_code or refresh_token)")
|
||||
}
|
||||
|
||||
data := &TokenData{
|
||||
PersistentCode: creds.PersistentCode,
|
||||
RefreshToken: creds.RefreshToken,
|
||||
RefreshExpAt: time.Now().Add(30 * 24 * time.Hour),
|
||||
CorpID: creds.CorpID,
|
||||
UserID: creds.UserID,
|
||||
UserName: creds.UserName,
|
||||
CorpName: creds.CorpName,
|
||||
}
|
||||
if err := SaveTokenData(configDir, data); err != nil {
|
||||
return "", fmt.Errorf("saving imported credentials: %w", err)
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
+104
-5
@@ -14,12 +14,14 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -27,13 +29,61 @@ const (
|
||||
lockRetryDelay = 50 * time.Millisecond
|
||||
)
|
||||
|
||||
// tokenFileLock provides cross-process file locking for token operations.
|
||||
// It prevents concurrent refresh from multiple CLI processes,
|
||||
// which can corrupt token data when two processes refresh simultaneously.
|
||||
//
|
||||
// ── Process-level lock ──────────────────────────────────────────────────
|
||||
// Prevents multiple goroutines within the same process from refreshing
|
||||
// simultaneously. Uses sync.Map with channel signaling for efficient waiting.
|
||||
|
||||
var processLocks sync.Map // map[string]chan struct{}
|
||||
|
||||
// processLockKey generates a unique key for process-level locking.
|
||||
func processLockKey(configDir string) string {
|
||||
return "refresh:" + configDir
|
||||
}
|
||||
|
||||
// acquireProcessLock attempts to acquire the process-level lock.
|
||||
// If another goroutine holds it, this blocks until that goroutine releases.
|
||||
// Returns a release function that MUST be called when done.
|
||||
func acquireProcessLock(ctx context.Context, configDir string) (release func(), waited bool, err error) {
|
||||
key := processLockKey(configDir)
|
||||
done := make(chan struct{})
|
||||
|
||||
for {
|
||||
// Try to store our channel; if successful, we own the lock
|
||||
if existing, loaded := processLocks.LoadOrStore(key, done); !loaded {
|
||||
// We got the lock
|
||||
return func() {
|
||||
close(done)
|
||||
processLocks.Delete(key)
|
||||
}, waited, nil
|
||||
} else {
|
||||
// Another goroutine holds the lock; wait for it
|
||||
ch, ok := existing.(chan struct{})
|
||||
if !ok {
|
||||
// Unexpected type; delete and retry
|
||||
processLocks.Delete(key)
|
||||
continue
|
||||
}
|
||||
waited = true
|
||||
select {
|
||||
case <-ch:
|
||||
// Lock released; retry to acquire
|
||||
continue
|
||||
case <-ctx.Done():
|
||||
return nil, waited, ctx.Err()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── File-level lock ─────────────────────────────────────────────────────
|
||||
// Prevents multiple CLI processes from refreshing simultaneously.
|
||||
// Platform support:
|
||||
// - Unix/macOS: flock(2) system call
|
||||
// - Windows: LockFileEx / UnlockFileEx from kernel32.dll
|
||||
|
||||
// tokenFileLock provides cross-process file locking for token operations.
|
||||
// It prevents concurrent refresh from multiple CLI processes,
|
||||
// which can corrupt token data when two processes refresh simultaneously.
|
||||
type tokenFileLock struct {
|
||||
path string
|
||||
file *os.File
|
||||
@@ -76,3 +126,52 @@ func (l *tokenFileLock) release() {
|
||||
l.file = nil
|
||||
}
|
||||
}
|
||||
|
||||
// ── Dual-layer lock ─────────────────────────────────────────────────────
|
||||
// Combines process-level and file-level locks for comprehensive protection.
|
||||
|
||||
// DualLock holds both process-level and file-level locks.
|
||||
type DualLock struct {
|
||||
processRelease func()
|
||||
fileLock *tokenFileLock
|
||||
Waited bool // true if we waited for another goroutine/process
|
||||
}
|
||||
|
||||
// AcquireDualLock acquires both process-level and file-level locks.
|
||||
// This provides comprehensive protection against:
|
||||
// 1. Multiple goroutines in the same process (sync.Map)
|
||||
// 2. Multiple CLI processes (file lock)
|
||||
//
|
||||
// The caller MUST call Release() when done.
|
||||
func AcquireDualLock(ctx context.Context, configDir string) (*DualLock, error) {
|
||||
// 1. Acquire process-level lock first (fast, in-memory)
|
||||
processRelease, waited, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("acquiring process lock: %w", err)
|
||||
}
|
||||
|
||||
// 2. Acquire file-level lock (cross-process)
|
||||
fileLock, err := acquireTokenLock(configDir)
|
||||
if err != nil {
|
||||
processRelease() // Release process lock on failure
|
||||
return nil, fmt.Errorf("acquiring file lock: %w", err)
|
||||
}
|
||||
|
||||
return &DualLock{
|
||||
processRelease: processRelease,
|
||||
fileLock: fileLock,
|
||||
Waited: waited,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Release releases both locks in reverse order.
|
||||
func (d *DualLock) Release() {
|
||||
if d.fileLock != nil {
|
||||
d.fileLock.release()
|
||||
d.fileLock = nil
|
||||
}
|
||||
if d.processRelease != nil {
|
||||
d.processRelease()
|
||||
d.processRelease = nil
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
@@ -145,3 +147,203 @@ func TestAcquireTokenLock_LockFilePermissions(t *testing.T) {
|
||||
t.Fatalf("lock file permissions = %o, want 0600", perm)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── Process-level lock tests ───────────────────────────────────────────
|
||||
|
||||
func TestAcquireProcessLock_Basic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
|
||||
release, waited, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("acquireProcessLock() error = %v", err)
|
||||
}
|
||||
if waited {
|
||||
t.Fatal("should not have waited on first acquisition")
|
||||
}
|
||||
|
||||
release()
|
||||
|
||||
// Should be able to re-acquire after release
|
||||
release2, waited2, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("re-acquire after release error = %v", err)
|
||||
}
|
||||
if waited2 {
|
||||
t.Fatal("should not have waited on re-acquisition")
|
||||
}
|
||||
release2()
|
||||
}
|
||||
|
||||
func TestAcquireProcessLock_Contention(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
|
||||
// Goroutine 1 acquires the lock first
|
||||
release1, _, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("acquireProcessLock() g1 error = %v", err)
|
||||
}
|
||||
|
||||
acquired := make(chan bool, 1)
|
||||
var g2Waited bool
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(1)
|
||||
|
||||
// Goroutine 2 tries to acquire — should block until g1 releases
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
release2, waited, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
acquired <- false
|
||||
return
|
||||
}
|
||||
g2Waited = waited
|
||||
acquired <- true
|
||||
release2()
|
||||
}()
|
||||
|
||||
// Give goroutine 2 a moment to start blocking
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Verify goroutine 2 has not acquired yet
|
||||
select {
|
||||
case <-acquired:
|
||||
t.Fatal("goroutine 2 should not have acquired the lock while goroutine 1 holds it")
|
||||
default:
|
||||
// Expected: goroutine 2 is still waiting
|
||||
}
|
||||
|
||||
// Release lock1 so goroutine 2 can proceed
|
||||
release1()
|
||||
|
||||
// Wait for goroutine 2 to finish
|
||||
wg.Wait()
|
||||
|
||||
if !g2Waited {
|
||||
t.Fatal("goroutine 2 should have reported that it waited")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAcquireProcessLock_ContextCancellation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
|
||||
// Goroutine 1 holds the lock
|
||||
release1, _, err := acquireProcessLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("acquireProcessLock() g1 error = %v", err)
|
||||
}
|
||||
defer release1()
|
||||
|
||||
// Goroutine 2 tries to acquire with a cancellable context
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, _, err := acquireProcessLock(ctx, configDir)
|
||||
done <- err
|
||||
}()
|
||||
|
||||
// Give goroutine 2 time to start waiting
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
// Cancel the context
|
||||
cancel()
|
||||
|
||||
// Goroutine 2 should return with context.Canceled
|
||||
select {
|
||||
case err := <-done:
|
||||
if err != context.Canceled {
|
||||
t.Fatalf("expected context.Canceled, got %v", err)
|
||||
}
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("goroutine 2 did not return after context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
// ─── Dual-layer lock tests ──────────────────────────────────────────────
|
||||
|
||||
func TestAcquireDualLock_Basic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
|
||||
lock, err := AcquireDualLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("AcquireDualLock() error = %v", err)
|
||||
}
|
||||
if lock.Waited {
|
||||
t.Fatal("should not have waited on first acquisition")
|
||||
}
|
||||
|
||||
lock.Release()
|
||||
|
||||
// Should be able to re-acquire after release
|
||||
lock2, err := AcquireDualLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("re-acquire after release error = %v", err)
|
||||
}
|
||||
lock2.Release()
|
||||
}
|
||||
|
||||
func TestAcquireDualLock_DoubleRelease(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
|
||||
lock, err := AcquireDualLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("AcquireDualLock() error = %v", err)
|
||||
}
|
||||
|
||||
// First release should work fine
|
||||
lock.Release()
|
||||
|
||||
// Second release should not panic
|
||||
lock.Release()
|
||||
}
|
||||
|
||||
func TestAcquireDualLock_ConcurrentGoroutines(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
configDir := t.TempDir()
|
||||
ctx := context.Background()
|
||||
|
||||
const numGoroutines = 10
|
||||
var counter int64
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(numGoroutines)
|
||||
|
||||
// Launch multiple goroutines that all try to increment a counter
|
||||
// while holding the dual lock. If locking works correctly,
|
||||
// the final counter value should be numGoroutines.
|
||||
for i := 0; i < numGoroutines; i++ {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
lock, err := AcquireDualLock(ctx, configDir)
|
||||
if err != nil {
|
||||
t.Errorf("AcquireDualLock() error = %v", err)
|
||||
return
|
||||
}
|
||||
defer lock.Release()
|
||||
|
||||
// Critical section: read-modify-write
|
||||
current := atomic.LoadInt64(&counter)
|
||||
time.Sleep(time.Millisecond) // Simulate some work
|
||||
atomic.StoreInt64(&counter, current+1)
|
||||
}()
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
|
||||
if counter != numGoroutines {
|
||||
t.Fatalf("counter = %d, want %d (race condition detected)", counter, numGoroutines)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
var (
|
||||
migrationOnce sync.Once
|
||||
migrationDone bool
|
||||
)
|
||||
|
||||
// SaveTokenDataKeychain saves TokenData to the platform keychain.
|
||||
// This is the new secure storage method using random master key.
|
||||
func SaveTokenDataKeychain(data *TokenData) error {
|
||||
jsonData, err := json.MarshalIndent(data, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal token data: %w", err)
|
||||
}
|
||||
// Zero sensitive data after use
|
||||
defer func() {
|
||||
for i := range jsonData {
|
||||
jsonData[i] = 0
|
||||
}
|
||||
}()
|
||||
|
||||
if err := keychain.Set(keychain.Service, keychain.AccountToken, string(jsonData)); err != nil {
|
||||
return fmt.Errorf("save to keychain: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadTokenDataKeychain loads TokenData from the platform keychain.
|
||||
func LoadTokenDataKeychain() (*TokenData, error) {
|
||||
jsonStr, err := keychain.Get(keychain.Service, keychain.AccountToken)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load from keychain: %w", err)
|
||||
}
|
||||
if jsonStr == "" {
|
||||
return nil, fmt.Errorf("no token data in keychain")
|
||||
}
|
||||
|
||||
var data TokenData
|
||||
if err := json.Unmarshal([]byte(jsonStr), &data); err != nil {
|
||||
return nil, fmt.Errorf("parse token data: %w", err)
|
||||
}
|
||||
return &data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenDataKeychain removes TokenData from the platform keychain.
|
||||
func DeleteTokenDataKeychain() error {
|
||||
return keychain.Remove(keychain.Service, keychain.AccountToken)
|
||||
}
|
||||
|
||||
// TokenDataExistsKeychain checks if token data exists in keychain.
|
||||
func TokenDataExistsKeychain() bool {
|
||||
return keychain.Exists(keychain.Service, keychain.AccountToken)
|
||||
}
|
||||
|
||||
// EnsureMigration performs one-time migration from legacy .data to keychain.
|
||||
// This should be called early in the auth flow (e.g., during GetAccessToken).
|
||||
// The migration is idempotent and thread-safe.
|
||||
func EnsureMigration(configDir string, logger *slog.Logger) {
|
||||
migrationOnce.Do(func() {
|
||||
result := keychain.MigrateFromLegacy(configDir)
|
||||
migrationDone = true
|
||||
|
||||
if result.Migrated {
|
||||
if logger != nil {
|
||||
logger.Info("migrated token data to secure keychain storage",
|
||||
"from", result.FromPath,
|
||||
"backup", result.BackupPath)
|
||||
}
|
||||
} else if result.NeedRelogin {
|
||||
if logger != nil {
|
||||
logger.Warn("cannot migrate legacy token data, please re-login",
|
||||
"error", result.Error)
|
||||
}
|
||||
} else if result.Error != nil {
|
||||
if logger != nil {
|
||||
logger.Error("migration failed", "error", result.Error)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// IsMigrationDone returns true if migration has been attempted.
|
||||
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"
|
||||
|
||||
+1280
-12
File diff suppressed because it is too large
Load Diff
+387
-24
@@ -15,15 +15,18 @@ package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
)
|
||||
|
||||
// oauthHTTPClient is a dedicated HTTP client for OAuth operations with
|
||||
@@ -91,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 {
|
||||
@@ -99,15 +119,80 @@ 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
|
||||
denialReason string
|
||||
}
|
||||
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("回调中未收到授权码")):
|
||||
@@ -117,14 +202,160 @@ 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)
|
||||
var denialReason string
|
||||
if statusErr != nil {
|
||||
denialReason = "unknown"
|
||||
} else {
|
||||
denialReason = classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
|
||||
}
|
||||
cliAuthEnabled := denialReason == ""
|
||||
|
||||
// Update CLI auth disabled state
|
||||
callbackTokenMu.Lock()
|
||||
callbackAuthDisabled = !cliAuthEnabled
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Display appropriate HTML based on auth status and denial reason
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
switch {
|
||||
case cliAuthEnabled:
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
case denialReason == "user_forbidden" || denialReason == "user_not_allowed":
|
||||
_, _ = fmt.Fprint(w, accessDeniedHTML)
|
||||
case denialReason == "channel_not_allowed" || denialReason == "channel_required":
|
||||
_, _ = fmt.Fprint(w, channelDeniedHTML)
|
||||
default:
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
}
|
||||
// Ensure response is flushed to client
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
// Notify main goroutine with full result
|
||||
select {
|
||||
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled, denialReason: denialReason}:
|
||||
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}
|
||||
@@ -160,9 +391,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:
|
||||
@@ -171,13 +402,99 @@ 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 - for terminal denial reasons, exit immediately
|
||||
// (page shows accessDeniedHTML/channelDeniedHTML with no apply button,
|
||||
// so polling for apply submission would hang forever).
|
||||
// Error messages are kept consistent with the text shown on the HTML pages.
|
||||
if result.cliAuthDisabled {
|
||||
switch result.denialReason {
|
||||
case "user_forbidden", "user_not_allowed":
|
||||
return nil, errors.New(i18n.T("您不在该组织的 CLI 授权人员范围内,请联系组织管理员将您加入授权名单"))
|
||||
case "channel_not_allowed", "channel_required":
|
||||
return nil, errors.New(i18n.T("当前渠道未获得该组织授权,或组织已开启渠道管控,请联系组织管理员开通渠道访问权限,或升级到最新版本的 CLI"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
|
||||
|
||||
// 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 && classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL")) == "" {
|
||||
_, _ = 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)
|
||||
}
|
||||
|
||||
// Always persist clientId to app.json so future process startups
|
||||
// can load it via ResolveAppCredentials and populate DWS_CLIENT_ID env.
|
||||
if p.clientID != "" {
|
||||
_ = os.Setenv("DWS_CLIENT_ID", p.clientID)
|
||||
_ = SaveAppConfig(p.configDir, &AppConfig{ClientID: p.clientID})
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
p.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
@@ -209,19 +526,34 @@ func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
|
||||
return "", errors.New(i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"))
|
||||
}
|
||||
|
||||
// lockedRefresh attempts to refresh the token while holding a cross-process file lock.
|
||||
// It uses a double-check pattern: after acquiring the lock it re-loads from disk,
|
||||
// because another process may have already completed the refresh while we waited.
|
||||
// This prevents the classic race where two CLI processes both see an expired token
|
||||
// and both call the refresh API, invalidating each other's refresh_token.
|
||||
// lockedRefresh attempts to refresh the token while holding dual-layer locks.
|
||||
// It uses a double-check pattern with both process-level and file-level locking:
|
||||
//
|
||||
// Layer 1 (Process Lock - sync.Map):
|
||||
//
|
||||
// Prevents multiple goroutines within the same process from refreshing simultaneously.
|
||||
// If another goroutine is already refreshing, we wait for it and then re-check.
|
||||
//
|
||||
// Layer 2 (File Lock - flock/LockFileEx):
|
||||
//
|
||||
// Prevents multiple CLI processes from refreshing simultaneously.
|
||||
// If another process is refreshing, we wait for the file lock and then re-check.
|
||||
//
|
||||
// Double-Check Pattern:
|
||||
//
|
||||
// After acquiring the lock, we re-load from disk because another goroutine/process
|
||||
// may have already completed the refresh while we were waiting. This prevents the
|
||||
// classic race where two callers both see an expired token and both call the
|
||||
// refresh API, invalidating each other's refresh_token.
|
||||
func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
|
||||
lock, err := acquireTokenLock(p.configDir)
|
||||
// Acquire dual-layer lock (process-level + file-level)
|
||||
lock, err := AcquireDualLock(ctx, p.configDir)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("acquiring token lock: %w", err)
|
||||
return nil, fmt.Errorf("acquiring dual lock: %w", err)
|
||||
}
|
||||
defer lock.release()
|
||||
defer lock.Release()
|
||||
|
||||
// Double-check: re-load from disk — another process may have refreshed
|
||||
// Double-check: re-load from disk — another goroutine/process may have refreshed
|
||||
// while we were waiting for the lock.
|
||||
data, err := LoadTokenData(p.configDir)
|
||||
if err != nil {
|
||||
@@ -229,7 +561,11 @@ func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
|
||||
}
|
||||
if data.IsAccessTokenValid() {
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("token already refreshed by another process")
|
||||
if lock.Waited {
|
||||
p.logger.Debug("token already refreshed by another goroutine/process")
|
||||
} else {
|
||||
p.logger.Debug("token still valid after acquiring lock")
|
||||
}
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
@@ -240,7 +576,7 @@ func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
|
||||
}
|
||||
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("refreshing token (locked)")
|
||||
p.logger.Debug("refreshing token (dual-locked)")
|
||||
}
|
||||
return p.refreshWithRefreshToken(ctx, data)
|
||||
}
|
||||
@@ -270,3 +606,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"
|
||||
@@ -101,13 +101,34 @@ func SaveSecureTokenData(configDir string, data *TokenData) error {
|
||||
finalPath := filepath.Join(configDir, secureDataFile)
|
||||
tmpPath := finalPath + ".tmp"
|
||||
|
||||
if err := os.WriteFile(tmpPath, ciphertext, config.FilePerm); err != nil {
|
||||
// Atomic write with fsync to ensure data durability
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, config.FilePerm)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating tmp file: %w", err)
|
||||
}
|
||||
|
||||
writeSuccess := false
|
||||
defer func() {
|
||||
if !writeSuccess {
|
||||
tmpFile.Close()
|
||||
_ = os.Remove(tmpPath)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := tmpFile.Write(ciphertext); err != nil {
|
||||
return fmt.Errorf("writing tmp file: %w", err)
|
||||
}
|
||||
if err := tmpFile.Sync(); err != nil {
|
||||
return fmt.Errorf("syncing tmp file: %w", err)
|
||||
}
|
||||
if err := tmpFile.Close(); err != nil {
|
||||
return fmt.Errorf("closing tmp file: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpPath, finalPath); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return fmt.Errorf("renaming tmp to final: %w", err)
|
||||
}
|
||||
writeSuccess = true
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
+135
-8
@@ -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,26 +65,104 @@ func (t *TokenData) HasPersistentCode() bool {
|
||||
return t != nil && t.PersistentCode != ""
|
||||
}
|
||||
|
||||
// SaveTokenData encrypts and saves TokenData to .data file.
|
||||
// Uses AES-256-GCM encryption with a key derived from device MAC address.
|
||||
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 {
|
||||
return SaveSecureTokenData(configDir, data)
|
||||
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 encrypted .data file.
|
||||
// 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) {
|
||||
return LoadSecureTokenData(configDir)
|
||||
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()
|
||||
}
|
||||
data, err := LoadSecureTokenData(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := SaveTokenDataKeychain(data); err == nil {
|
||||
_ = DeleteSecureData(configDir)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes encrypted .data file from configDir.
|
||||
// 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 {
|
||||
return DeleteSecureData(configDir)
|
||||
if h := edition.Get(); h.DeleteToken != nil {
|
||||
return h.DeleteToken(configDir)
|
||||
}
|
||||
keychainErr := DeleteTokenDataKeychain()
|
||||
legacyErr := DeleteSecureData(configDir)
|
||||
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)
|
||||
@@ -114,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
|
||||
}
|
||||
|
||||
+59
-46
@@ -14,19 +14,22 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
func setupTestMAC(t *testing.T) {
|
||||
// cleanupKeychain removes test data from keychain after test completes.
|
||||
func cleanupKeychain(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Cleanup(func() {
|
||||
_ = keychain.Remove(keychain.Service, keychain.AccountToken)
|
||||
})
|
||||
}
|
||||
|
||||
func TestTokenSaveLoadAndDelete(t *testing.T) {
|
||||
setupTestMAC(t)
|
||||
cleanupKeychain(t)
|
||||
|
||||
configDir := t.TempDir()
|
||||
now := time.Now().UTC()
|
||||
@@ -42,33 +45,17 @@ func TestTokenSaveLoadAndDelete(t *testing.T) {
|
||||
CorpName: "测试科技",
|
||||
}
|
||||
|
||||
// Save to keychain
|
||||
if err := SaveTokenData(configDir, original); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify .data file was created with correct permissions.
|
||||
dataPath := filepath.Join(configDir, secureDataFile)
|
||||
info, err := os.Stat(dataPath)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat(.data) error = %v", err)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Fatalf(".data perms = %o, want 600", info.Mode().Perm())
|
||||
}
|
||||
if _, err := os.Stat(dataPath + ".tmp"); !os.IsNotExist(err) {
|
||||
t.Fatalf(".data.tmp should not remain, stat err = %v", err)
|
||||
}
|
||||
|
||||
// Verify .data file is NOT valid plaintext JSON (it's encrypted).
|
||||
raw, err := os.ReadFile(dataPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(.data) error = %v", err)
|
||||
}
|
||||
var probe map[string]any
|
||||
if json.Unmarshal(raw, &probe) == nil {
|
||||
t.Fatal("saved .data should be encrypted, not plain JSON")
|
||||
// Verify data exists in keychain
|
||||
if !TokenDataExistsKeychain() {
|
||||
t.Fatal("TokenDataExistsKeychain() should be true after save")
|
||||
}
|
||||
|
||||
// Load and verify
|
||||
loaded, err := LoadTokenData(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenData() error = %v", err)
|
||||
@@ -83,47 +70,71 @@ func TestTokenSaveLoadAndDelete(t *testing.T) {
|
||||
t.Fatalf("loaded corp_id = %q, want %q", loaded.CorpID, original.CorpID)
|
||||
}
|
||||
|
||||
// Delete and verify
|
||||
if err := DeleteTokenData(configDir); err != nil {
|
||||
t.Fatalf("DeleteTokenData() error = %v", err)
|
||||
}
|
||||
if TokenDataExistsKeychain() {
|
||||
t.Fatal("TokenDataExistsKeychain() should be false after delete")
|
||||
}
|
||||
if _, err := LoadTokenData(configDir); err == nil {
|
||||
t.Fatal("LoadTokenData() error = nil after delete, want failure")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenDecryptionFailsWithCorruptedData(t *testing.T) {
|
||||
func TestTokenOverwrite(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
|
||||
configDir := t.TempDir()
|
||||
data := &TokenData{
|
||||
AccessToken: "at_test",
|
||||
RefreshToken: "rt_test",
|
||||
|
||||
// Save first version
|
||||
data1 := &TokenData{
|
||||
AccessToken: "at_v1",
|
||||
RefreshToken: "rt_v1",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
CorpID: "corp_v1",
|
||||
}
|
||||
if err := SaveTokenData(configDir, data); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
if err := SaveTokenData(configDir, data1); err != nil {
|
||||
t.Fatalf("SaveTokenData(v1) error = %v", err)
|
||||
}
|
||||
|
||||
dataPath := filepath.Join(configDir, secureDataFile)
|
||||
raw, err := os.ReadFile(dataPath)
|
||||
// Save second version (overwrite)
|
||||
data2 := &TokenData{
|
||||
AccessToken: "at_v2",
|
||||
RefreshToken: "rt_v2",
|
||||
ExpiresAt: time.Now().Add(2 * time.Hour),
|
||||
RefreshExpAt: time.Now().Add(48 * time.Hour),
|
||||
CorpID: "corp_v2",
|
||||
}
|
||||
if err := SaveTokenData(configDir, data2); err != nil {
|
||||
t.Fatalf("SaveTokenData(v2) error = %v", err)
|
||||
}
|
||||
|
||||
// Load should return v2
|
||||
loaded, err := LoadTokenData(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(.data) error = %v", err)
|
||||
t.Fatalf("LoadTokenData() error = %v", err)
|
||||
}
|
||||
raw[len(raw)-1] ^= 0xFF
|
||||
if err := os.WriteFile(dataPath, raw, 0o600); err != nil {
|
||||
t.Fatalf("WriteFile(.data) error = %v", err)
|
||||
if loaded.AccessToken != "at_v2" {
|
||||
t.Fatalf("access_token = %q, want %q", loaded.AccessToken, "at_v2")
|
||||
}
|
||||
|
||||
if _, err := LoadTokenData(configDir); err == nil {
|
||||
t.Fatal("LoadTokenData with corrupted ciphertext should fail")
|
||||
if loaded.CorpID != "corp_v2" {
|
||||
t.Fatalf("corp_id = %q, want %q", loaded.CorpID, "corp_v2")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSecureDataExists(t *testing.T) {
|
||||
func TestTokenDataExistsKeychain(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
|
||||
configDir := t.TempDir()
|
||||
if SecureDataExists(configDir) {
|
||||
t.Fatal("SecureDataExists() should be false before save")
|
||||
|
||||
// Should be false before save
|
||||
if TokenDataExistsKeychain() {
|
||||
t.Fatal("TokenDataExistsKeychain() should be false before save")
|
||||
}
|
||||
|
||||
// Save data
|
||||
data := &TokenData{
|
||||
AccessToken: "at_test",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
@@ -131,8 +142,10 @@ func TestSecureDataExists(t *testing.T) {
|
||||
if err := SaveTokenData(configDir, data); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
if !SecureDataExists(configDir) {
|
||||
t.Fatal("SecureDataExists() should be true after save")
|
||||
|
||||
// Should be true after save
|
||||
if !TokenDataExistsKeychain() {
|
||||
t.Fatal("TokenDataExistsKeychain() should be true after save")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Vendored
+31
-1
@@ -258,7 +258,37 @@ func (s *Store) saveJSON(path string, value any) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(path, data, 0o600)
|
||||
|
||||
// Atomic write with fsync to ensure data durability
|
||||
tmpPath := path + ".tmp"
|
||||
tmpFile, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o600)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
writeSuccess := false
|
||||
defer func() {
|
||||
if !writeSuccess {
|
||||
tmpFile.Close()
|
||||
_ = os.Remove(tmpPath)
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := tmpFile.Write(data); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tmpFile.Sync(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tmpFile.Close(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := os.Rename(tmpPath, path); err != nil {
|
||||
_ = os.Remove(tmpPath)
|
||||
return err
|
||||
}
|
||||
writeSuccess = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) loadJSON(path string, out any) error {
|
||||
|
||||
+226
-42
@@ -20,16 +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"
|
||||
)
|
||||
|
||||
@@ -56,7 +58,7 @@ type FlagSpec struct {
|
||||
Description string
|
||||
}
|
||||
|
||||
func NewMCPCommand(ctx context.Context, loader CatalogLoader, runner executor.Runner) *cobra.Command {
|
||||
func NewMCPCommand(ctx context.Context, loader CatalogLoader, runner executor.Runner, engine *pipeline.Engine) *cobra.Command {
|
||||
catalog, loadErr := loader.Load(ctx)
|
||||
|
||||
longDescription := "Reserved canonical runtime surface. Tools are generated from the shared Tool IR under dws mcp."
|
||||
@@ -94,43 +96,68 @@ func NewMCPCommand(ctx context.Context, loader CatalogLoader, runner executor.Ru
|
||||
if product.CLI != nil && product.CLI.Skip {
|
||||
continue
|
||||
}
|
||||
productCommand := newProductCommand(product, runner)
|
||||
productCommand := newProductCommand(product, runner, engine)
|
||||
cmd.AddCommand(productCommand)
|
||||
addGroupedProductAlias(cmd, product, runner)
|
||||
addGroupedProductAlias(cmd, product, runner, engine)
|
||||
}
|
||||
return cmd
|
||||
}
|
||||
|
||||
func NewSchemaCommand(loader CatalogLoader) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "schema [canonical-product.tool]",
|
||||
Short: "Inspect canonical schema metadata",
|
||||
return &cobra.Command{
|
||||
Use: "schema [product.tool]",
|
||||
Short: "查看 MCP 工具 Schema (产品列表 / 工具参数)",
|
||||
Long: `查看已发现的 MCP 产品和工具的 Schema 元数据。
|
||||
|
||||
不带参数时列出所有产品及其工具数量;带 product.tool 路径时
|
||||
输出该工具的完整输入 Schema(JSON Schema 格式)。
|
||||
|
||||
示例:
|
||||
dws schema # 列出所有产品
|
||||
dws schema aitable.query_records # 查看 aitable query_records 的参数 Schema
|
||||
dws schema --fields id,tools # 只显示 id 和 tools 字段
|
||||
dws schema --jq '.products[].id' # 用 jq 提取所有产品 ID`,
|
||||
Args: cobra.MaximumNArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
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
|
||||
}
|
||||
|
||||
jsonOut, err := cmd.Flags().GetBool("json")
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read schema flags")
|
||||
}
|
||||
|
||||
payload, err := schemaPayload(catalog, args)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if jsonOut {
|
||||
return output.WriteJSON(cmd.OutOrStdout(), payload)
|
||||
}
|
||||
return writeSchemaText(cmd.OutOrStdout(), payload)
|
||||
return output.WriteFiltered(
|
||||
cmd.OutOrStdout(),
|
||||
output.ResolveFormat(cmd, output.FormatJSON),
|
||||
payload,
|
||||
output.ResolveFields(cmd),
|
||||
output.ResolveJQ(cmd),
|
||||
)
|
||||
},
|
||||
}
|
||||
cmd.Flags().Bool("json", false, "Emit schema metadata as JSON")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func BuildFlagSpecs(schema map[string]any, hints map[string]ir.CLIFlagHint) []FlagSpec {
|
||||
@@ -169,7 +196,7 @@ func BuildFlagSpecs(schema map[string]any, hints map[string]ir.CLIFlagHint) []Fl
|
||||
return specs
|
||||
}
|
||||
|
||||
func newProductCommand(product ir.CanonicalProduct, runner executor.Runner) *cobra.Command {
|
||||
func newProductCommand(product ir.CanonicalProduct, runner executor.Runner, engine *pipeline.Engine) *cobra.Command {
|
||||
shortDescription := product.DisplayName
|
||||
if strings.TrimSpace(product.Description) != "" {
|
||||
shortDescription = product.Description
|
||||
@@ -204,12 +231,33 @@ func newProductCommand(product ir.CanonicalProduct, runner executor.Runner) *cob
|
||||
}
|
||||
|
||||
for _, tool := range product.Tools {
|
||||
cmd.AddCommand(newToolCommand(product, tool, runner))
|
||||
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
|
||||
}
|
||||
|
||||
func addGroupedProductAlias(root *cobra.Command, product ir.CanonicalProduct, runner executor.Runner) {
|
||||
func addGroupedProductAlias(root *cobra.Command, product ir.CanonicalProduct, runner executor.Runner, engine *pipeline.Engine) {
|
||||
if root == nil || product.CLI == nil {
|
||||
return
|
||||
}
|
||||
@@ -260,7 +308,7 @@ func addGroupedProductAlias(root *cobra.Command, product ir.CanonicalProduct, ru
|
||||
cliCopy.Group = ""
|
||||
aliasProduct.CLI = &cliCopy
|
||||
}
|
||||
productCommand := newProductCommand(aliasProduct, runner)
|
||||
productCommand := newProductCommand(aliasProduct, runner, engine)
|
||||
productCommand.Use = leaf
|
||||
productCommand.Aliases = nil
|
||||
if leaf != aliasProduct.ID {
|
||||
@@ -269,7 +317,7 @@ func addGroupedProductAlias(root *cobra.Command, product ir.CanonicalProduct, ru
|
||||
parent.AddCommand(productCommand)
|
||||
}
|
||||
|
||||
func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner executor.Runner) *cobra.Command {
|
||||
func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner executor.Runner, engine *pipeline.Engine) *cobra.Command {
|
||||
shortDescription := tool.Title
|
||||
if strings.TrimSpace(tool.Description) != "" {
|
||||
shortDescription = tool.Description
|
||||
@@ -303,43 +351,143 @@ func newToolCommand(product ir.CanonicalProduct, tool ir.ToolDescriptor, runner
|
||||
}
|
||||
dryRun = value
|
||||
}
|
||||
|
||||
// One guard per invocation ensures stdin is read at most once.
|
||||
guard := NewStdinGuard()
|
||||
|
||||
jsonPayload, err := cmd.Flags().GetString("json")
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read --json")
|
||||
}
|
||||
|
||||
// Resolve @file / @- for --json flag.
|
||||
jsonPayload, err = ResolveInputSource(jsonPayload, "json", guard)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
paramsPayload, err := cmd.Flags().GetString("params")
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read --params")
|
||||
}
|
||||
overrides, err := collectOverrides(cmd, specs)
|
||||
|
||||
// Resolve @file / @- for all string-typed override flags BEFORE
|
||||
// the implicit stdin fallback, so explicit @- in any flag takes
|
||||
// priority over the implicit pipe read.
|
||||
overrides, err := collectOverrides(cmd, specs, guard)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Implicit stdin fallback (lowest priority): if no --json was
|
||||
// given and no flag claimed stdin via @-, read from pipe.
|
||||
if jsonPayload == "" && !guard.Claimed() && StdinIsPipe() {
|
||||
if claimErr := guard.Claim("implicit stdin (pipe)"); claimErr != nil {
|
||||
return claimErr
|
||||
}
|
||||
stdinData, stdinErr := ReadStdin()
|
||||
if stdinErr != nil {
|
||||
return stdinErr
|
||||
}
|
||||
jsonPayload = stdinData
|
||||
}
|
||||
|
||||
params, err := executor.MergePayloads(jsonPayload, paramsPayload, overrides)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// PostParse: normalise parameter values (date formats,
|
||||
// booleans, enums) using the tool's input schema.
|
||||
if engine != nil && engine.HasHandlers(pipeline.PostParse) {
|
||||
pctx := &pipeline.Context{
|
||||
Command: tool.CanonicalPath,
|
||||
Params: params,
|
||||
Schema: tool.InputSchema,
|
||||
}
|
||||
if pipeErr := engine.RunPhase(pipeline.PostParse, pctx); pipeErr != nil {
|
||||
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 {
|
||||
return err
|
||||
}
|
||||
if !dryRun {
|
||||
if err := confirmSensitiveTool(cmd, tool); err != nil {
|
||||
if err := confirmSensitiveTool(cmd, tool, guard); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// PreRequest: last chance to inspect/mutate payload before
|
||||
// the JSON-RPC call is dispatched.
|
||||
if engine != nil && engine.HasHandlers(pipeline.PreRequest) {
|
||||
pctx := &pipeline.Context{
|
||||
Command: tool.CanonicalPath,
|
||||
Params: params,
|
||||
Schema: tool.InputSchema,
|
||||
Payload: params,
|
||||
}
|
||||
if pipeErr := engine.RunPhase(pipeline.PreRequest, pctx); pipeErr != nil {
|
||||
return pipeErr
|
||||
}
|
||||
params = pctx.Params
|
||||
slog.Debug("pipeline pre-request",
|
||||
"command", tool.CanonicalPath,
|
||||
"param_count", len(params),
|
||||
)
|
||||
}
|
||||
|
||||
invocation := executor.NewInvocation(product, tool, params)
|
||||
invocation.DryRun = dryRun
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// PostResponse: transform or enrich the response before
|
||||
// writing it to stdout.
|
||||
if engine != nil && engine.HasHandlers(pipeline.PostResponse) {
|
||||
pctx := &pipeline.Context{
|
||||
Command: tool.CanonicalPath,
|
||||
Params: params,
|
||||
Schema: tool.InputSchema,
|
||||
Response: result.Response,
|
||||
}
|
||||
if pipeErr := engine.RunPhase(pipeline.PostResponse, pctx); pipeErr != nil {
|
||||
return pipeErr
|
||||
}
|
||||
result.Response = pctx.Response
|
||||
slog.Debug("pipeline post-response",
|
||||
"command", tool.CanonicalPath,
|
||||
"has_response", result.Response != nil,
|
||||
)
|
||||
}
|
||||
|
||||
if warning := lifecycleWarning(product); warning != "" {
|
||||
if result.Response == nil {
|
||||
result.Response = map[string]any{}
|
||||
}
|
||||
result.Response["warning"] = warning
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), result)
|
||||
return output.WriteFiltered(
|
||||
cmd.OutOrStdout(),
|
||||
output.ResolveFormat(cmd, output.FormatJSON),
|
||||
result,
|
||||
output.ResolveFields(cmd),
|
||||
output.ResolveJQ(cmd),
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
@@ -399,7 +547,7 @@ func applyFlagSpecs(cmd *cobra.Command, specs []FlagSpec) {
|
||||
}
|
||||
}
|
||||
|
||||
func collectOverrides(cmd *cobra.Command, specs []FlagSpec) (map[string]any, error) {
|
||||
func collectOverrides(cmd *cobra.Command, specs []FlagSpec, guard *StdinGuard) (map[string]any, error) {
|
||||
overrides := make(map[string]any)
|
||||
for _, spec := range specs {
|
||||
flagName := strings.TrimSpace(spec.FlagName)
|
||||
@@ -417,7 +565,12 @@ func collectOverrides(cmd *cobra.Command, specs []FlagSpec) (map[string]any, err
|
||||
if err != nil {
|
||||
return nil, apperrors.NewInternal(fmt.Sprintf("failed to read --%s", flagName))
|
||||
}
|
||||
overrides[spec.PropertyName] = value
|
||||
// Resolve @file / @- for all string-typed flags.
|
||||
resolved, resolveErr := ResolveInputSource(value, flagName, guard)
|
||||
if resolveErr != nil {
|
||||
return nil, resolveErr
|
||||
}
|
||||
overrides[spec.PropertyName] = resolved
|
||||
case flagJSON:
|
||||
value, err := cmd.Flags().GetString(flagName)
|
||||
if err != nil {
|
||||
@@ -491,10 +644,23 @@ func collectOverrides(cmd *cobra.Command, specs []FlagSpec) (map[string]any, err
|
||||
|
||||
func schemaPayload(catalog ir.Catalog, args []string) (map[string]any, error) {
|
||||
if len(args) == 0 {
|
||||
products := make([]map[string]any, 0, len(catalog.Products))
|
||||
for _, p := range catalog.Products {
|
||||
tools := make([]map[string]any, 0, len(p.Tools))
|
||||
for _, t := range p.Tools {
|
||||
tools = append(tools, compactTool(t))
|
||||
}
|
||||
products = append(products, map[string]any{
|
||||
"id": p.ID,
|
||||
"name": p.DisplayName,
|
||||
"description": p.Description,
|
||||
"tools": tools,
|
||||
})
|
||||
}
|
||||
return map[string]any{
|
||||
"kind": "schema",
|
||||
"products": catalog.Products,
|
||||
"count": len(catalog.Products),
|
||||
"count": len(products),
|
||||
"products": products,
|
||||
}, nil
|
||||
}
|
||||
|
||||
@@ -503,24 +669,35 @@ func schemaPayload(catalog ir.Catalog, args []string) (map[string]any, error) {
|
||||
return nil, apperrors.NewValidation(fmt.Sprintf("unknown canonical schema path %q", args[0]))
|
||||
}
|
||||
return map[string]any{
|
||||
"kind": "schema",
|
||||
"path": args[0],
|
||||
"product": product,
|
||||
"tool": tool,
|
||||
"required": requiredFields(tool.InputSchema),
|
||||
"kind": "schema",
|
||||
"path": args[0],
|
||||
"product": map[string]any{"id": product.ID, "name": product.DisplayName},
|
||||
"tool": compactTool(tool),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func writeSchemaText(w io.Writer, payload map[string]any) error {
|
||||
if path, ok := payload["path"].(string); ok && path != "" {
|
||||
_, err := fmt.Fprintf(w, "schema for %s\n", path)
|
||||
return err
|
||||
// compactTool returns a lean representation of a tool for schema
|
||||
// output, keeping only the fields that AI agents and developers
|
||||
// need: name, description, parameters, and sensitivity flag.
|
||||
func compactTool(t ir.ToolDescriptor) map[string]any {
|
||||
tool := map[string]any{
|
||||
"name": t.RPCName,
|
||||
"title": t.Title,
|
||||
"description": t.Description,
|
||||
"sensitive": t.Sensitive,
|
||||
}
|
||||
_, err := fmt.Fprintln(w, "canonical schema catalog")
|
||||
return err
|
||||
|
||||
if props, ok := t.InputSchema["properties"]; ok {
|
||||
tool["parameters"] = props
|
||||
}
|
||||
if req := requiredFields(t.InputSchema); len(req) > 0 {
|
||||
tool["required"] = req
|
||||
}
|
||||
|
||||
return tool
|
||||
}
|
||||
|
||||
func confirmSensitiveTool(cmd *cobra.Command, tool ir.ToolDescriptor) error {
|
||||
func confirmSensitiveTool(cmd *cobra.Command, tool ir.ToolDescriptor, guard *StdinGuard) error {
|
||||
if !tool.Sensitive {
|
||||
return nil
|
||||
}
|
||||
@@ -537,6 +714,13 @@ func confirmSensitiveTool(cmd *cobra.Command, tool ir.ToolDescriptor) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stdin was consumed for data input — interactive confirmation is impossible.
|
||||
if guard != nil && guard.Claimed() {
|
||||
return apperrors.NewValidation(
|
||||
"stdin used for data input; pass --yes to confirm sensitive operation",
|
||||
)
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintf(cmd.ErrOrStderr(), "tool %s is sensitive, continue? [y/N]: ", tool.CanonicalPath)
|
||||
confirmed, err := readYesNo(cmd.InOrStdin())
|
||||
if err != nil {
|
||||
|
||||
+607
-11
@@ -25,6 +25,7 @@ import (
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/ir"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func TestBuildFlagSpecsGeneratesOnlySupportedTopLevelFlags(t *testing.T) {
|
||||
@@ -124,7 +125,7 @@ func TestNewMCPCommandReturnsLoaderErrorForInvocations(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
wantErr := errors.New("fixture missing")
|
||||
cmd := NewMCPCommand(context.Background(), errorLoader{err: wantErr}, executor.EchoRunner{})
|
||||
cmd := NewMCPCommand(context.Background(), errorLoader{err: wantErr}, executor.EchoRunner{}, nil)
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
@@ -153,7 +154,7 @@ func TestNewMCPCommandSkipsProductsMarkedSkip(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, executor.EchoRunner{})
|
||||
}, executor.EchoRunner{}, nil)
|
||||
|
||||
if got := cmd.Commands(); len(got) != 1 || got[0].Name() != "drive" {
|
||||
t.Fatalf("mcp commands = %#v, want only drive", got)
|
||||
@@ -178,7 +179,7 @@ func TestProductCommandUsesCLICommandAlias(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
}, runner, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -212,7 +213,7 @@ func TestNewMCPCommandAddsGroupedRoutesFromCLIMetadata(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
}, runner, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -257,7 +258,7 @@ func TestToolCommandUsesCLINameAndFlagHints(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
}, runner, nil)
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
@@ -300,7 +301,7 @@ func TestToolCommandValidatesInputSchemaBeforeRun(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
}, runner, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -340,7 +341,7 @@ func TestToolCommandSupportsDryRunWithoutSensitiveConfirmation(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
}, runner, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -380,7 +381,7 @@ func TestDeprecatedLifecycleAddsWarningToResult(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, executor.EchoRunner{})
|
||||
}, executor.EchoRunner{}, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
var errOut bytes.Buffer
|
||||
@@ -424,7 +425,7 @@ func TestDeprecatedLifecyclePrintsWarningToStderr(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, executor.EchoRunner{})
|
||||
}, executor.EchoRunner{}, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
var errOut bytes.Buffer
|
||||
@@ -462,7 +463,7 @@ func TestSensitiveToolConfirmationWorksWithoutYesFlag(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, executor.EchoRunner{})
|
||||
}, executor.EchoRunner{}, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -491,7 +492,7 @@ func TestLegacyCandidateLifecycleAddsWarningToResult(t *testing.T) {
|
||||
},
|
||||
},
|
||||
},
|
||||
}, executor.EchoRunner{})
|
||||
}, executor.EchoRunner{}, nil)
|
||||
|
||||
var out bytes.Buffer
|
||||
var errOut bytes.Buffer
|
||||
@@ -514,6 +515,601 @@ func TestLegacyCandidateLifecycleAddsWarningToResult(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Input source resolution: @file for string flags
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestToolCommandResolvesAtFileForStringFlag(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "msg.md")
|
||||
if err := os.WriteFile(filePath, []byte("Hello from file"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
"user_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--text", "@" + filePath, "--user-id", "u001"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["text"] != "Hello from file" {
|
||||
t.Errorf("params[text] = %q, want %q", runner.last.Params["text"], "Hello from file")
|
||||
}
|
||||
if runner.last.Params["user_id"] != "u001" {
|
||||
t.Errorf("params[user_id] = %q, want %q", runner.last.Params["user_id"], "u001")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandResolvesAtFileForJsonFlag(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "payload.json")
|
||||
payload := `{"text":"from json file","user_id":"u002"}`
|
||||
if err := os.WriteFile(filePath, []byte(payload), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
"user_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--json", "@" + filePath})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["text"] != "from json file" {
|
||||
t.Errorf("params[text] = %q, want %q", runner.last.Params["text"], "from json file")
|
||||
}
|
||||
if runner.last.Params["user_id"] != "u002" {
|
||||
t.Errorf("params[user_id] = %q, want %q", runner.last.Params["user_id"], "u002")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandMultipleAtFileFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
titlePath := filepath.Join(dir, "title.txt")
|
||||
bodyPath := filepath.Join(dir, "body.md")
|
||||
if err := os.WriteFile(titlePath, []byte("My Title"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(bodyPath, []byte("# Body\n\nContent here"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "doc",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "create_document",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"title": map[string]any{"type": "string"},
|
||||
"body": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"doc", "create_document", "--title", "@" + titlePath, "--body", "@" + bodyPath})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["title"] != "My Title" {
|
||||
t.Errorf("params[title] = %q, want %q", runner.last.Params["title"], "My Title")
|
||||
}
|
||||
if runner.last.Params["body"] != "# Body\n\nContent here" {
|
||||
t.Errorf("params[body] = %q, want %q", runner.last.Params["body"], "# Body\n\nContent here")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandAtFileMissingReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--text", "@/nonexistent/file.txt"})
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() should fail for missing @file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--text") {
|
||||
t.Errorf("error should mention flag name, got: %v", err)
|
||||
}
|
||||
if runner.called != 0 {
|
||||
t.Error("runner should not be called on @file error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandAtFileForJsonMissingReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--json", "@/nonexistent/payload.json"})
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() should fail for missing @file on --json")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--json") {
|
||||
t.Errorf("error should mention --json, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandAtFileUTF8ContentPreserved(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "chinese.txt")
|
||||
content := "你好世界 🌍\n第二行"
|
||||
if err := os.WriteFile(filePath, []byte(content), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--text", "@" + filePath})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["text"] != content {
|
||||
t.Errorf("params[text] = %q, want %q", runner.last.Params["text"], content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolCommandPlainAtValueNotResolvedForNonStringFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Integer and boolean flags should NOT resolve @file syntax.
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "todo",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "create_task",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"title": map[string]any{"type": "string"},
|
||||
"priority": map[string]any{"type": "integer"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"todo", "create_task", "--title", "test", "--priority", "3"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["priority"] != 3 {
|
||||
t.Errorf("params[priority] = %v, want 3", runner.last.Params["priority"])
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Input source resolution: --json @file override priority
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestToolCommandJsonFlagOverridesOverrideFlags(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "base.json")
|
||||
if err := os.WriteFile(filePath, []byte(`{"text":"from-json","user_id":"json-user"}`), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "chat",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "send_message",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"text": map[string]any{"type": "string"},
|
||||
"user_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
// --text override should win over --json base payload.
|
||||
cmd.SetArgs([]string{"chat", "send_message", "--json", "@" + filePath, "--text", "override"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["text"] != "override" {
|
||||
t.Errorf("params[text] = %q, want %q (override should win)", runner.last.Params["text"], "override")
|
||||
}
|
||||
if runner.last.Params["user_id"] != "json-user" {
|
||||
t.Errorf("params[user_id] = %q, want %q (from json base)", runner.last.Params["user_id"], "json-user")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Sensitive tool + stdin guard interaction
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSensitiveToolWithStdinClaimedRequiresYes(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "msg.txt")
|
||||
if err := os.WriteFile(filePath, []byte("content"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
// Sensitive tool + @file (does NOT claim stdin) → should still prompt.
|
||||
// We provide "yes" on stdin to pass confirmation.
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "doc",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "delete_document",
|
||||
Sensitive: true,
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"doc_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetIn(strings.NewReader("yes\n"))
|
||||
cmd.SetArgs([]string{"doc", "delete_document", "--doc-id", "DOC001"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.called != 1 {
|
||||
t.Errorf("runner called = %d, want 1", runner.called)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSensitiveToolDeniedOnStdinWithNoYesFlag(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "doc",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "delete_document",
|
||||
Sensitive: true,
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"doc_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetIn(strings.NewReader("no\n"))
|
||||
cmd.SetArgs([]string{"doc", "delete_document", "--doc-id", "DOC001"})
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() should fail when user denies confirmation")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "cancelled") {
|
||||
t.Errorf("error should mention cancellation, got: %v", err)
|
||||
}
|
||||
if runner.called != 0 {
|
||||
t.Error("runner should not be called when confirmation denied")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSensitiveToolWithYesFlagSkipsConfirmation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "doc",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "delete_document",
|
||||
Sensitive: true,
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"doc_id": map[string]any{"type": "string"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.PersistentFlags().Bool("yes", false, "Skip confirmation")
|
||||
cmd.SetArgs([]string{"doc", "delete_document", "--doc-id", "DOC001", "--yes"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.called != 1 {
|
||||
t.Errorf("runner called = %d, want 1", runner.called)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// collectOverrides: @file does not affect non-string flag types
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCollectOverridesResolvesAtFileOnlyForStringKind(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "name.txt")
|
||||
if err := os.WriteFile(filePath, []byte("resolved name"), 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newTestMCPCommand(t, ir.Catalog{
|
||||
Products: []ir.CanonicalProduct{
|
||||
{
|
||||
ID: "contact",
|
||||
Tools: []ir.ToolDescriptor{
|
||||
{
|
||||
RPCName: "search_user",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"keyword": map[string]any{"type": "string"},
|
||||
"active": map[string]any{"type": "boolean"},
|
||||
"limit": map[string]any{"type": "integer"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}, runner)
|
||||
|
||||
cmd.SetArgs([]string{"contact", "search_user",
|
||||
"--keyword", "@" + filePath,
|
||||
"--active=true",
|
||||
"--limit", "10",
|
||||
})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
if runner.last.Params["keyword"] != "resolved name" {
|
||||
t.Errorf("params[keyword] = %q, want %q", runner.last.Params["keyword"], "resolved name")
|
||||
}
|
||||
if runner.last.Params["active"] != true {
|
||||
t.Errorf("params[active] = %v, want true", runner.last.Params["active"])
|
||||
}
|
||||
if runner.last.Params["limit"] != 10 {
|
||||
t.Errorf("params[limit] = %v, want 10", runner.last.Params["limit"])
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Test helper
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func newTestMCPCommand(t *testing.T, catalog ir.Catalog, runner executor.Runner) *cobra.Command {
|
||||
t.Helper()
|
||||
cmd := NewMCPCommand(context.Background(), StaticLoader{Catalog: catalog}, runner, nil)
|
||||
var out, errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
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
|
||||
}
|
||||
|
||||
+126
-9
@@ -17,25 +17,107 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/ir"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CACHE_DIR",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "覆盖缓存目录",
|
||||
DefaultValue: "~/.dws/cache",
|
||||
Example: "/tmp/dws-cache",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_CATALOG_FIXTURE",
|
||||
Category: configmeta.CategoryDebug,
|
||||
Description: "使用本地 JSON 文件替代在线目录发现",
|
||||
Example: "/path/to/catalog.json",
|
||||
Hidden: true,
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_PLUGIN_COLD_TIMEOUT",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "插件 MCP 冷启动发现的超时时长(Go duration 格式,如 2s / 1500ms)。设置后同时覆盖 HTTP 与 stdio 插件的冷启动预算;未设置时使用内置默认值(HTTP 无鉴权 1s / 有鉴权 1.5s / stdio 2s)。",
|
||||
DefaultValue: "",
|
||||
Example: "3s",
|
||||
})
|
||||
}
|
||||
|
||||
// CatalogDegradedReason identifies why catalog discovery returned empty.
|
||||
type CatalogDegradedReason string
|
||||
|
||||
const (
|
||||
DegradedUnauthenticated CatalogDegradedReason = "unauthenticated"
|
||||
DegradedMarketUnreachable CatalogDegradedReason = "market_unreachable"
|
||||
DegradedRuntimeAllFailed CatalogDegradedReason = "runtime_all_failed"
|
||||
)
|
||||
|
||||
// CatalogDegraded is returned by EnvironmentLoader.Load when discovery
|
||||
// fails for a diagnosable reason. Callers that need graceful degradation
|
||||
// (e.g. the runtime runner) can check errors.As and fall back to an
|
||||
// empty catalog; callers like the schema command can surface the hint.
|
||||
type CatalogDegraded struct {
|
||||
Reason CatalogDegradedReason
|
||||
Hint string
|
||||
ServerCount int // number of servers discovered (only set for runtime_all_failed)
|
||||
}
|
||||
|
||||
func (e *CatalogDegraded) Error() string { return string(e.Reason) + ": " + e.Hint }
|
||||
|
||||
func degradedHint(reason CatalogDegradedReason, serverCount int) string {
|
||||
embedded := edition.Get().IsEmbedded
|
||||
switch reason {
|
||||
case DegradedUnauthenticated:
|
||||
if embedded {
|
||||
return "未登录,请重新认证"
|
||||
}
|
||||
return "未登录,无法发现 MCP 服务。请先执行: dws auth login"
|
||||
case DegradedMarketUnreachable:
|
||||
if embedded {
|
||||
return "无法连接 MCP 市场,请检查网络"
|
||||
}
|
||||
return "无法连接 MCP 市场 (mcp.dingtalk.com),请检查网络"
|
||||
case DegradedRuntimeAllFailed:
|
||||
if embedded {
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试", serverCount)
|
||||
}
|
||||
return fmt.Sprintf("已发现 %d 个服务但连接全部失败,请稍后重试或执行: dws cache refresh", serverCount)
|
||||
default:
|
||||
return "MCP 服务发现失败"
|
||||
}
|
||||
}
|
||||
|
||||
func newCatalogDegraded(reason CatalogDegradedReason, serverCount int) *CatalogDegraded {
|
||||
return &CatalogDegraded{
|
||||
Reason: reason,
|
||||
Hint: degradedHint(reason, serverCount),
|
||||
ServerCount: serverCount,
|
||||
}
|
||||
}
|
||||
|
||||
const (
|
||||
CatalogFixtureEnv = "DWS_CATALOG_FIXTURE"
|
||||
CacheDirEnv = "DWS_CACHE_DIR"
|
||||
PluginColdTimeoutEnv = "DWS_PLUGIN_COLD_TIMEOUT"
|
||||
DefaultMarketBaseURL = "https://mcp.dingtalk.com"
|
||||
|
||||
// defaultDiscoveryTimeout bounds the time spent on live registry discovery.
|
||||
defaultDiscoveryTimeout = 10 * time.Second
|
||||
// Tightened to 4s so a slow/unreachable discovery endpoint cannot block
|
||||
// every CLI command invocation. See issue #119.
|
||||
defaultDiscoveryTimeout = 4 * time.Second
|
||||
)
|
||||
|
||||
type CatalogLoader interface {
|
||||
@@ -50,6 +132,22 @@ func (l StaticLoader) Load(_ context.Context) (ir.Catalog, error) {
|
||||
return l.Catalog, nil
|
||||
}
|
||||
|
||||
// CatalogLoaderFrom creates a CatalogLoader that returns a
|
||||
// pre-loaded catalog and error. This allows multiple consumers
|
||||
// (schema command, MCP command tree) to share one discovery result.
|
||||
func CatalogLoaderFrom(catalog ir.Catalog, err error) CatalogLoader {
|
||||
return &preloadedLoader{catalog: catalog, err: err}
|
||||
}
|
||||
|
||||
type preloadedLoader struct {
|
||||
catalog ir.Catalog
|
||||
err error
|
||||
}
|
||||
|
||||
func (l *preloadedLoader) Load(_ context.Context) (ir.Catalog, error) {
|
||||
return l.catalog, l.err
|
||||
}
|
||||
|
||||
type FixtureLoader struct {
|
||||
Path string
|
||||
}
|
||||
@@ -73,6 +171,13 @@ type EnvironmentLoader struct {
|
||||
// DiscoveryTimeout overrides the default timeout for live registry discovery.
|
||||
// Zero means use defaultDiscoveryTimeout.
|
||||
DiscoveryTimeout time.Duration
|
||||
// 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 {
|
||||
@@ -104,11 +209,22 @@ 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
|
||||
@@ -123,14 +239,15 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
transportClient,
|
||||
store,
|
||||
)
|
||||
if l.LoggerFunc != nil {
|
||||
service.Logger = l.LoggerFunc()
|
||||
}
|
||||
response, err := service.MarketClient.FetchServers(discoverCtx, 200)
|
||||
if err != nil {
|
||||
// Graceful degradation: return empty catalog on discovery failure.
|
||||
// The runtime runner will fall back to EchoRunner for unknown products.
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedMarketUnreachable, 0)
|
||||
}
|
||||
|
||||
servers := market.NormalizeServers(response, "live_market")
|
||||
@@ -160,10 +277,10 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
|
||||
refreshed, failures := service.DiscoverAllRuntime(discoverCtx, toRefresh)
|
||||
if len(unchangedRuntime) == 0 && len(refreshed) == 0 && len(failures) > 0 {
|
||||
if cached.Available {
|
||||
if cached.Available && len(cached.Catalog.Products) > 0 {
|
||||
return cached.Catalog, nil
|
||||
}
|
||||
return ir.Catalog{}, nil
|
||||
return ir.Catalog{}, newCatalogDegraded(DegradedRuntimeAllFailed, len(servers))
|
||||
}
|
||||
|
||||
refreshedByKey := make(map[string]discovery.RuntimeServer, len(refreshed))
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
// maxStdinSize limits the amount of data read from stdin or @file
|
||||
// to prevent memory exhaustion from accidental large pipes.
|
||||
maxStdinSize = 10 * 1024 * 1024 // 10 MB
|
||||
)
|
||||
|
||||
// StdinGuard ensures stdin is consumed at most once per command invocation.
|
||||
// Multiple flags using @- or implicit stdin fallback would race on the same
|
||||
// reader; StdinGuard detects and rejects the second claim with a clear error.
|
||||
type StdinGuard struct {
|
||||
mu sync.Mutex
|
||||
claimed bool
|
||||
claimBy string
|
||||
}
|
||||
|
||||
// NewStdinGuard creates a fresh guard for one command invocation.
|
||||
func NewStdinGuard() *StdinGuard {
|
||||
return &StdinGuard{}
|
||||
}
|
||||
|
||||
// Claim marks stdin as consumed by the named source (e.g. "--text @-").
|
||||
// Returns an error if stdin was already claimed.
|
||||
func (g *StdinGuard) Claim(source string) error {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
if g.claimed {
|
||||
return apperrors.NewValidation(fmt.Sprintf(
|
||||
"stdin already consumed by %s; cannot also read stdin for %s",
|
||||
g.claimBy, source,
|
||||
))
|
||||
}
|
||||
g.claimed = true
|
||||
g.claimBy = source
|
||||
return nil
|
||||
}
|
||||
|
||||
// Claimed reports whether stdin has been consumed.
|
||||
func (g *StdinGuard) Claimed() bool {
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
return g.claimed
|
||||
}
|
||||
|
||||
// StdinIsPipe reports whether stdin is a pipe (not a terminal).
|
||||
// This is a non-consuming check — it only inspects file mode via stat.
|
||||
func StdinIsPipe() bool {
|
||||
info, err := os.Stdin.Stat()
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return info.Mode()&os.ModeCharDevice == 0
|
||||
}
|
||||
|
||||
// ReadStdinIfPiped reads all data from stdin if it is a pipe (not a terminal).
|
||||
// Returns empty string if stdin is a terminal or has no data.
|
||||
func ReadStdinIfPiped() (string, error) {
|
||||
if !StdinIsPipe() {
|
||||
return "", nil
|
||||
}
|
||||
return readStdinBounded()
|
||||
}
|
||||
|
||||
// ReadStdin reads all data from stdin unconditionally (up to maxStdinSize).
|
||||
// Use this when the caller has explicitly requested stdin via @-.
|
||||
func ReadStdin() (string, error) {
|
||||
return readStdinBounded()
|
||||
}
|
||||
|
||||
// readFileBounded opens a file and reads up to maxStdinSize bytes.
|
||||
// Uses io.LimitReader to avoid TOCTOU between stat and read.
|
||||
func readFileBounded(path string) ([]byte, error) {
|
||||
f, err := os.Open(path)
|
||||
if err != nil {
|
||||
return nil, apperrors.NewValidation("@file: " + err.Error())
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
data, err := io.ReadAll(io.LimitReader(f, maxStdinSize+1))
|
||||
if err != nil {
|
||||
return nil, apperrors.NewValidation("@file: " + err.Error())
|
||||
}
|
||||
if int64(len(data)) > maxStdinSize {
|
||||
return nil, apperrors.NewValidation("@file: file exceeds 10 MB limit")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// readStdinBounded reads from os.Stdin up to maxStdinSize bytes.
|
||||
func readStdinBounded() (string, error) {
|
||||
data, err := io.ReadAll(io.LimitReader(os.Stdin, maxStdinSize+1))
|
||||
if err != nil {
|
||||
return "", apperrors.NewValidation("failed to read stdin: " + err.Error())
|
||||
}
|
||||
if int64(len(data)) > maxStdinSize {
|
||||
return "", apperrors.NewValidation("stdin input exceeds 10 MB limit")
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
// ReadFileArg reads the contents of a file referenced by the @filename syntax.
|
||||
// Returns the original value unchanged if it does not start with "@".
|
||||
// Returns an error if the file cannot be read or exceeds the size limit.
|
||||
//
|
||||
// Note: @- (stdin) is NOT handled here; use ResolveInputSource instead.
|
||||
func ReadFileArg(value string) (string, bool, error) {
|
||||
if !strings.HasPrefix(value, "@") {
|
||||
return value, false, nil
|
||||
}
|
||||
path := value[1:]
|
||||
if path == "" {
|
||||
return "", false, apperrors.NewValidation("@file: filename must not be empty")
|
||||
}
|
||||
// @- is stdin, not a file — callers should use ResolveInputSource.
|
||||
if path == "-" {
|
||||
return value, false, nil
|
||||
}
|
||||
|
||||
data, err := readFileBounded(path)
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
return string(data), true, nil
|
||||
}
|
||||
|
||||
// ResolveInputSource resolves a flag value that may reference an external
|
||||
// input source. It supports three forms:
|
||||
//
|
||||
// - "@-" reads from stdin (requires StdinGuard claim)
|
||||
// - "@<path>" reads from the named file
|
||||
// - anything else returned unchanged
|
||||
//
|
||||
// The flagName parameter is used only for error messages and StdinGuard tracking.
|
||||
func ResolveInputSource(value string, flagName string, guard *StdinGuard) (string, error) {
|
||||
if !strings.HasPrefix(value, "@") {
|
||||
return value, nil
|
||||
}
|
||||
|
||||
path := value[1:]
|
||||
if path == "" {
|
||||
return "", apperrors.NewValidation(fmt.Sprintf("--%s: @file filename must not be empty", flagName))
|
||||
}
|
||||
|
||||
// @- reads from stdin.
|
||||
if path == "-" {
|
||||
if guard == nil {
|
||||
return "", apperrors.NewValidation(fmt.Sprintf("--%s: stdin (@-) not available in this context", flagName))
|
||||
}
|
||||
if err := guard.Claim(fmt.Sprintf("--%s @-", flagName)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return ReadStdin()
|
||||
}
|
||||
|
||||
// @<path> reads from file.
|
||||
data, err := readFileBounded(path)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("--%s: %w", flagName, err)
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
@@ -0,0 +1,684 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package cli
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// writeTestFile is a test helper that writes data to a file and fails the
|
||||
// test immediately if the write fails, preventing confusing downstream errors.
|
||||
func writeTestFile(t *testing.T, path string, data []byte) {
|
||||
t.Helper()
|
||||
if err := os.WriteFile(path, data, 0o644); err != nil {
|
||||
t.Fatalf("failed to write test file %s: %v", path, err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ReadFileArg
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestReadFileArgPlainValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
val, isFile, err := ReadFileArg("hello world")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if isFile {
|
||||
t.Error("plain value should not be detected as file")
|
||||
}
|
||||
if val != "hello world" {
|
||||
t.Errorf("got %q, want %q", val, "hello world")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileArgReadsFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "input.txt")
|
||||
writeTestFile(t, path, []byte(`{"title":"test"}`))
|
||||
|
||||
val, isFile, err := ReadFileArg("@" + path)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !isFile {
|
||||
t.Error("@file value should be detected as file")
|
||||
}
|
||||
if val != `{"title":"test"}` {
|
||||
t.Errorf("got %q, want %q", val, `{"title":"test"}`)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileArgEmptyFilename(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, _, err := ReadFileArg("@")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty filename")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "must not be empty") {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileArgMissingFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, _, err := ReadFileArg("@/nonexistent/path/file.txt")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileArgSizeLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "huge.txt")
|
||||
// Create a file slightly over maxStdinSize (write 10MB + 1 byte)
|
||||
f, err := os.Create(path)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create test file: %v", err)
|
||||
}
|
||||
data := strings.Repeat("x", maxStdinSize+1)
|
||||
if _, err := f.WriteString(data); err != nil {
|
||||
f.Close()
|
||||
t.Fatalf("failed to write test data: %v", err)
|
||||
}
|
||||
f.Close()
|
||||
|
||||
_, _, err = ReadFileArg("@" + path)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for oversized file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "10 MB") {
|
||||
t.Errorf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileArgAtDashPassThrough(t *testing.T) {
|
||||
t.Parallel()
|
||||
// @- means stdin — ReadFileArg should NOT handle it, just pass through.
|
||||
val, isFile, err := ReadFileArg("@-")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if isFile {
|
||||
t.Error("@- should not be treated as a file by ReadFileArg")
|
||||
}
|
||||
if val != "@-" {
|
||||
t.Errorf("got %q, want %q", val, "@-")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadStdinIfPipedReturnsEmptyForTerminal(t *testing.T) {
|
||||
// This test runs in a terminal context (go test), so stdin is a terminal.
|
||||
// ReadStdinIfPiped should return empty string.
|
||||
val, err := ReadStdinIfPiped()
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "" {
|
||||
t.Errorf("expected empty string for terminal stdin, got %q", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdinIsPipeReturnsFalseForTerminal(t *testing.T) {
|
||||
// go test runs with stdin as a terminal, so StdinIsPipe should return false.
|
||||
if StdinIsPipe() {
|
||||
t.Error("expected StdinIsPipe() == false in terminal context")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// readFileBounded (via ReadFileArg)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestReadFileBoundedEmptyFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "empty.txt")
|
||||
writeTestFile(t, path, []byte(""))
|
||||
|
||||
val, isFile, err := ReadFileArg("@" + path)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !isFile {
|
||||
t.Error("should be detected as file")
|
||||
}
|
||||
if val != "" {
|
||||
t.Errorf("expected empty content, got %q", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadFileBoundedPermissionDenied(t *testing.T) {
|
||||
t.Parallel()
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "noperm.txt")
|
||||
writeTestFile(t, path, []byte("secret"))
|
||||
os.Chmod(path, 0o000)
|
||||
t.Cleanup(func() { os.Chmod(path, 0o644) })
|
||||
|
||||
_, _, err := ReadFileArg("@" + path)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for permission denied")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// StdinGuard
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStdinGuardFirstClaimSucceeds(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
if err := guard.Claim("--text @-"); err != nil {
|
||||
t.Fatalf("first claim should succeed: %v", err)
|
||||
}
|
||||
if !guard.Claimed() {
|
||||
t.Error("guard should report claimed after successful claim")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdinGuardSecondClaimFails(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
_ = guard.Claim("--text @-")
|
||||
|
||||
err := guard.Claim("--body @-")
|
||||
if err == nil {
|
||||
t.Fatal("second claim should fail")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--text @-") {
|
||||
t.Errorf("error should mention first claimer, got: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--body @-") {
|
||||
t.Errorf("error should mention second claimer, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdinGuardNotClaimedInitially(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
if guard.Claimed() {
|
||||
t.Error("fresh guard should not be claimed")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ResolveInputSource
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestResolveInputSourcePlainValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("hello", "text", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "hello" {
|
||||
t.Errorf("got %q, want %q", val, "hello")
|
||||
}
|
||||
if guard.Claimed() {
|
||||
t.Error("plain value should not claim stdin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceEmptyValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("", "text", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "" {
|
||||
t.Errorf("got %q, want empty", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileReadsFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "msg.txt")
|
||||
writeTestFile(t, path, []byte("file content here"))
|
||||
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("@"+path, "text", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "file content here" {
|
||||
t.Errorf("got %q, want %q", val, "file content here")
|
||||
}
|
||||
// @file should NOT claim stdin.
|
||||
if guard.Claimed() {
|
||||
t.Error("@file should not claim stdin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileEmptyName(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
_, err := ResolveInputSource("@", "json", guard)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty @file name")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--json") {
|
||||
t.Errorf("error should mention flag name, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileMissing(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
_, err := ResolveInputSource("@/no/such/file.txt", "data", guard)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--data") {
|
||||
t.Errorf("error should mention flag name, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileSizeLimit(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "huge.bin")
|
||||
f, err := os.Create(path)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create test file: %v", err)
|
||||
}
|
||||
if _, err := f.WriteString(strings.Repeat("x", maxStdinSize+1)); err != nil {
|
||||
f.Close()
|
||||
t.Fatalf("failed to write test data: %v", err)
|
||||
}
|
||||
f.Close()
|
||||
|
||||
guard := NewStdinGuard()
|
||||
_, resolveErr := ResolveInputSource("@"+path, "body", guard)
|
||||
if resolveErr == nil {
|
||||
t.Fatal("expected error for oversized file")
|
||||
}
|
||||
if !strings.Contains(resolveErr.Error(), "10 MB") {
|
||||
t.Errorf("error should mention size limit, got: %v", resolveErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtDashNilGuardFails(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := ResolveInputSource("@-", "text", nil)
|
||||
if err == nil {
|
||||
t.Fatal("expected error when guard is nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "not available") {
|
||||
t.Errorf("error should mention stdin unavailability, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtDashDoubleClaimFails(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
// Simulate first claim from another source.
|
||||
_ = guard.Claim("--json @-")
|
||||
|
||||
_, err := ResolveInputSource("@-", "text", guard)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for double stdin claim")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "already consumed") {
|
||||
t.Errorf("error should mention stdin conflict, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileMultiline(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "multiline.md")
|
||||
content := "# Title\n\nLine 1\nLine 2\nLine 3\n"
|
||||
writeTestFile(t, path, []byte(content))
|
||||
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("@"+path, "text", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != content {
|
||||
t.Errorf("multiline content mismatch:\ngot: %q\nwant: %q", val, content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileUTF8(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "chinese.txt")
|
||||
content := "你好世界 🌍"
|
||||
writeTestFile(t, path, []byte(content))
|
||||
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("@"+path, "text", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != content {
|
||||
t.Errorf("UTF-8 content mismatch:\ngot: %q\nwant: %q", val, content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceAtFileJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "payload.json")
|
||||
content := `{"title":"meeting","startTime":"2026-03-29T10:00:00Z"}`
|
||||
writeTestFile(t, path, []byte(content))
|
||||
|
||||
guard := NewStdinGuard()
|
||||
val, err := ResolveInputSource("@"+path, "json", guard)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != content {
|
||||
t.Errorf("JSON content mismatch:\ngot: %q\nwant: %q", val, content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveInputSourceValueStartingWithAtSign(t *testing.T) {
|
||||
t.Parallel()
|
||||
// A value like "@mention" that looks like @file but the file doesn't exist
|
||||
// should return an error (user likely intended file input).
|
||||
guard := NewStdinGuard()
|
||||
_, err := ResolveInputSource("@mention_someone", "text", guard)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for non-existent @file path")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ResolveInputSource: table-driven edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestResolveInputSourceTable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
existingFile := filepath.Join(dir, "exists.txt")
|
||||
writeTestFile(t, existingFile, []byte("file-data"))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value string
|
||||
flagName string
|
||||
guard *StdinGuard
|
||||
wantVal string
|
||||
wantErr string // substring to match in error, "" means no error
|
||||
wantClaim bool // expect guard to be claimed after call
|
||||
}{
|
||||
{
|
||||
name: "plain string unchanged",
|
||||
value: "hello",
|
||||
flagName: "text",
|
||||
guard: NewStdinGuard(),
|
||||
wantVal: "hello",
|
||||
},
|
||||
{
|
||||
name: "empty string unchanged",
|
||||
value: "",
|
||||
flagName: "text",
|
||||
guard: NewStdinGuard(),
|
||||
wantVal: "",
|
||||
},
|
||||
{
|
||||
name: "plain string with special chars",
|
||||
value: "hello@world.com",
|
||||
flagName: "email",
|
||||
guard: NewStdinGuard(),
|
||||
wantVal: "hello@world.com",
|
||||
},
|
||||
{
|
||||
name: "@file reads content",
|
||||
value: "@" + existingFile,
|
||||
flagName: "data",
|
||||
guard: NewStdinGuard(),
|
||||
wantVal: "file-data",
|
||||
},
|
||||
{
|
||||
name: "@file does not claim stdin",
|
||||
value: "@" + existingFile,
|
||||
flagName: "data",
|
||||
guard: NewStdinGuard(),
|
||||
wantVal: "file-data",
|
||||
wantClaim: false,
|
||||
},
|
||||
{
|
||||
name: "bare @ is error",
|
||||
value: "@",
|
||||
flagName: "body",
|
||||
guard: NewStdinGuard(),
|
||||
wantErr: "must not be empty",
|
||||
},
|
||||
{
|
||||
name: "missing file is error",
|
||||
value: "@/tmp/does-not-exist-" + t.Name(),
|
||||
flagName: "file",
|
||||
guard: NewStdinGuard(),
|
||||
wantErr: "--file",
|
||||
},
|
||||
{
|
||||
name: "@- with nil guard is error",
|
||||
value: "@-",
|
||||
flagName: "text",
|
||||
guard: nil,
|
||||
wantErr: "not available",
|
||||
},
|
||||
{
|
||||
name: "@- with pre-claimed guard is error",
|
||||
value: "@-",
|
||||
flagName: "body",
|
||||
guard: func() *StdinGuard {
|
||||
g := NewStdinGuard()
|
||||
_ = g.Claim("--json @-")
|
||||
return g
|
||||
}(),
|
||||
wantErr: "already consumed",
|
||||
},
|
||||
{
|
||||
name: "error message includes flag name",
|
||||
value: "@/nonexistent",
|
||||
flagName: "my-flag",
|
||||
guard: NewStdinGuard(),
|
||||
wantErr: "--my-flag",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
val, err := ResolveInputSource(tt.value, tt.flagName, tt.guard)
|
||||
if tt.wantErr != "" {
|
||||
if err == nil {
|
||||
t.Fatalf("expected error containing %q, got nil", tt.wantErr)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Errorf("error = %q, want substring %q", err.Error(), tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != tt.wantVal {
|
||||
t.Errorf("value = %q, want %q", val, tt.wantVal)
|
||||
}
|
||||
if tt.guard != nil && tt.wantClaim != tt.guard.Claimed() {
|
||||
t.Errorf("guard.Claimed() = %v, want %v", tt.guard.Claimed(), tt.wantClaim)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ReadFileArg: table-driven edge cases
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestReadFileArgTable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
existingFile := filepath.Join(dir, "data.txt")
|
||||
writeTestFile(t, existingFile, []byte("content"))
|
||||
emptyFile := filepath.Join(dir, "empty.txt")
|
||||
writeTestFile(t, emptyFile, []byte(""))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
value string
|
||||
wantVal string
|
||||
wantIsFile bool
|
||||
wantErr string
|
||||
}{
|
||||
{"plain string", "hello", "hello", false, ""},
|
||||
{"empty string", "", "", false, ""},
|
||||
{"email-like value", "user@domain.com", "user@domain.com", false, ""},
|
||||
{"@file reads content", "@" + existingFile, "content", true, ""},
|
||||
{"@file empty content", "@" + emptyFile, "", true, ""},
|
||||
{"@- passes through", "@-", "@-", false, ""},
|
||||
{"bare @ is error", "@", "", false, "must not be empty"},
|
||||
{"missing file is error", "@/nonexistent", "", false, "@file"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
val, isFile, err := ReadFileArg(tt.value)
|
||||
if tt.wantErr != "" {
|
||||
if err == nil {
|
||||
t.Fatalf("expected error containing %q, got nil", tt.wantErr)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tt.wantErr) {
|
||||
t.Errorf("error = %q, want substring %q", err.Error(), tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != tt.wantVal {
|
||||
t.Errorf("value = %q, want %q", val, tt.wantVal)
|
||||
}
|
||||
if isFile != tt.wantIsFile {
|
||||
t.Errorf("isFile = %v, want %v", isFile, tt.wantIsFile)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// StdinGuard: table-driven claim sequences
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStdinGuardClaimSequences(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
claims []string // sources to claim in order
|
||||
wantFails int // how many claims should fail
|
||||
wantErrSub string // substring expected in first failure
|
||||
}{
|
||||
{
|
||||
name: "single claim succeeds",
|
||||
claims: []string{"--json @-"},
|
||||
wantFails: 0,
|
||||
},
|
||||
{
|
||||
name: "second claim fails",
|
||||
claims: []string{"--json @-", "--text @-"},
|
||||
wantFails: 1,
|
||||
wantErrSub: "already consumed",
|
||||
},
|
||||
{
|
||||
name: "third claim also fails",
|
||||
claims: []string{"--json @-", "--text @-", "--body @-"},
|
||||
wantFails: 2,
|
||||
wantErrSub: "already consumed",
|
||||
},
|
||||
{
|
||||
name: "error names both sources",
|
||||
claims: []string{"implicit stdin (pipe)", "--text @-"},
|
||||
wantFails: 1,
|
||||
wantErrSub: "implicit stdin (pipe)",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
fails := 0
|
||||
var firstErr error
|
||||
for _, source := range tt.claims {
|
||||
if err := guard.Claim(source); err != nil {
|
||||
fails++
|
||||
if firstErr == nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
}
|
||||
if fails != tt.wantFails {
|
||||
t.Errorf("failures = %d, want %d", fails, tt.wantFails)
|
||||
}
|
||||
if tt.wantErrSub != "" && firstErr != nil {
|
||||
if !strings.Contains(firstErr.Error(), tt.wantErrSub) {
|
||||
t.Errorf("first error = %q, want substring %q", firstErr.Error(), tt.wantErrSub)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// StdinGuard concurrency safety
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestStdinGuardConcurrentClaims(t *testing.T) {
|
||||
t.Parallel()
|
||||
guard := NewStdinGuard()
|
||||
|
||||
const goroutines = 50
|
||||
results := make(chan error, goroutines)
|
||||
for i := 0; i < goroutines; i++ {
|
||||
go func(id int) {
|
||||
results <- guard.Claim("goroutine")
|
||||
}(i)
|
||||
}
|
||||
|
||||
successCount := 0
|
||||
for i := 0; i < goroutines; i++ {
|
||||
if err := <-results; err == nil {
|
||||
successCount++
|
||||
}
|
||||
}
|
||||
if successCount != 1 {
|
||||
t.Errorf("exactly one goroutine should succeed, got %d", successCount)
|
||||
}
|
||||
}
|
||||
@@ -35,7 +35,12 @@ import (
|
||||
//
|
||||
// Conversion rules reference: docs/mcp-to-cli-conversion.md
|
||||
func BuildDynamicCommands(servers []market.ServerDescriptor, runner executor.Runner, detailsByID map[string][]market.DetailTool) []*cobra.Command {
|
||||
var commands []*cobra.Command
|
||||
type builtCmd struct {
|
||||
cmd *cobra.Command
|
||||
parent string // cli.Parent: attach as sub-command of this top-level command
|
||||
}
|
||||
|
||||
var built []builtCmd
|
||||
for _, server := range servers {
|
||||
cli := server.CLI
|
||||
// §1.5: cli.skip → skip entire service
|
||||
@@ -90,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
|
||||
@@ -136,7 +145,41 @@ func BuildDynamicCommands(servers []market.ServerDescriptor, runner executor.Run
|
||||
}
|
||||
}
|
||||
|
||||
commands = append(commands, rootCmd)
|
||||
built = append(built, builtCmd{cmd: rootCmd, parent: strings.TrimSpace(cli.Parent)})
|
||||
}
|
||||
|
||||
// Collect top-level commands first, then attach child commands via cli.Parent.
|
||||
topLevel := make(map[string]*cobra.Command)
|
||||
var topOrder []string
|
||||
var children []builtCmd
|
||||
|
||||
for _, b := range built {
|
||||
if b.parent == "" {
|
||||
name := b.cmd.Name()
|
||||
if _, exists := topLevel[name]; !exists {
|
||||
topOrder = append(topOrder, name)
|
||||
}
|
||||
topLevel[name] = b.cmd
|
||||
} else {
|
||||
children = append(children, b)
|
||||
}
|
||||
}
|
||||
for _, child := range children {
|
||||
if parent, ok := topLevel[child.parent]; ok {
|
||||
parent.AddCommand(child.cmd)
|
||||
} else {
|
||||
// Parent not found among dynamic commands; emit as top-level.
|
||||
name := child.cmd.Name()
|
||||
if _, exists := topLevel[name]; !exists {
|
||||
topOrder = append(topOrder, name)
|
||||
}
|
||||
topLevel[name] = child.cmd
|
||||
}
|
||||
}
|
||||
|
||||
commands := make([]*cobra.Command, 0, len(topLevel))
|
||||
for _, name := range topOrder {
|
||||
commands = append(commands, topLevel[name])
|
||||
}
|
||||
return commands
|
||||
}
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
// 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 compat
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
)
|
||||
|
||||
func TestBuildDynamicCommands_ParentNesting(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
servers := []market.ServerDescriptor{
|
||||
{
|
||||
Endpoint: "https://endpoint-chat",
|
||||
CLI: market.CLIOverlay{
|
||||
ID: "group-chat",
|
||||
Command: "chat",
|
||||
ToolOverrides: map[string]market.CLIToolOverride{
|
||||
"list_conversations": {CLIName: "list"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
Endpoint: "https://endpoint-bot",
|
||||
CLI: market.CLIOverlay{
|
||||
ID: "bot",
|
||||
Command: "bot",
|
||||
Parent: "chat",
|
||||
ToolOverrides: map[string]market.CLIToolOverride{
|
||||
"send_robot_message": {CLIName: "send"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
cmds := BuildDynamicCommands(servers, executor.EchoRunner{}, nil)
|
||||
|
||||
// Should produce only one top-level command: "chat"
|
||||
if len(cmds) != 1 {
|
||||
names := make([]string, len(cmds))
|
||||
for i, c := range cmds {
|
||||
names[i] = c.Name()
|
||||
}
|
||||
t.Fatalf("expected 1 top-level command, got %d: %v", len(cmds), names)
|
||||
}
|
||||
if cmds[0].Name() != "chat" {
|
||||
t.Fatalf("expected top-level command 'chat', got %q", cmds[0].Name())
|
||||
}
|
||||
|
||||
// "bot" should be a sub-command of "chat"
|
||||
found := false
|
||||
for _, sub := range cmds[0].Commands() {
|
||||
if sub.Name() == "bot" {
|
||||
found = true
|
||||
// "bot" should have its own sub-command "send"
|
||||
hasSend := false
|
||||
for _, leaf := range sub.Commands() {
|
||||
if leaf.Name() == "send" {
|
||||
hasSend = true
|
||||
}
|
||||
}
|
||||
if !hasSend {
|
||||
t.Fatal("expected 'bot' to have sub-command 'send'")
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("expected 'bot' as sub-command of 'chat'")
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDynamicCommands_ParentNotFound(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
servers := []market.ServerDescriptor{
|
||||
{
|
||||
Endpoint: "https://endpoint-orphan",
|
||||
CLI: market.CLIOverlay{
|
||||
ID: "orphan",
|
||||
Command: "orphan",
|
||||
Parent: "nonexistent",
|
||||
ToolOverrides: map[string]market.CLIToolOverride{
|
||||
"do_something": {CLIName: "do"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
cmds := BuildDynamicCommands(servers, executor.EchoRunner{}, nil)
|
||||
|
||||
// Parent not found, should fall back to top-level
|
||||
if len(cmds) != 1 {
|
||||
t.Fatalf("expected 1 top-level command, got %d", len(cmds))
|
||||
}
|
||||
if cmds[0].Name() != "orphan" {
|
||||
t.Fatalf("expected top-level command 'orphan', got %q", cmds[0].Name())
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildDynamicCommands_NoParent(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
servers := []market.ServerDescriptor{
|
||||
{
|
||||
Endpoint: "https://endpoint-a",
|
||||
CLI: market.CLIOverlay{
|
||||
ID: "svc-a",
|
||||
Command: "alpha",
|
||||
ToolOverrides: map[string]market.CLIToolOverride{
|
||||
"tool_a": {CLIName: "run"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
Endpoint: "https://endpoint-b",
|
||||
CLI: market.CLIOverlay{
|
||||
ID: "svc-b",
|
||||
Command: "beta",
|
||||
ToolOverrides: map[string]market.CLIToolOverride{
|
||||
"tool_b": {CLIName: "exec"},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
cmds := BuildDynamicCommands(servers, executor.EchoRunner{}, nil)
|
||||
|
||||
if len(cmds) != 2 {
|
||||
t.Fatalf("expected 2 top-level commands, got %d", len(cmds))
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
+130
-15
@@ -21,12 +21,30 @@ import (
|
||||
"log/slog"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/configmeta"
|
||||
)
|
||||
|
||||
func init() {
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_TENANT",
|
||||
Category: configmeta.CategoryCore,
|
||||
Description: "缓存分区的租户标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
configmeta.Register(configmeta.ConfigItem{
|
||||
Name: "DWS_AUTH_IDENTITY",
|
||||
Category: configmeta.CategorySecurity,
|
||||
Description: "缓存分区的认证身份标识",
|
||||
DefaultValue: "default",
|
||||
})
|
||||
}
|
||||
|
||||
const (
|
||||
tenantEnv = "DWS_TENANT"
|
||||
authIdentityEnv = "DWS_AUTH_IDENTITY"
|
||||
@@ -41,6 +59,11 @@ type Service struct {
|
||||
Tenant string
|
||||
AuthIdentity string
|
||||
Logger *slog.Logger
|
||||
// PerServerTimeout overrides the default per-server discovery timeout
|
||||
// when greater than zero. Useful for tests and for callers that need a
|
||||
// tighter or looser bound. When zero, defaultPerServerDiscoveryTimeout
|
||||
// applies.
|
||||
PerServerTimeout time.Duration
|
||||
}
|
||||
|
||||
type RuntimeServer struct {
|
||||
@@ -152,29 +175,121 @@ func (s *Service) DiscoverServerRuntime(ctx context.Context, server market.Serve
|
||||
}, nil
|
||||
}
|
||||
|
||||
// defaultPerServerDiscoveryTimeout bounds the time spent discovering tools on
|
||||
// a single registry-listed server. Tightened to 2s so a slow/unreachable
|
||||
// server cannot stall every CLI command — a healthy MCP endpoint negotiates
|
||||
// well under a second. See issue #119.
|
||||
const defaultPerServerDiscoveryTimeout = 2 * time.Second
|
||||
|
||||
func (s *Service) DiscoverAllRuntime(ctx context.Context, servers []market.ServerDescriptor) ([]RuntimeServer, []RuntimeFailure) {
|
||||
results := make([]RuntimeServer, 0, len(servers))
|
||||
failures := make([]RuntimeFailure, 0)
|
||||
for _, server := range servers {
|
||||
if server.CLI.Skip {
|
||||
continue
|
||||
type discoveryResult struct {
|
||||
server RuntimeServer
|
||||
failure *RuntimeFailure
|
||||
}
|
||||
|
||||
perServerTimeout := defaultPerServerDiscoveryTimeout
|
||||
if s.PerServerTimeout > 0 {
|
||||
perServerTimeout = s.PerServerTimeout
|
||||
}
|
||||
|
||||
filtered := make([]market.ServerDescriptor, 0, len(servers))
|
||||
for _, srv := range servers {
|
||||
if !srv.CLI.Skip {
|
||||
filtered = append(filtered, srv)
|
||||
}
|
||||
runtimeServer, err := s.DiscoverServerRuntime(ctx, server)
|
||||
if err != nil {
|
||||
if errors.Is(err, errCLIServerSkipped) {
|
||||
continue
|
||||
}
|
||||
if len(filtered) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
ch := make(chan discoveryResult, len(filtered))
|
||||
var wg sync.WaitGroup
|
||||
for _, srv := range filtered {
|
||||
wg.Add(1)
|
||||
go func(server market.ServerDescriptor) {
|
||||
defer wg.Done()
|
||||
serverCtx, cancel := context.WithTimeout(ctx, perServerTimeout)
|
||||
defer cancel()
|
||||
start := time.Now()
|
||||
rs, err := s.DiscoverServerRuntime(serverCtx, server)
|
||||
elapsed := time.Since(start)
|
||||
if err != nil {
|
||||
if errors.Is(err, errCLIServerSkipped) {
|
||||
return
|
||||
}
|
||||
if s.Logger != nil {
|
||||
s.Logger.Warn("server_discovery_failed",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
|
||||
slog.String("error", err.Error()),
|
||||
slog.Bool("is_timeout", errors.Is(err, context.DeadlineExceeded)),
|
||||
)
|
||||
}
|
||||
// Per-server sub-context timed out but parent is still alive:
|
||||
// try cache fallback instead of reporting a hard failure.
|
||||
if errors.Is(err, context.DeadlineExceeded) && ctx.Err() == nil {
|
||||
if cached, cacheErr := s.loadServerFromCache(server); cacheErr == nil {
|
||||
if s.Logger != nil {
|
||||
s.Logger.Info("server_discovery_cache_fallback",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("source", cached.Source),
|
||||
)
|
||||
}
|
||||
ch <- discoveryResult{server: cached}
|
||||
return
|
||||
}
|
||||
}
|
||||
ch <- discoveryResult{failure: &RuntimeFailure{ServerKey: server.Key, Err: err}}
|
||||
return
|
||||
}
|
||||
failures = append(failures, RuntimeFailure{
|
||||
ServerKey: server.Key,
|
||||
Err: err,
|
||||
})
|
||||
continue
|
||||
if s.Logger != nil {
|
||||
s.Logger.Debug("server_discovery_ok",
|
||||
slog.String("server_key", server.Key),
|
||||
slog.String("duration", elapsed.Truncate(time.Millisecond).String()),
|
||||
slog.String("source", rs.Source),
|
||||
)
|
||||
}
|
||||
ch <- discoveryResult{server: rs}
|
||||
}(srv)
|
||||
}
|
||||
go func() {
|
||||
wg.Wait()
|
||||
close(ch)
|
||||
}()
|
||||
|
||||
results := make([]RuntimeServer, 0, len(filtered))
|
||||
failures := make([]RuntimeFailure, 0)
|
||||
for dr := range ch {
|
||||
if dr.failure != nil {
|
||||
failures = append(failures, *dr.failure)
|
||||
} else {
|
||||
results = append(results, dr.server)
|
||||
}
|
||||
results = append(results, runtimeServer)
|
||||
}
|
||||
return results, failures
|
||||
}
|
||||
|
||||
// loadServerFromCache tries to load a server's tools from cache, returning a
|
||||
// degraded RuntimeServer. Used as fallback when a per-server discovery timeout
|
||||
// fires but the parent context is still alive.
|
||||
func (s *Service) loadServerFromCache(server market.ServerDescriptor) (RuntimeServer, error) {
|
||||
partition := s.partition()
|
||||
snapshot, freshness, err := s.Cache.LoadTools(partition, server.Key)
|
||||
if err != nil {
|
||||
return RuntimeServer{}, err
|
||||
}
|
||||
server.NegotiatedProtocolVersion = snapshot.ProtocolVersion
|
||||
server.Source = string(freshness) + "_cache"
|
||||
server.Degraded = true
|
||||
return RuntimeServer{
|
||||
Server: server,
|
||||
NegotiatedProtocolVersion: snapshot.ProtocolVersion,
|
||||
Tools: snapshot.Tools,
|
||||
Source: string(freshness) + "_cache",
|
||||
Degraded: true,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Service) DiscoverDetail(ctx context.Context, server market.ServerDescriptor) (market.DetailResponse, error) {
|
||||
partition := s.partition()
|
||||
var fetchErr error
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
+120
-35
@@ -14,13 +14,14 @@
|
||||
package errors
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"bytes"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
// Category represents a stable error class with a documented exit code.
|
||||
@@ -36,18 +37,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 +200,31 @@ func NewInternal(message string, opts ...Option) error {
|
||||
return newError(CategoryInternal, message, opts...)
|
||||
}
|
||||
|
||||
// ExitCoder is implemented by errors that provide their own exit code.
|
||||
// Edition-specific error types (e.g. PATError, CLIError) implement this
|
||||
// so the framework can resolve exit codes without importing edition packages.
|
||||
type ExitCoder interface {
|
||||
ExitCode() int
|
||||
}
|
||||
|
||||
// RawStderrError is implemented by errors that must output raw content
|
||||
// directly to stderr, bypassing all CLI formatting (e.g. "Error:" prefix).
|
||||
// PAT authorization errors use this to pass JSON through to the desktop runtime.
|
||||
type RawStderrError interface {
|
||||
error
|
||||
RawStderr() string
|
||||
}
|
||||
|
||||
// ExitCode maps any error to a stable exit code.
|
||||
func ExitCode(err error) int {
|
||||
var typed *Error
|
||||
if stderrors.As(err, &typed) {
|
||||
return typed.ExitCode()
|
||||
}
|
||||
var ec ExitCoder
|
||||
if stderrors.As(err, &ec) {
|
||||
return ec.ExitCode()
|
||||
}
|
||||
return 5
|
||||
}
|
||||
|
||||
@@ -247,6 +268,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"] = config.GetDeveloperSettingsURL()
|
||||
}
|
||||
}
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
errorPayload["technical_detail"] = typed.ServerDiag.TechnicalDetail
|
||||
}
|
||||
}
|
||||
if typed.Cause != nil {
|
||||
errorPayload["cause"] = typed.Cause.Error()
|
||||
}
|
||||
@@ -263,8 +301,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 +330,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: 开启地址: "+config.GetDeveloperSettingsURL())
|
||||
}
|
||||
|
||||
if len(typed.Actions) > 0 {
|
||||
for _, action := range typed.Actions {
|
||||
if strings.TrimSpace(action) == "" {
|
||||
@@ -298,22 +355,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())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,289 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package errors
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ExitCodePermission is the process exit code for PAT authorisation failures.
|
||||
const ExitCodePermission = 4
|
||||
|
||||
// PATError represents a PAT (Personal Action Token) authorization failure
|
||||
// that should be passed through to stderr as raw JSON without any CLI-layer
|
||||
// wrapping. The host application (e.g. RewindDesktop) parses the JSON to
|
||||
// display its own authorisation UI.
|
||||
type PATError struct {
|
||||
RawJSON string
|
||||
}
|
||||
|
||||
func (e *PATError) Error() string { return e.RawJSON }
|
||||
|
||||
// ExitCode returns the documented exit code for PAT permission errors (4).
|
||||
func (e *PATError) ExitCode() int { return ExitCodePermission }
|
||||
|
||||
// RawStderr returns the raw JSON to be written directly to stderr.
|
||||
func (e *PATError) RawStderr() string { return e.RawJSON }
|
||||
|
||||
// patNoPermissionCodes are PAT error codes that should be passed through
|
||||
// as transparent PATError without CLI-level wrapping.
|
||||
var patNoPermissionCodes = map[string]bool{
|
||||
"PAT_NO_PERMISSION": true,
|
||||
"PAT_LOW_RISK_NO_PERMISSION": true,
|
||||
"PAT_MEDIUM_RISK_NO_PERMISSION": true,
|
||||
"PAT_HIGH_RISK_NO_PERMISSION": true,
|
||||
}
|
||||
|
||||
// patAuthRequiredCodes are error codes that trigger the PAT authorization
|
||||
// flow (e.g. the server auto-created a CLI app and returned auth details).
|
||||
var patAuthRequiredCodes = map[string]bool{
|
||||
"AGENT_CODE_NOT_EXISTS": true,
|
||||
}
|
||||
|
||||
// IsPATError reports whether err is a *PATError.
|
||||
func IsPATError(err error) bool {
|
||||
_, ok := err.(*PATError)
|
||||
return ok
|
||||
}
|
||||
|
||||
// IsPATNoPermissionCode reports whether code is a known PAT permission error code.
|
||||
func IsPATNoPermissionCode(code string) bool {
|
||||
return patNoPermissionCodes[code]
|
||||
}
|
||||
|
||||
// ---- DWS gateway auth errors (shared between PAT & general auth) ----------
|
||||
|
||||
// dwsGatewayErrors is the set of DWS gateway-level auth error codes.
|
||||
var dwsGatewayErrors = map[string]bool{
|
||||
"DWS_SERVICE_UNAUTHORIZED": true,
|
||||
"DWS_AUTH_SERVICE_FAILED": true,
|
||||
}
|
||||
|
||||
// getDWSGatewayErrorCode extracts a DWS gateway error code from errBody
|
||||
// (supports both errorCode and error_code field names).
|
||||
func getDWSGatewayErrorCode(errBody map[string]any) (string, bool) {
|
||||
for _, key := range []string{"errorCode", "error_code"} {
|
||||
if code, ok := errBody[key].(string); ok && dwsGatewayErrors[code] {
|
||||
return code, true
|
||||
}
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// isNotLoggedInError checks if the error body indicates missing authentication.
|
||||
func isNotLoggedInError(body map[string]any) bool {
|
||||
if errMsg, ok := body["error"].(string); ok {
|
||||
if strings.Contains(errMsg, "Missing service_id or access_key") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// isBusinessError checks if a parsed JSON body represents a business-level error.
|
||||
func isBusinessError(body map[string]any) bool {
|
||||
if _, ok := body["error"].(string); ok {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(bool); ok && !v {
|
||||
return true
|
||||
}
|
||||
if v, ok := body["success"].(string); ok && strings.EqualFold(v, "false") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ---- Classification functions -----------------------------------------------
|
||||
|
||||
// ClassifyToolResultContent checks a raw MCP tool result content map for
|
||||
// DWS gateway auth errors and PAT permission error codes. This is intended
|
||||
// for use as the edition.Hooks.ClassifyToolResult callback so the framework's
|
||||
// runner returns a typed error before its generic business-error classification.
|
||||
//
|
||||
// Check order: DWS gateway auth > PAT permission.
|
||||
func ClassifyToolResultContent(content map[string]any) error {
|
||||
if _, ok := getDWSGatewayErrorCode(content); ok {
|
||||
raw, _ := json.Marshal(content)
|
||||
return NewAuth(string(raw),
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ClassifyMCPResponseText classifies a text response returned by an MCP tool call.
|
||||
// Returns a typed error for known gateway auth failures, PAT interceptions,
|
||||
// and business-level errors embedded in HTTP-200 JSON bodies.
|
||||
//
|
||||
// Check order: DWS gateway > PAT permission > generic business error.
|
||||
func ClassifyMCPResponseText(text string) error {
|
||||
var body map[string]any
|
||||
if json.Unmarshal([]byte(text), &body) != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
if _, ok := getDWSGatewayErrorCode(body); ok {
|
||||
return NewAuth(text,
|
||||
WithReason("gateway_auth_expired"),
|
||||
WithHint(authExpiredHint()),
|
||||
)
|
||||
}
|
||||
|
||||
if isNotLoggedInError(body) {
|
||||
return NewAuth("当前未登录",
|
||||
WithReason("not_configured"),
|
||||
WithHint(notLoggedInHint()),
|
||||
WithActions("dws auth login"),
|
||||
)
|
||||
}
|
||||
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := body[key].(string); ok && patNoPermissionCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(body, code)}
|
||||
}
|
||||
}
|
||||
|
||||
if isBusinessError(body) {
|
||||
return NewAPI(text,
|
||||
WithReason("business_error"),
|
||||
WithHint(suggestForBusinessErrorText(body)),
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ---- Hints -----------------------------------------------------------------
|
||||
|
||||
func authExpiredHint() string {
|
||||
return "Re-authenticate: dws auth login"
|
||||
}
|
||||
|
||||
func notLoggedInHint() string {
|
||||
return "请先登录:dws auth login"
|
||||
}
|
||||
|
||||
func suggestForBusinessErrorText(body map[string]any) string {
|
||||
msg := ""
|
||||
if v, ok := body["errorMsg"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["message"].(string); ok {
|
||||
msg = v
|
||||
} else if v, ok := body["error"].(string); ok {
|
||||
msg = v
|
||||
}
|
||||
switch {
|
||||
case strings.Contains(msg, "搜索内容不能为空"):
|
||||
return "请提供非空搜索关键词: dws doc search --query \"关键词\""
|
||||
case strings.Contains(msg, "User has no permission to access this email"):
|
||||
return "请确认邮箱地址正确,查看可用邮箱: dws mail mailbox list"
|
||||
case strings.Contains(msg, "频率超限") || strings.Contains(msg, "rate limit"):
|
||||
return "API rate limit exceeded, wait a moment and retry"
|
||||
case strings.Contains(msg, "参数错误") || strings.Contains(msg, "param error"):
|
||||
return "Check input parameters. Use --help for available flags"
|
||||
default:
|
||||
return "MCP tool returned a business error; check parameters and refer to skill documentation."
|
||||
}
|
||||
}
|
||||
|
||||
// ---- PAT JSON helpers ------------------------------------------------------
|
||||
|
||||
var patTopLevelStrip = map[string]bool{
|
||||
"success": true, "code": true, "errorCode": true, "error_code": true,
|
||||
"message": true, "error": true, "trace_id": true, "class": true,
|
||||
}
|
||||
|
||||
func cleanPATJSON(body map[string]any, code string) string {
|
||||
out := map[string]any{
|
||||
"success": false,
|
||||
"code": code,
|
||||
}
|
||||
if data, ok := body["data"]; ok {
|
||||
out["data"] = stripClassFields(data)
|
||||
} else {
|
||||
fallback := map[string]any{}
|
||||
for k, v := range body {
|
||||
if !patTopLevelStrip[k] {
|
||||
fallback[k] = v
|
||||
}
|
||||
}
|
||||
if len(fallback) > 0 {
|
||||
out["data"] = stripClassFields(fallback)
|
||||
}
|
||||
}
|
||||
b, err := json.MarshalIndent(out, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Sprintf(`{"success":false,"code":"%s"}`, code)
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
|
||||
// ---- Runner adapter functions ------------------------------------------------
|
||||
// These match the function signatures referenced by runner.go's PAT check
|
||||
// framework (ClassifyPatAuthCheck / AsPatAuthCheckError).
|
||||
|
||||
// ClassifyPatAuthCheck is the open-source fallback that checks a tool-call
|
||||
// Content map for PAT permission codes and auth-required codes. Returns a
|
||||
// non-nil *PATError when the content carries a recognised PAT/auth error.
|
||||
func ClassifyPatAuthCheck(content map[string]any) *PATError {
|
||||
for _, key := range []string{"code", "errorCode"} {
|
||||
if code, ok := content[key].(string); ok {
|
||||
if patNoPermissionCodes[code] || patAuthRequiredCodes[code] {
|
||||
return &PATError{RawJSON: cleanPATJSON(content, code)}
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// AsPatAuthCheckError extracts a *PATError from an error chain.
|
||||
func AsPatAuthCheckError(err error) *PATError {
|
||||
var patErr *PATError
|
||||
if stderrors.As(err, &patErr) {
|
||||
return patErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func stripClassFields(v any) any {
|
||||
switch val := v.(type) {
|
||||
case map[string]any:
|
||||
clean := make(map[string]any, len(val))
|
||||
for k, item := range val {
|
||||
if k == "class" {
|
||||
continue
|
||||
}
|
||||
clean[k] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
case []any:
|
||||
clean := make([]any, len(val))
|
||||
for i, item := range val {
|
||||
clean[i] = stripClassFields(item)
|
||||
}
|
||||
return clean
|
||||
default:
|
||||
return v
|
||||
}
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user