Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
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 | ||
|
|
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: 52.6%"><title>coverage: 52.6%</title><linearGradient id="s" x2="0" y2="100%"><stop offset="0" stop-color="#bbb" stop-opacity=".1"/><stop offset="1" stop-opacity=".1"/></linearGradient><clipPath id="r"><rect width="108" height="20" rx="3" fill="#fff"/></clipPath><g clip-path="url(#r)"><rect width="61" height="20" fill="#555"/><rect x="61" width="47" height="20" fill="#e05d44"/><rect width="108" height="20" fill="url(#s)"/></g><g fill="#fff" text-anchor="middle" font-family="Verdana,Geneva,DejaVu Sans,sans-serif" text-rendering="geometricPrecision" font-size="110"><text aria-hidden="true" x="315" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="510">coverage</text><text x="315" y="140" transform="scale(.1)" fill="#fff" textLength="510">coverage</text><text aria-hidden="true" x="835" y="150" fill="#010101" fill-opacity=".3" transform="scale(.1)" textLength="370">52.6%</text><text x="835" y="140" transform="scale(.1)" fill="#fff" textLength="370">52.6%</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)
|
||||
}
|
||||
};
|
||||
|
||||
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 }}
|
||||
+20
-41
@@ -1,51 +1,30 @@
|
||||
# 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
|
||||
|
||||
@@ -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
|
||||
@@ -4,6 +4,54 @@ 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.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,421 @@
|
||||
<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:
|
||||
**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
|
||||
|
||||
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 | `chat` | 10 | `message` `group` `search` | Group CRUD, member management, bot messaging, webhook |
|
||||
| Bot | `chat bot` | 6 | `bot` `group` `message` `search` | Robot creation/search, group/single messaging, webhook, message recall |
|
||||
| Calendar | `calendar` | 13 | `event` `room` `participant` `busy` | Events CRUD, meeting room booking, free-busy query, participant management |
|
||||
| Todo | `todo` | 6 | `task` | Create, list, update, done, get detail, delete |
|
||||
| Approval | `oa` | 9 | `approval` | Approve/reject/revoke, pending tasks, initiated instances, process list |
|
||||
| Attendance | `attendance` | 4 | `record` `shift` `summary` `rules` | Clock-in records, shift schedules, attendance summary, group rules |
|
||||
| Ding | `ding` | 2 | `message` | Send/recall DING messages |
|
||||
| Report | `report` | 7 | `create` `list` `detail` `template` `stats` `sent` | Create reports, sent/received list, templates, statistics |
|
||||
| AITable | `aitable` | 20 | `base` `table` `record` `field` `attachment` `template` | Full CRUD for bases/tables/records/fields, templates |
|
||||
| Workbench | `workbench` | 2 | `app` | Batch query app details |
|
||||
| DevDoc | `devdoc` | 1 | `article` | Search platform docs and error codes |
|
||||
|
||||
## 退出码
|
||||
> 86 commands across 12 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` 字段),便于机器消费。
|
||||
`doc` (documents) · `mail` (email) · `minutes` (AI transcription) · `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
|
||||
+423
@@ -0,0 +1,423 @@
|
||||
<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>
|
||||
|
||||
**预编译二进制文件**:从 [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>
|
||||
|
||||
## 升级
|
||||
|
||||
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
|
||||
@@ -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=
|
||||
|
||||
+169
-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)
|
||||
@@ -105,6 +124,14 @@ func newAuthLoginCommand() *cobra.Command {
|
||||
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 +157,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 +188,30 @@ 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"))
|
||||
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 +226,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 +265,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",
|
||||
@@ -329,7 +350,9 @@ func newAuthResetCommand() *cobra.Command {
|
||||
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 +438,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")
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -16,6 +16,8 @@ package app
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
// Build-time variables injected via ldflags when available.
|
||||
@@ -28,7 +30,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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -23,7 +23,9 @@ type GlobalFlags struct {
|
||||
ClientSecret string
|
||||
Debug bool
|
||||
DryRun bool
|
||||
Fields string
|
||||
Format string
|
||||
JQ string
|
||||
Mock bool
|
||||
Output string
|
||||
Timeout int
|
||||
@@ -37,7 +39,9 @@ func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
|
||||
cmd.PersistentFlags().StringVar(&flags.ClientSecret, "client-secret", "", "Override OAuth client secret (DingTalk 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")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+71
-10
@@ -16,6 +16,7 @@ package app
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
@@ -29,16 +30,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 +56,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,6 +89,12 @@ func newLegacyPublicCommands(ctx context.Context, runner executor.Runner) []*cob
|
||||
// Tests may override discoveryBaseURLOverride to redirect to a local server;
|
||||
// in that case the registry cache is always bypassed.
|
||||
func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.Command {
|
||||
totalStart := time.Now()
|
||||
defer func() {
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] loadDynamicCommands total: %v\n", time.Since(totalStart))
|
||||
}
|
||||
}()
|
||||
|
||||
store := cacheStoreFromEnv()
|
||||
partition := config.DefaultPartition
|
||||
@@ -70,13 +106,20 @@ 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)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache load: %v (err=%v)\n", time.Since(cacheLoadStart), cacheErr)
|
||||
}
|
||||
|
||||
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)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] using cached registry: servers=%d, freshness=%s\n", len(snapshot.Servers), freshness)
|
||||
}
|
||||
servers = snapshot.Servers
|
||||
// Only trigger async revalidation in production (no URL override).
|
||||
// Tests set discoveryBaseURLOverride and control cache expiry directly,
|
||||
@@ -92,9 +135,15 @@ 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)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] fetching from market API: %s\n", baseURL)
|
||||
}
|
||||
fetchStart := time.Now()
|
||||
client := market.NewClient(baseURL, ipv4OnlyHTTPClient())
|
||||
resp, fetchErr := client.FetchServers(ctx, config.DefaultFetchServersLimit)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] market API fetch: %v (err=%v)\n", time.Since(fetchStart), fetchErr)
|
||||
}
|
||||
if fetchErr != nil {
|
||||
slog.Debug("loadDynamicCommands: market API fetch failed", "error", fetchErr)
|
||||
// Degrade to stale cache if available (production only).
|
||||
@@ -106,12 +155,18 @@ func loadDynamicCommands(ctx context.Context, runner executor.Runner) []*cobra.C
|
||||
}
|
||||
} else {
|
||||
servers = market.NormalizeServers(resp, "market")
|
||||
slog.Debug("loadDynamicCommands: normalized servers", "count", len(servers))
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] normalized servers: %d\n", len(servers))
|
||||
}
|
||||
// Persist fresh data (only in non-test mode).
|
||||
if useCache {
|
||||
saveStart := time.Now()
|
||||
if saveErr := store.SaveRegistry(partition, cache.RegistrySnapshot{Servers: servers}); saveErr != nil {
|
||||
slog.Debug("loadDynamicCommands: failed to save registry cache", "error", saveErr)
|
||||
}
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cache save: %v\n", time.Since(saveStart))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -122,9 +177,17 @@ 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)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] load details: %v\n", time.Since(detailStart))
|
||||
}
|
||||
|
||||
buildStart := time.Now()
|
||||
cmds := compat.BuildDynamicCommands(servers, runner, detailsByID)
|
||||
slog.Debug("loadDynamicCommands: built dynamic commands", "commands", len(cmds))
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] build commands: %v (count=%d)\n", time.Since(buildStart), len(cmds))
|
||||
}
|
||||
|
||||
return cmds
|
||||
}
|
||||
@@ -332,10 +395,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,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
|
||||
}
|
||||
+239
-28
@@ -29,27 +29,66 @@ import (
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cache"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/discovery"
|
||||
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/generator"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/market"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pipeline/handlers"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/recovery"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
)
|
||||
|
||||
type outputFileContextKey struct{}
|
||||
|
||||
const recoveryEventStderrPrefix = "RECOVERY_EVENT_ID="
|
||||
|
||||
// Execute runs the root command and returns the process exit code.
|
||||
func Execute() int {
|
||||
totalStart := time.Now()
|
||||
timing := NewTimingCollector()
|
||||
defer func() {
|
||||
timing.PrintIfEnabled()
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] Execute total: %v\n", time.Since(totalStart))
|
||||
}
|
||||
}()
|
||||
|
||||
ctx, cancel := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
|
||||
defer cancel()
|
||||
|
||||
root := NewRootCommand(ctx)
|
||||
// Attach timing collector to context for use by child components
|
||||
ctx = WithTimingCollector(ctx, timing)
|
||||
|
||||
initStart := time.Now()
|
||||
recovery.ResetRuntimeState()
|
||||
engine := newPipelineEngine()
|
||||
root := NewRootCommandWithEngine(ctx, engine)
|
||||
initDuration := time.Since(initStart)
|
||||
timing.Record("cmd_init", initDuration)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] command init: %v\n", initDuration)
|
||||
}
|
||||
|
||||
// Run PreParse handlers on raw argv before Cobra parses flags.
|
||||
// This corrects model-generated errors like --userId → --user-id
|
||||
// and --limit100 → --limit 100.
|
||||
pipeline.RunPreParse(root, engine)
|
||||
|
||||
execStart := time.Now()
|
||||
executed, err := root.ExecuteC()
|
||||
execDuration := time.Since(execStart)
|
||||
timing.Record("cobra_exec", execDuration)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] cobra ExecuteC: %v\n", execDuration)
|
||||
}
|
||||
if err != nil {
|
||||
if executed == nil {
|
||||
executed = root
|
||||
@@ -60,6 +99,9 @@ func Execute() int {
|
||||
_, _ = fmt.Fprintln(os.Stderr)
|
||||
}
|
||||
_ = printExecutionError(executed, os.Stdout, os.Stderr, err)
|
||||
if last := recovery.LatestCapture(); last != nil && last.EventID != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "%s%s\n", recoveryEventStderrPrefix, last.EventID)
|
||||
}
|
||||
return apperrors.ExitCode(err)
|
||||
}
|
||||
return 0
|
||||
@@ -69,11 +111,52 @@ func isUnknownCommandError(err error) bool {
|
||||
return err != nil && strings.Contains(err.Error(), "unknown command")
|
||||
}
|
||||
|
||||
// flagErrorWithSuggestions provides helpful suggestions for common flag mistakes.
|
||||
func flagErrorWithSuggestions(cmd *cobra.Command, err error) error {
|
||||
errMsg := err.Error()
|
||||
|
||||
// Common flag aliases and suggestions
|
||||
suggestions := map[string]string{
|
||||
"--json": "提示: 请使用 --format json 或 -f json 来输出 JSON 格式",
|
||||
"--method": "提示: dws auth login 默认使用 OAuth 设备流登录,无需指定 --method",
|
||||
"--device-flow": "提示: dws auth login 默认已使用设备流,无需 --device-flow 参数",
|
||||
"--email": "提示: dws 不支持邮箱/密码登录,请使用 dws auth login 进行扫码登录",
|
||||
"--code": "提示: dws 不支持验证码登录,请使用 dws auth login 进行扫码登录",
|
||||
"--corp-id": "提示: corp-id 会在登录时自动获取,无需手动指定",
|
||||
"--password": "提示: dws 不支持密码登录,请使用 dws auth login 进行扫码登录",
|
||||
"--phone": "提示: dws 不支持手机号登录,请使用 dws auth login 进行扫码登录",
|
||||
"--app-key": "提示: 请使用环境变量 DWS_CLIENT_ID 或 --client-id 设置 AppKey",
|
||||
"--app-secret": "提示: 请使用环境变量 DWS_CLIENT_SECRET 或 --client-secret 设置 AppSecret",
|
||||
}
|
||||
|
||||
for flag, suggestion := range suggestions {
|
||||
if strings.Contains(errMsg, "unknown flag: "+flag) {
|
||||
return fmt.Errorf("%w\n%s", err, suggestion)
|
||||
}
|
||||
}
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
func printExecutionError(root *cobra.Command, stdout, stderr io.Writer, err error) error {
|
||||
if wantsJSONErrors(root) {
|
||||
return apperrors.PrintJSON(stdout, err)
|
||||
}
|
||||
return apperrors.PrintHuman(stderr, err)
|
||||
return apperrors.PrintHumanAt(stderr, err, resolveVerbosity(root))
|
||||
}
|
||||
|
||||
// resolveVerbosity derives the error verbosity level from the root command's flags.
|
||||
func resolveVerbosity(cmd *cobra.Command) apperrors.Verbosity {
|
||||
if cmd == nil {
|
||||
return apperrors.VerbosityNormal
|
||||
}
|
||||
if debug, err := cmd.Flags().GetBool("debug"); err == nil && debug {
|
||||
return apperrors.VerbosityDebug
|
||||
}
|
||||
if verbose, err := cmd.Flags().GetBool("verbose"); err == nil && verbose {
|
||||
return apperrors.VerbosityVerbose
|
||||
}
|
||||
return apperrors.VerbosityNormal
|
||||
}
|
||||
|
||||
func wantsJSONErrors(root *cobra.Command) bool {
|
||||
@@ -130,13 +213,24 @@ func NewRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
var rootCtx context.Context
|
||||
if len(ctx) > 0 && ctx[0] != nil {
|
||||
rootCtx = ctx[0]
|
||||
} else {
|
||||
}
|
||||
return NewRootCommandWithEngine(rootCtx, nil)
|
||||
}
|
||||
|
||||
// NewRootCommandWithEngine constructs the root CLI command with an
|
||||
// optional pipeline engine for input correction. When engine is nil,
|
||||
// no pipeline processing is applied.
|
||||
func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine) *cobra.Command {
|
||||
if rootCtx == nil {
|
||||
rootCtx = context.Background()
|
||||
}
|
||||
flags := &GlobalFlags{}
|
||||
loader := cli.EnvironmentLoader{
|
||||
LookupEnv: os.LookupEnv,
|
||||
CatalogBaseURLOverride: DiscoveryBaseURL(),
|
||||
AuthTokenFunc: func(ctx context.Context) string {
|
||||
return resolveRuntimeAuthToken(ctx, "")
|
||||
},
|
||||
}
|
||||
runner := newCommandRunnerWithFlags(loader, flags)
|
||||
|
||||
@@ -166,6 +260,7 @@ func NewRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
return configureOutputSink(cmd)
|
||||
},
|
||||
PersistentPostRunE: func(cmd *cobra.Command, args []string) error {
|
||||
CloseFileLogger()
|
||||
return closeOutputSink(cmd)
|
||||
},
|
||||
}
|
||||
@@ -173,16 +268,18 @@ func NewRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
bindPersistentFlags(root, flags)
|
||||
|
||||
schemaCmd := newSchemaCommand(loader)
|
||||
schemaCmd.Hidden = true
|
||||
genSkillsCmd := newGenerateSkillsCommand()
|
||||
genSkillsCmd.Hidden = true
|
||||
mcpCmd := newMCPCommand(rootCtx, loader, runner)
|
||||
mcpCmd := newMCPCommand(rootCtx, loader, runner, engine)
|
||||
mcpCmd.Hidden = true
|
||||
|
||||
utilityCommands := []*cobra.Command{
|
||||
newAuthCommand(),
|
||||
newSkillCommand(),
|
||||
newCacheCommand(),
|
||||
newCompletionCommand(root),
|
||||
newRecoveryCommand(rootCtx, loader, flags),
|
||||
newUpgradeCommand(),
|
||||
newVersionCommand(),
|
||||
schemaCmd,
|
||||
genSkillsCmd,
|
||||
@@ -192,8 +289,16 @@ func NewRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
root.AddCommand(newLegacyPublicCommands(rootCtx, runner)...)
|
||||
root.AddCommand(newLegacyHiddenCommands(runner)...)
|
||||
|
||||
if fn := edition.Get().RegisterExtraCommands; fn != nil {
|
||||
caller := newToolCallerAdapter(runner, flags)
|
||||
fn(root, caller)
|
||||
deduplicateCommands(root)
|
||||
}
|
||||
|
||||
hideNonDirectRuntimeCommands(root)
|
||||
configureRootHelp(root)
|
||||
// Set custom flag error handler for better UX
|
||||
root.SetFlagErrorFunc(flagErrorWithSuggestions)
|
||||
root.SetContext(rootCtx)
|
||||
|
||||
return root
|
||||
@@ -203,6 +308,10 @@ func newAuthCommand() *cobra.Command {
|
||||
return buildAuthCommand()
|
||||
}
|
||||
|
||||
func newSkillCommand() *cobra.Command {
|
||||
return buildSkillCommand()
|
||||
}
|
||||
|
||||
func newCacheCommand() *cobra.Command {
|
||||
cacheCmd := newPlaceholderParent("cache", "缓存管理")
|
||||
|
||||
@@ -364,24 +473,51 @@ func newVersionCommand() *cobra.Command {
|
||||
Example: " dws version\n dws version --format json",
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
format, err := cmd.Flags().GetString("format")
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read format flag")
|
||||
wantJSON := cmd.Flags().Changed("format")
|
||||
if wantJSON {
|
||||
format, _ := cmd.Flags().GetString("format")
|
||||
wantJSON = (format == "json")
|
||||
}
|
||||
payload := map[string]any{
|
||||
"version": Version(),
|
||||
"go": "1.24+",
|
||||
|
||||
editionName := edition.Get().Name
|
||||
if editionName == "" {
|
||||
editionName = "open"
|
||||
}
|
||||
if format == "json" {
|
||||
ver := RawVersion()
|
||||
bt := BuildTime()
|
||||
gc := GitCommit()
|
||||
goVer := "1.24+"
|
||||
|
||||
arch := "MCP Dynamic Aggregation"
|
||||
|
||||
if wantJSON {
|
||||
payload := map[string]any{
|
||||
"version": ver,
|
||||
"edition": editionName,
|
||||
"architecture": arch,
|
||||
"go": goVer,
|
||||
}
|
||||
if bt != "unknown" {
|
||||
payload["build"] = bt
|
||||
}
|
||||
if gc != "unknown" {
|
||||
payload["commit"] = gc
|
||||
}
|
||||
return output.WriteJSON(cmd.OutOrStdout(), payload)
|
||||
}
|
||||
_, err = fmt.Fprintf(
|
||||
cmd.OutOrStdout(),
|
||||
"版本: %s\nGo: %s\n",
|
||||
Version(),
|
||||
"1.24+",
|
||||
)
|
||||
return err
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Version:", ver)
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Edition:", editionName)
|
||||
if bt != "unknown" {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Build:", bt)
|
||||
}
|
||||
if gc != "unknown" {
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Commit:", gc)
|
||||
}
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Architecture:", arch)
|
||||
fmt.Fprintf(w, "%-16s%s\n", "Go:", goVer)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -473,22 +609,35 @@ func newGenerateSkillsCommand() *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newMCPCommand(ctx context.Context, loader cli.CatalogLoader, runner executor.Runner) *cobra.Command {
|
||||
return cli.NewMCPCommand(ctx, loader, runner)
|
||||
func newMCPCommand(ctx context.Context, loader cli.CatalogLoader, runner executor.Runner, engine *pipeline.Engine) *cobra.Command {
|
||||
return cli.NewMCPCommand(ctx, loader, runner, engine)
|
||||
}
|
||||
|
||||
// hideNonDirectRuntimeCommands marks top-level product commands as hidden
|
||||
// unless they correspond to a product discovered via dynamic server discovery.
|
||||
// unless they correspond to a product discovered via dynamic server discovery
|
||||
// or listed in the edition's VisibleProducts hook.
|
||||
// Public utility commands (auth, cache, completion, version) are always kept
|
||||
// visible; explicitly hidden commands stay hidden.
|
||||
func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
allowedProducts := DirectRuntimeProductIDs()
|
||||
var allowedProducts map[string]bool
|
||||
if fn := edition.Get().VisibleProducts; fn != nil {
|
||||
products := fn()
|
||||
allowedProducts = make(map[string]bool, len(products))
|
||||
for _, p := range products {
|
||||
allowedProducts[p] = true
|
||||
}
|
||||
} else {
|
||||
allowedProducts = DirectRuntimeProductIDs()
|
||||
}
|
||||
staticCommands := map[string]bool{
|
||||
"auth": true,
|
||||
"cache": true,
|
||||
"completion": true,
|
||||
"version": true,
|
||||
"help": true,
|
||||
"recovery": true,
|
||||
"schema": true,
|
||||
"mcp": true,
|
||||
}
|
||||
for _, cmd := range root.Commands() {
|
||||
name := cmd.Name()
|
||||
@@ -505,6 +654,24 @@ func hideNonDirectRuntimeCommands(root *cobra.Command) {
|
||||
}
|
||||
}
|
||||
|
||||
// deduplicateCommands removes duplicate top-level commands, keeping the last
|
||||
// registered one. This ensures overlay commands take precedence over
|
||||
// open-source defaults when both register the same product name.
|
||||
func deduplicateCommands(root *cobra.Command) {
|
||||
seen := make(map[string]*cobra.Command)
|
||||
var dups []*cobra.Command
|
||||
for _, cmd := range root.Commands() {
|
||||
name := cmd.Name()
|
||||
if prev, ok := seen[name]; ok {
|
||||
dups = append(dups, prev)
|
||||
}
|
||||
seen[name] = cmd
|
||||
}
|
||||
for _, dup := range dups {
|
||||
root.RemoveCommand(dup)
|
||||
}
|
||||
}
|
||||
|
||||
func cacheStoreFromEnv() *cache.Store {
|
||||
cacheDir := strings.TrimSpace(os.Getenv(cli.CacheDirEnv))
|
||||
return cache.NewStore(cacheDir)
|
||||
@@ -721,7 +888,12 @@ func cleanCacheFiles(root, product string, staleOnly bool) (int, error) {
|
||||
return removed, nil
|
||||
}
|
||||
|
||||
// configureLogLevel sets the global slog level based on --debug and --verbose flags.
|
||||
// fileLogger holds the package-level file logger for diagnostics.
|
||||
// It is initialized by configureLogLevel and closed by CloseFileLogger.
|
||||
var fileLogger *logging.FileLogger
|
||||
|
||||
// configureLogLevel sets the global slog level based on --debug and --verbose flags
|
||||
// and initializes the file logger for diagnostics.
|
||||
// --debug → slog.LevelDebug; --verbose → slog.LevelInfo; default → slog.LevelWarn.
|
||||
func configureLogLevel(flags *GlobalFlags) {
|
||||
if flags == nil {
|
||||
@@ -736,7 +908,46 @@ func configureLogLevel(flags *GlobalFlags) {
|
||||
default:
|
||||
level = slog.LevelWarn
|
||||
}
|
||||
slog.SetDefault(slog.New(slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{
|
||||
Level: level,
|
||||
})))
|
||||
stderrHandler := slog.NewTextHandler(os.Stderr, &slog.HandlerOptions{Level: level})
|
||||
|
||||
// Initialize file logger — writes to ~/.dws/logs/dws.log at DEBUG level
|
||||
// regardless of stderr level. All slog calls are captured for diagnostics.
|
||||
fileLogger = logging.Setup(defaultConfigDir())
|
||||
fileHandler := slog.NewJSONHandler(fileLogger.Writer(), &slog.HandlerOptions{Level: slog.LevelDebug})
|
||||
|
||||
slog.SetDefault(slog.New(logging.NewMultiHandler(stderrHandler, fileHandler)))
|
||||
}
|
||||
|
||||
// FileLoggerInstance returns the package-level file logger, or nil if not initialized.
|
||||
func FileLoggerInstance() *slog.Logger {
|
||||
if fileLogger == nil {
|
||||
return nil
|
||||
}
|
||||
return fileLogger.Logger
|
||||
}
|
||||
|
||||
// CloseFileLogger flushes and closes the file logger.
|
||||
func CloseFileLogger() {
|
||||
if fileLogger != nil {
|
||||
fileLogger.Close()
|
||||
}
|
||||
}
|
||||
|
||||
// newPipelineEngine creates and configures the pipeline engine with
|
||||
// the standard set of handlers for model input correction.
|
||||
func newPipelineEngine() *pipeline.Engine {
|
||||
engine := pipeline.NewEngine()
|
||||
engine.RegisterAll(
|
||||
// PreParse handlers run in order: alias → sticky → paramname.
|
||||
// Alias normalises case first (--userId → --user-id), then
|
||||
// sticky splits glued values (--limit100 → --limit 100), then
|
||||
// paramname fixes near-miss typos (--limt → --limit).
|
||||
handlers.AliasHandler{},
|
||||
handlers.StickyHandler{},
|
||||
handlers.ParamNameHandler{},
|
||||
|
||||
// PostParse handlers normalise structured values.
|
||||
handlers.ParamValueHandler{},
|
||||
)
|
||||
return engine
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ import (
|
||||
mockmcp "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/test/mock_mcp"
|
||||
)
|
||||
|
||||
func TestPrintExecutionErrorDefaultsToHumanReadable(t *testing.T) {
|
||||
func TestPrintExecutionErrorDefaultsToJSON(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
root := NewRootCommand()
|
||||
@@ -43,14 +43,11 @@ func TestPrintExecutionErrorDefaultsToHumanReadable(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("printExecutionError() error = %v", err)
|
||||
}
|
||||
if stdout.Len() != 0 {
|
||||
t.Fatalf("stdout = %q, want empty for human-readable error output", stdout.String())
|
||||
if stderr.Len() != 0 {
|
||||
t.Fatalf("stderr = %q, want empty for JSON error output", stderr.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(stdout.String(), "\"category\": \"validation\"") {
|
||||
t.Fatalf("stdout = %q, want JSON error payload", stdout.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -175,8 +172,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 +213,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())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"strings"
|
||||
"text/tabwriter"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -53,7 +54,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
|
||||
}
|
||||
|
||||
+184
-21
@@ -22,6 +22,7 @@ import (
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
@@ -30,6 +31,7 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"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/edition"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -51,6 +53,7 @@ func newCommandRunnerWithFlags(loader cli.CatalogLoader, flags *GlobalFlags) exe
|
||||
}
|
||||
transportClient := transport.NewClient(httpClient)
|
||||
transportClient.ExtraHeaders = resolveIdentityHeaders()
|
||||
transportClient.FileLogger = FileLoggerInstance()
|
||||
return &runtimeRunner{
|
||||
loader: loader,
|
||||
transport: transportClient,
|
||||
@@ -73,6 +76,13 @@ type runtimeRunner struct {
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
totalStart := time.Now()
|
||||
defer func() {
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] runtimeRunner.Run total: %v\n", time.Since(totalStart))
|
||||
}
|
||||
}()
|
||||
|
||||
if r.loader == nil || r.transport == nil {
|
||||
return r.fallback.Run(ctx, invocation)
|
||||
}
|
||||
@@ -87,12 +97,14 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -116,7 +128,19 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
|
||||
tc := r.transport.WithAuth(r.resolveAuthToken(ctx), resolveIdentityHeaders())
|
||||
// Lazy bind FileLogger: it may be nil at construction time because
|
||||
// configureLogLevel runs later in PersistentPreRunE.
|
||||
if r.transport.FileLogger == nil {
|
||||
r.transport.FileLogger = FileLoggerInstance()
|
||||
}
|
||||
|
||||
authStart := time.Now()
|
||||
authToken := r.resolveAuthToken(ctx)
|
||||
authDuration := time.Since(authStart)
|
||||
RecordTiming(ctx, "auth_token", authDuration)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] resolveAuthToken: %v\n", authDuration)
|
||||
}
|
||||
|
||||
if invocation.DryRun {
|
||||
return executor.Result{
|
||||
@@ -147,19 +171,48 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// 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"),
|
||||
)
|
||||
}
|
||||
|
||||
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
|
||||
callStart := time.Now()
|
||||
callResult, err := tc.CallTool(ctx, endpoint, invocation.Tool, invocation.Params)
|
||||
callDuration := time.Since(callStart)
|
||||
RecordTiming(ctx, "mcp_call", callDuration)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] MCP CallTool: %v\n", callDuration)
|
||||
}
|
||||
if err != nil {
|
||||
if isAuthError(err) {
|
||||
if fn := edition.Get().OnAuthError; fn != nil {
|
||||
_ = fn(defaultConfigDir(), err)
|
||||
}
|
||||
}
|
||||
captureRuntimeFailure(invocation, err, err)
|
||||
return executor.Result{}, err
|
||||
}
|
||||
|
||||
if callResult.IsError {
|
||||
diag := transport.ExtractServerDiagnosticsFromMap(callResult.Content)
|
||||
logBusinessError(r.transport.FileLogger, "mcp_tool_error", invocation, callResult.Content, diag)
|
||||
mcpErr := apperrors.NewAPI(
|
||||
extractMCPErrorMessage(callResult),
|
||||
apperrors.WithOperation("tools/call"),
|
||||
apperrors.WithReason("mcp_tool_error"),
|
||||
apperrors.WithServerKey(invocation.CanonicalProduct),
|
||||
apperrors.WithHint("MCP tool returned a business error; check tool parameters and refer to skill documentation."),
|
||||
apperrors.WithServerDiag(diag),
|
||||
)
|
||||
captureRuntimeFailure(invocation, mcpErr, mcpErr)
|
||||
return executor.Result{}, mcpErr
|
||||
}
|
||||
|
||||
@@ -168,6 +221,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),
|
||||
@@ -191,25 +256,58 @@ 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() {
|
||||
loadDuration := time.Since(loadStart)
|
||||
RecordTiming(ctx, "keychain_load", loadDuration)
|
||||
if os.Getenv("DWS_PERF_DEBUG") != "" {
|
||||
_, _ = fmt.Fprintf(os.Stderr, "[PERF] getCachedRuntimeToken (first load): %v\n", loadDuration)
|
||||
}
|
||||
}()
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
cachedRuntimeToken = strings.TrimSpace(token)
|
||||
return
|
||||
}
|
||||
// If the error is a decryption failure (corrupted data), log and bail out
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return
|
||||
}
|
||||
// Try legacy manager as fallback
|
||||
manager := authpkg.NewManager(configDir, nil)
|
||||
configureLegacyAuthManagerCompatibility(manager)
|
||||
if token, _, err := manager.GetToken(); err == nil && strings.TrimSpace(token) != "" {
|
||||
cachedRuntimeToken = strings.TrimSpace(token)
|
||||
return
|
||||
}
|
||||
})
|
||||
return cachedRuntimeToken
|
||||
}
|
||||
|
||||
// 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 +341,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))
|
||||
@@ -273,9 +379,33 @@ func resolveIdentityHeaders() map[string]string {
|
||||
headers[k] = 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 +425,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...)
|
||||
}
|
||||
|
||||
@@ -45,7 +45,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 +90,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 +121,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 +255,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 +446,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 +496,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)
|
||||
@@ -628,7 +659,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,377 @@
|
||||
// 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"
|
||||
"net/http"
|
||||
"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/spf13/cobra"
|
||||
)
|
||||
|
||||
const (
|
||||
// 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"`
|
||||
}
|
||||
|
||||
// agentSkillPaths maps target names to their relative skill installation paths.
|
||||
// These paths are relative to the user's home directory.
|
||||
var agentSkillPaths = map[string]string{
|
||||
"qoder": ".qoder/skills",
|
||||
"claude": ".claude/skills",
|
||||
"cursor": ".cursor/skills",
|
||||
"codex": ".codex/skills",
|
||||
"opencode": filepath.Join(".config", "opencode", "skills"),
|
||||
}
|
||||
|
||||
// supportedTargets returns a comma-separated list of supported targets.
|
||||
func supportedTargets() string {
|
||||
targets := make([]string, 0, len(agentSkillPaths)+1)
|
||||
for target := range agentSkillPaths {
|
||||
targets = append(targets, target)
|
||||
}
|
||||
targets = append(targets, ".")
|
||||
return strings.Join(targets, ", ")
|
||||
}
|
||||
|
||||
func buildSkillCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "skill",
|
||||
Short: "技能管理",
|
||||
Long: "管理钉钉技能市场的技能。支持下载和安装技能到指定 Agent 目录。",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(newSkillAddCommand())
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newSkillAddCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "add <skillId> <target>",
|
||||
Short: "下载并安装技能到指定目录",
|
||||
Long: fmt.Sprintf(`从钉钉技能市场下载技能并安装到指定 Agent 目录。
|
||||
|
||||
参数:
|
||||
skillId 技能 ID(必填),可从钉钉技能市场获取
|
||||
target 安装目标(必填),支持: %s
|
||||
|
||||
安装路径:
|
||||
qoder -> ~/.qoder/skills/
|
||||
claude -> ~/.claude/skills/
|
||||
cursor -> ~/.cursor/skills/
|
||||
codex -> ~/.codex/skills/
|
||||
opencode -> ~/.config/opencode/skills/
|
||||
. -> 当前目录
|
||||
|
||||
示例:
|
||||
dws skill add skill-123 qoder # 安装到 ~/.qoder/skills/
|
||||
dws skill add skill-123 claude # 安装到 ~/.claude/skills/
|
||||
dws skill add skill-123 . # 安装到当前目录`, supportedTargets()),
|
||||
Args: cobra.ExactArgs(2),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: runSkillAdd,
|
||||
}
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
func 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()))
|
||||
}
|
||||
|
||||
// Load auth token
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
return apperrors.NewAuth("not logged in or token expired. Please run 'dws auth login' first",
|
||||
apperrors.WithHint("请先执行 'dws auth login' 登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
}
|
||||
|
||||
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, tokenData.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
|
||||
}
|
||||
|
||||
// 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, apperrors.NewAuth("authentication failed. Please run 'dws auth login' to refresh your token",
|
||||
apperrors.WithHint("请执行 'dws auth login' 重新登录"),
|
||||
apperrors.WithActions("dws auth login"))
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
// 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,732 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"archive/zip"
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
)
|
||||
|
||||
func TestResolveSkillTargetPath(t *testing.T) {
|
||||
homeDir, err := os.UserHomeDir()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get home directory: %v", err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
target string
|
||||
wantSuffix string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "qoder target",
|
||||
target: "qoder",
|
||||
wantSuffix: filepath.Join(".qoder", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "claude target",
|
||||
target: "claude",
|
||||
wantSuffix: filepath.Join(".claude", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "cursor target",
|
||||
target: "cursor",
|
||||
wantSuffix: filepath.Join(".cursor", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "codex target",
|
||||
target: "codex",
|
||||
wantSuffix: filepath.Join(".codex", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "opencode target",
|
||||
target: "opencode",
|
||||
wantSuffix: filepath.Join(".config", "opencode", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "case insensitive - QODER",
|
||||
target: "QODER",
|
||||
wantSuffix: filepath.Join(".qoder", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "case insensitive - Claude",
|
||||
target: "Claude",
|
||||
wantSuffix: filepath.Join(".claude", "skills"),
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid target",
|
||||
target: "invalid",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "empty target",
|
||||
target: "",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "whitespace only",
|
||||
target: " ",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := resolveSkillTargetPath(tt.target)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("resolveSkillTargetPath() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if !tt.wantErr {
|
||||
expected := filepath.Join(homeDir, tt.wantSuffix)
|
||||
if got != expected {
|
||||
t.Errorf("resolveSkillTargetPath() = %v, want %v", got, expected)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveSkillTargetPathCurrentDir(t *testing.T) {
|
||||
// Test "." target returns current working directory
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
t.Fatalf("failed to get working directory: %v", err)
|
||||
}
|
||||
|
||||
got, err := resolveSkillTargetPath(".")
|
||||
if err != nil {
|
||||
t.Errorf("resolveSkillTargetPath(\".\") error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if got != cwd {
|
||||
t.Errorf("resolveSkillTargetPath(\".\") = %v, want %v", got, cwd)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseDownloadSkillResponse(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
jsonInput string
|
||||
wantSuccess bool
|
||||
wantURL string
|
||||
wantFile string
|
||||
wantErrCode string
|
||||
wantErrMsg string
|
||||
}{
|
||||
{
|
||||
name: "successful response",
|
||||
jsonInput: `{
|
||||
"success": true,
|
||||
"result": {
|
||||
"downloadUrl": "https://example.com/skill.zip",
|
||||
"fileName": "my-skill.zip"
|
||||
}
|
||||
}`,
|
||||
wantSuccess: true,
|
||||
wantURL: "https://example.com/skill.zip",
|
||||
wantFile: "my-skill.zip",
|
||||
},
|
||||
{
|
||||
name: "error response",
|
||||
jsonInput: `{
|
||||
"success": false,
|
||||
"errorCode": "SKILL_NOT_FOUND",
|
||||
"errorMsg": "The skill does not exist"
|
||||
}`,
|
||||
wantSuccess: false,
|
||||
wantErrCode: "SKILL_NOT_FOUND",
|
||||
wantErrMsg: "The skill does not exist",
|
||||
},
|
||||
{
|
||||
name: "success with empty result",
|
||||
jsonInput: `{
|
||||
"success": true,
|
||||
"result": null
|
||||
}`,
|
||||
wantSuccess: true,
|
||||
wantURL: "",
|
||||
wantFile: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var resp downloadSkillResponse
|
||||
if err := json.Unmarshal([]byte(tt.jsonInput), &resp); err != nil {
|
||||
t.Fatalf("failed to unmarshal JSON: %v", err)
|
||||
}
|
||||
|
||||
if resp.Success != tt.wantSuccess {
|
||||
t.Errorf("Success = %v, want %v", resp.Success, tt.wantSuccess)
|
||||
}
|
||||
|
||||
if tt.wantSuccess && resp.Result != nil {
|
||||
if resp.Result.DownloadURL != tt.wantURL {
|
||||
t.Errorf("DownloadURL = %v, want %v", resp.Result.DownloadURL, tt.wantURL)
|
||||
}
|
||||
if resp.Result.FileName != tt.wantFile {
|
||||
t.Errorf("FileName = %v, want %v", resp.Result.FileName, tt.wantFile)
|
||||
}
|
||||
}
|
||||
|
||||
if !tt.wantSuccess {
|
||||
if resp.ErrorCode != tt.wantErrCode {
|
||||
t.Errorf("ErrorCode = %v, want %v", resp.ErrorCode, tt.wantErrCode)
|
||||
}
|
||||
if resp.ErrorMsg != tt.wantErrMsg {
|
||||
t.Errorf("ErrorMsg = %v, want %v", resp.ErrorMsg, tt.wantErrMsg)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZip(t *testing.T) {
|
||||
// Create a temporary zip file with test content
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "test.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create zip file with test content
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
|
||||
// Add a file to the zip
|
||||
fileContent := []byte("test content")
|
||||
writer, err := zipWriter.Create("test-file.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write(fileContent); err != nil {
|
||||
t.Fatalf("failed to write file content: %v", err)
|
||||
}
|
||||
|
||||
// Add a subdirectory with a file
|
||||
writer, err = zipWriter.Create("subdir/nested-file.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create nested file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write([]byte("nested content")); err != nil {
|
||||
t.Fatalf("failed to write nested file content: %v", err)
|
||||
}
|
||||
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract the zip
|
||||
if err := extractSkillZip(zipPath, destDir); err != nil {
|
||||
t.Fatalf("extractSkillZip() error = %v", err)
|
||||
}
|
||||
|
||||
// Verify extracted files
|
||||
extractedFile := filepath.Join(destDir, "test-file.txt")
|
||||
content, err := os.ReadFile(extractedFile)
|
||||
if err != nil {
|
||||
t.Errorf("failed to read extracted file: %v", err)
|
||||
}
|
||||
if string(content) != "test content" {
|
||||
t.Errorf("extracted content = %v, want %v", string(content), "test content")
|
||||
}
|
||||
|
||||
// Verify nested file
|
||||
nestedFile := filepath.Join(destDir, "subdir", "nested-file.txt")
|
||||
content, err = os.ReadFile(nestedFile)
|
||||
if err != nil {
|
||||
t.Errorf("failed to read nested file: %v", err)
|
||||
}
|
||||
if string(content) != "nested content" {
|
||||
t.Errorf("nested content = %v, want %v", string(content), "nested content")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSkillZipPreventZipSlip(t *testing.T) {
|
||||
tempDir := t.TempDir()
|
||||
zipPath := filepath.Join(tempDir, "malicious.zip")
|
||||
destDir := filepath.Join(tempDir, "extracted")
|
||||
|
||||
// Create a zip file with a path traversal attempt
|
||||
zipFile, err := os.Create(zipPath)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create zip file: %v", err)
|
||||
}
|
||||
|
||||
zipWriter := zip.NewWriter(zipFile)
|
||||
|
||||
// Try to create a file with path traversal
|
||||
writer, err := zipWriter.Create("../../../etc/passwd")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create malicious file in zip: %v", err)
|
||||
}
|
||||
if _, err := writer.Write([]byte("malicious content")); err != nil {
|
||||
t.Fatalf("failed to write malicious content: %v", err)
|
||||
}
|
||||
|
||||
if err := zipWriter.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip writer: %v", err)
|
||||
}
|
||||
if err := zipFile.Close(); err != nil {
|
||||
t.Fatalf("failed to close zip file: %v", err)
|
||||
}
|
||||
|
||||
// Extract should fail due to zip slip protection
|
||||
err = extractSkillZip(zipPath, destDir)
|
||||
if err == nil {
|
||||
t.Error("extractSkillZip() should have failed for zip slip attack")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "invalid file path") {
|
||||
t.Errorf("error should mention invalid file path, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandValidation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
args []string
|
||||
wantErr bool
|
||||
errMsg string
|
||||
}{
|
||||
{
|
||||
name: "missing arguments",
|
||||
args: []string{"skill", "add"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "missing target",
|
||||
args: []string{"skill", "add", "skill-123"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
{
|
||||
name: "too many arguments",
|
||||
args: []string{"skill", "add", "skill-123", "qoder", "extra"},
|
||||
wantErr: true,
|
||||
errMsg: "accepts 2 arg(s)",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs(tt.args)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Execute() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
if tt.wantErr && !strings.Contains(err.Error(), tt.errMsg) {
|
||||
t.Errorf("error = %v, should contain %v", err, tt.errMsg)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddInvalidTarget(t *testing.T) {
|
||||
// Setup: Create config directory with valid token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
// Save a valid token
|
||||
err := authpkg.SaveTokenData(configDir, &authpkg.TokenData{
|
||||
AccessToken: "test-token",
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("failed to save token data: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "invalid-target"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err = cmd.Execute()
|
||||
if err == nil {
|
||||
t.Error("Execute() should have failed for invalid target")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "invalid target") {
|
||||
t.Errorf("error should mention invalid target, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddRequiresAuth(t *testing.T) {
|
||||
// Setup: Create config directory without token
|
||||
tempDir := t.TempDir()
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
// Ensure the config directory exists but has no token
|
||||
if err := os.MkdirAll(configDir, 0755); err != nil {
|
||||
t.Fatalf("failed to create config dir: %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "skill-123", "qoder"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Error("Execute() should have failed without auth")
|
||||
}
|
||||
// Check for authentication-related error (English or Chinese)
|
||||
errStr := err.Error()
|
||||
if !strings.Contains(errStr, "not logged in") && !strings.Contains(errStr, "token") && !strings.Contains(errStr, "未登录") && !strings.Contains(errStr, "auth") {
|
||||
t.Errorf("error should mention authentication, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFetchSkillDownloadInfoUnauthorized(t *testing.T) {
|
||||
// Create mock server that returns 401
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// We can't easily test the actual fetchSkillDownloadInfo function
|
||||
// because it uses a hardcoded URL. This test verifies HTTP 401 handling pattern.
|
||||
client := &http.Client{Timeout: 5 * time.Second}
|
||||
resp, err := client.Get(server.URL)
|
||||
if err != nil {
|
||||
t.Fatalf("request failed: %v", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusUnauthorized {
|
||||
t.Errorf("expected 401, got %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSupportedTargets(t *testing.T) {
|
||||
targets := supportedTargets()
|
||||
|
||||
// Should contain all predefined targets
|
||||
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode", "."}
|
||||
for _, expected := range expectedTargets {
|
||||
if !strings.Contains(targets, expected) {
|
||||
t.Errorf("supportedTargets() should contain %s, got: %s", expected, targets)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAgentSkillPathsCrossPlatform(t *testing.T) {
|
||||
// Verify that paths use platform-appropriate separators
|
||||
for target, path := range agentSkillPaths {
|
||||
if runtime.GOOS == "windows" {
|
||||
if strings.Contains(path, "/") && !strings.Contains(path, "\\") {
|
||||
// On Windows, filepath.Join should use backslashes
|
||||
// But raw map values may use forward slashes
|
||||
t.Logf("Note: %s path '%s' uses forward slashes (will be converted by filepath.Join)", target, path)
|
||||
}
|
||||
}
|
||||
|
||||
// Test that resolveSkillTargetPath produces valid paths
|
||||
resolved, err := resolveSkillTargetPath(target)
|
||||
if err != nil {
|
||||
t.Errorf("resolveSkillTargetPath(%s) failed: %v", target, err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Path should be absolute
|
||||
if !filepath.IsAbs(resolved) {
|
||||
t.Errorf("resolveSkillTargetPath(%s) returned non-absolute path: %s", target, resolved)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupTempFile(t *testing.T) {
|
||||
// Create a temporary file
|
||||
tempFile, err := os.CreateTemp("", "test-cleanup-*.txt")
|
||||
if err != nil {
|
||||
t.Fatalf("failed to create temp file: %v", err)
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
tempFile.Close()
|
||||
|
||||
// Verify file exists
|
||||
if _, err := os.Stat(tempPath); os.IsNotExist(err) {
|
||||
t.Fatalf("temp file should exist before cleanup")
|
||||
}
|
||||
|
||||
// Clean up
|
||||
cleanupTempFile(tempPath)
|
||||
|
||||
// Verify file is deleted
|
||||
if _, err := os.Stat(tempPath); !os.IsNotExist(err) {
|
||||
t.Errorf("temp file should be deleted after cleanup")
|
||||
}
|
||||
|
||||
// Cleanup should not panic on empty path
|
||||
cleanupTempFile("")
|
||||
|
||||
// Cleanup should not panic on non-existent file
|
||||
cleanupTempFile("/nonexistent/path/file.txt")
|
||||
}
|
||||
|
||||
func TestDownloadSkillResponseJSON(t *testing.T) {
|
||||
// Test JSON marshaling/unmarshaling round-trip
|
||||
original := downloadSkillResponse{
|
||||
Success: true,
|
||||
Result: &downloadSkillResult{
|
||||
DownloadURL: "https://example.com/skill.zip",
|
||||
FileName: "skill.zip",
|
||||
},
|
||||
}
|
||||
|
||||
data, err := json.Marshal(original)
|
||||
if err != nil {
|
||||
t.Fatalf("failed to marshal: %v", err)
|
||||
}
|
||||
|
||||
var parsed downloadSkillResponse
|
||||
if err := json.Unmarshal(data, &parsed); err != nil {
|
||||
t.Fatalf("failed to unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if parsed.Success != original.Success {
|
||||
t.Errorf("Success mismatch: got %v, want %v", parsed.Success, original.Success)
|
||||
}
|
||||
if parsed.Result.DownloadURL != original.Result.DownloadURL {
|
||||
t.Errorf("DownloadURL mismatch: got %v, want %v", parsed.Result.DownloadURL, original.Result.DownloadURL)
|
||||
}
|
||||
if parsed.Result.FileName != original.Result.FileName {
|
||||
t.Errorf("FileName mismatch: got %v, want %v", parsed.Result.FileName, original.Result.FileName)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
output := out.String()
|
||||
// Check for the Long description which is shown in help
|
||||
if !strings.Contains(output, "技能") {
|
||||
t.Errorf("help should mention '技能', got: %s", output)
|
||||
}
|
||||
if !strings.Contains(output, "add") {
|
||||
t.Errorf("help should mention 'add' subcommand, got: %s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSkillAddCommandHelp(t *testing.T) {
|
||||
cmd := NewRootCommand()
|
||||
cmd.SetArgs([]string{"skill", "add", "--help"})
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
|
||||
output := out.String()
|
||||
// Should mention supported targets
|
||||
expectedTargets := []string{"qoder", "claude", "cursor", "codex", "opencode"}
|
||||
for _, target := range expectedTargets {
|
||||
if !strings.Contains(output, target) {
|
||||
t.Errorf("help should mention target '%s', got: %s", target, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func 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,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 (
|
||||
"context"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Environment variable to enable performance timing output.
|
||||
const PerfTimingEnv = "DWS_PERF_TIMING"
|
||||
|
||||
// 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
|
||||
}
|
||||
|
||||
// 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()
|
||||
if len(entries) == 0 {
|
||||
fmt.Fprintf(w, "\n[Timing] Total: %v (no detailed entries)\n", tc.Total().Truncate(time.Millisecond))
|
||||
return
|
||||
}
|
||||
|
||||
fmt.Fprintln(w)
|
||||
fmt.Fprintln(w, "[Timing] Execution breakdown:")
|
||||
for _, e := range entries {
|
||||
fmt.Fprintf(w, " %-30s %v\n", e.Name, e.Duration.Truncate(time.Millisecond))
|
||||
}
|
||||
fmt.Fprintf(w, " %-30s %v\n", "──────────────────────────────", "──────────")
|
||||
fmt.Fprintf(w, " %-30s %v\n", "Total", tc.Total().Truncate(time.Millisecond))
|
||||
}
|
||||
|
||||
// PrintIfEnabled prints timing info to stderr if DWS_PERF_TIMING is set.
|
||||
func (tc *TimingCollector) PrintIfEnabled() {
|
||||
if tc == nil {
|
||||
return
|
||||
}
|
||||
if os.Getenv(PerfTimingEnv) == "" {
|
||||
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)
|
||||
}
|
||||
|
||||
// IsPerfTimingEnabled returns true if performance timing output is enabled.
|
||||
func IsPerfTimingEnabled() bool {
|
||||
return os.Getenv(PerfTimingEnv) != ""
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
// 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"
|
||||
"os"
|
||||
"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, "[Timing]") {
|
||||
t.Error("output should contain [Timing] 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(PerfTimingEnv, "1")
|
||||
defer os.Unsetenv(PerfTimingEnv)
|
||||
|
||||
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) {
|
||||
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 TestIsPerfTimingEnabled(t *testing.T) {
|
||||
// Clear the env var first
|
||||
os.Unsetenv(PerfTimingEnv)
|
||||
|
||||
if IsPerfTimingEnabled() {
|
||||
t.Error("IsPerfTimingEnabled should return false when env var is not set")
|
||||
}
|
||||
|
||||
os.Setenv(PerfTimingEnv, "1")
|
||||
defer os.Unsetenv(PerfTimingEnv)
|
||||
|
||||
if !IsPerfTimingEnabled() {
|
||||
t.Error("IsPerfTimingEnabled should return true when env var is set")
|
||||
}
|
||||
}
|
||||
@@ -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,720 @@
|
||||
// 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()
|
||||
)
|
||||
|
||||
func newUpgradeCommand() *cobra.Command {
|
||||
var (
|
||||
flagCheck bool
|
||||
flagList bool
|
||||
flagVersion string
|
||||
flagRollback bool
|
||||
flagForce bool
|
||||
flagSkipSkills 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 --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 {
|
||||
return runUpgradeList(cmd, format)
|
||||
}
|
||||
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().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 ---
|
||||
|
||||
func runUpgradeList(cmd *cobra.Command, format string) error {
|
||||
client := upgrade.NewClient()
|
||||
|
||||
if format != "json" {
|
||||
fmt.Printf(" %s\n", ugDim("获取版本列表..."))
|
||||
}
|
||||
|
||||
versions, err := client.FetchAllReleases()
|
||||
if err != nil {
|
||||
return fmt.Errorf("获取版本列表失败: %w", err)
|
||||
}
|
||||
|
||||
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),
|
||||
})
|
||||
}
|
||||
return writeJSON(cmd.OutOrStdout(), map[string]any{
|
||||
"current_version": ensureV(version),
|
||||
"versions": items,
|
||||
})
|
||||
}
|
||||
|
||||
if len(versions) == 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)))
|
||||
fmt.Printf(" %s\n", ugDim("提示: 使用 dws upgrade --version v1.0.5 安装指定版本"))
|
||||
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,840 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// setupMCPConfigDir creates a temp config dir with mcp_url pointing to the
|
||||
// given test server and sets DWS_CONFIG_DIR via t.Setenv.
|
||||
// NOTE: tests calling this must NOT use t.Parallel().
|
||||
func setupMCPConfigDir(t *testing.T, srvURL string) string {
|
||||
t.Helper()
|
||||
dir := t.TempDir()
|
||||
os.WriteFile(filepath.Join(dir, "mcp_url"), []byte(srvURL), 0o600)
|
||||
t.Setenv("DWS_CONFIG_DIR", dir)
|
||||
return dir
|
||||
}
|
||||
|
||||
// resetClientIDFromMCP clears the MCP-sourced flag (test helper).
|
||||
func resetClientIDFromMCP() {
|
||||
clientMu.Lock()
|
||||
defer clientMu.Unlock()
|
||||
clientIDFromMCP = false
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 1. CheckCLIAuthEnabled: interface error → fail-closed with retry
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error from CheckCLIAuthEnabled when server returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Server 500 → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
configDir := setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 2 * time.Second},
|
||||
}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ Connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `{this is not valid json}`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ Malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Timeout_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(5 * time.Second)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{
|
||||
configDir: configDir,
|
||||
httpClient: &http.Client{Timeout: 200 * time.Millisecond},
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
_, err := p.CheckCLIAuthEnabled(ctx, "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error on request timeout, got nil")
|
||||
}
|
||||
t.Logf("✅ Timeout → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 2. CheckCLIAuthEnabled: transient error then recovery → succeeds
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !status.Success || !status.Result.CLIAuthEnabled {
|
||||
t.Fatalf("expected CLIAuthEnabled=true, got %+v", status)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ Transient error then success: attempts=%d, enabled=%v", calls.Load(), status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 3. CheckCLIAuthEnabled: normal responses (pass-through)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestCheckCLIAuthEnabled_Enabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=true")
|
||||
}
|
||||
t.Logf("✅ Normal enabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
func TestCheckCLIAuthEnabled_Disabled(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
p := &OAuthProvider{configDir: configDir, httpClient: srv.Client()}
|
||||
|
||||
status, err := p.CheckCLIAuthEnabled(context.Background(), "fake-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if status.Result.CLIAuthEnabled {
|
||||
t.Fatal("expected CLIAuthEnabled=false")
|
||||
}
|
||||
t.Logf("✅ Normal disabled response: success=%v, enabled=%v", status.Success, status.Result.CLIAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 4. OAuth callback: simulates the fail-closed logic at the /callback level
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestOAuthCallback_CLIAuthError_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error = fmt.Errorf("simulated network error")
|
||||
var authStatus *CLIAuthStatus
|
||||
_ = authStatus
|
||||
|
||||
// This is the exact expression used in oauth_provider.go callback:
|
||||
// cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
cliAuthEnabled := statusErr == nil // false
|
||||
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when statusErr != nil")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=%v → cliAuthEnabled=%v → shows notEnabledHTML (fail-closed)", statusErr, cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthEnabled_ShowsSuccessPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
if !cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be true when API returns enabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=true → cliAuthEnabled=%v → shows successHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
func TestOAuthCallback_CLIAuthDisabledByServer_ShowsNotEnabledPage(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var statusErr error
|
||||
authStatus := &CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
}
|
||||
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
if cliAuthEnabled {
|
||||
t.Fatal("cliAuthEnabled should be false when server says disabled")
|
||||
}
|
||||
t.Logf("✅ OAuth callback: statusErr=nil, enabled=false → cliAuthEnabled=%v → shows notEnabledHTML", cliAuthEnabled)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 5. Device Flow: loginOnce with broken cliAuthEnabled endpoint
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthError_FailClosed(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
|
||||
default:
|
||||
http.Error(w, "not found: "+r.URL.Path, http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth check fails, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "检查 CLI 授权状态失败") && !strings.Contains(err.Error(), "Failed to check CLI auth status") {
|
||||
t.Fatalf("unexpected error message: %s", err)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth check error → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthDisabled_ShowsError(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: false},
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, SuperAdminPath):
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "张三"}},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
_, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected loginOnce to fail when CLI auth is disabled, got nil")
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth disabled by server → login blocked: %s", err)
|
||||
}
|
||||
|
||||
func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
SetClientIDFromMCP("test-client-id")
|
||||
t.Cleanup(func() {
|
||||
SetClientID("")
|
||||
resetClientIDFromMCP()
|
||||
})
|
||||
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
|
||||
switch {
|
||||
case strings.HasSuffix(r.URL.Path, DeviceCodePath):
|
||||
writeServiceResult(w, true, DeviceAuthResponse{
|
||||
DeviceCode: "dc-test",
|
||||
UserCode: "TEST-CODE",
|
||||
VerificationURI: "https://example.com/verify",
|
||||
ExpiresIn: 60,
|
||||
Interval: 1,
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, DeviceTokenPath):
|
||||
writeServiceResult(w, true, DeviceTokenResponse{
|
||||
AuthCode: "test-auth-code",
|
||||
}, "", "")
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, MCPOAuthTokenPath):
|
||||
json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "test-access-token",
|
||||
"refreshToken": "test-refresh-token",
|
||||
"expiresIn": 7200,
|
||||
"corpId": "corp123",
|
||||
})
|
||||
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
json.NewEncoder(w).Encode(CLIAuthStatus{
|
||||
Success: true,
|
||||
Result: struct {
|
||||
CLIAuthEnabled bool `json:"cliAuthEnabled"`
|
||||
}{CLIAuthEnabled: true},
|
||||
})
|
||||
|
||||
default:
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
}
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
configDir := setupMCPConfigDir(t, srv.URL)
|
||||
provider := &DeviceFlowProvider{
|
||||
configDir: configDir,
|
||||
clientID: "test-client-id",
|
||||
scope: DefaultScopes,
|
||||
baseURL: srv.URL,
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
}
|
||||
|
||||
token, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected loginOnce to succeed, got error: %v", err)
|
||||
}
|
||||
if token.AccessToken != "test-access-token" {
|
||||
t.Fatalf("unexpected token: %s", token.AccessToken)
|
||||
}
|
||||
t.Logf("✅ Device Flow: CLI auth enabled → login succeeded, token=%s", token.AccessToken)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 6. FetchClientIDFromMCP: /cli/clientId error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestFetchClientIDFromMCP_ServerError_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/clientId returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId 500 → fail-closed with retry: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_ConnectionRefused_FailClosed(t *testing.T) {
|
||||
setupMCPConfigDir(t, "http://127.0.0.1:1")
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when connection is refused, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId connection refused → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_MalformedJSON_FailClosed(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
fmt.Fprint(w, `not json at all`)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error for malformed JSON, got nil")
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId malformed JSON → fail-closed: error=%q, attempts=%d", err, calls.Load())
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_BusinessError_FailClosed(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: false,
|
||||
ErrorCode: "forbidden",
|
||||
ErrorMsg: "access denied",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when server returns success=false, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "access denied") {
|
||||
t.Fatalf("error should contain server error message, got: %s", err)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId business error → fail-closed: error=%q", err)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "recovered-client-id",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if id != "recovered-client-id" {
|
||||
t.Fatalf("expected client ID 'recovered-client-id', got %q", id)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts (2 failures + 1 success), got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId transient then success: attempts=%d, id=%s", calls.Load(), id)
|
||||
}
|
||||
|
||||
func TestFetchClientIDFromMCP_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != ClientIDPath {
|
||||
http.Error(w, "not found", http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(ClientIDResponse{
|
||||
Success: true,
|
||||
Result: "my-client-id-123",
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
id, err := FetchClientIDFromMCP(context.Background())
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if id != "my-client-id-123" {
|
||||
t.Fatalf("expected 'my-client-id-123', got %q", id)
|
||||
}
|
||||
t.Logf("✅ /cli/clientId normal success: id=%s", id)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 7. GetSuperAdmins: /cli/superAdmin error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestGetSuperAdmins_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/superAdmin returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 2 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "a1", Name: "张三"}, {StaffID: "a2", Name: "李四"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "fake-token")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failures, got error: %v", err)
|
||||
}
|
||||
if !result.Success || len(result.Result) != 2 {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin transient then success: attempts=%d, admins=%v", calls.Load(), result.Result)
|
||||
}
|
||||
|
||||
func TestGetSuperAdmins_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SuperAdminResponse{
|
||||
Success: true,
|
||||
Result: []SuperAdmin{{StaffID: "admin1", Name: "王五"}},
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := GetSuperAdmins(context.Background(), "good-token")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if len(result.Result) != 1 || result.Result[0].Name != "王五" {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/superAdmin normal success: admins=%v", result.Result)
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// 8. SendCliAuthApply: /cli/sendCliAuthApply error handling
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendCliAuthApply_ServerError_RetriesAndFails(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
calls.Add(1)
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
_, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err == nil {
|
||||
t.Fatal("expected error when /cli/sendCliAuthApply returns 500, got nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "failed after 3 attempts") {
|
||||
t.Fatalf("error should mention retry exhaustion, got: %s", err)
|
||||
}
|
||||
if c := calls.Load(); c != 3 {
|
||||
t.Fatalf("expected 3 retry attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply 500 → retried 3 times: error=%q", err)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_TransientThenSuccess(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := calls.Add(1)
|
||||
if n <= 1 {
|
||||
http.Error(w, "Service Unavailable", http.StatusServiceUnavailable)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "admin1")
|
||||
|
||||
if err != nil {
|
||||
t.Fatalf("expected success after transient failure, got error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
if c := calls.Load(); c != 2 {
|
||||
t.Fatalf("expected 2 attempts, got %d", c)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply transient then success: attempts=%d", calls.Load())
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_Success(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("x-user-access-token") != "good-token" {
|
||||
t.Errorf("missing access token header")
|
||||
}
|
||||
if !strings.Contains(r.URL.RawQuery, "adminStaffId=admin123") {
|
||||
t.Errorf("missing or wrong adminStaffId param: %s", r.URL.RawQuery)
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{Success: true, Result: true})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "good-token", "admin123")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !result.Success || !result.Result {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply normal success: result=%+v", result)
|
||||
}
|
||||
|
||||
func TestSendCliAuthApply_BusinessError(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(SendApplyResponse{
|
||||
Success: false,
|
||||
ErrorCode: "invalid_admin",
|
||||
ErrorMsg: "admin not found",
|
||||
Result: false,
|
||||
})
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
setupMCPConfigDir(t, srv.URL)
|
||||
|
||||
result, err := SendCliAuthApply(context.Background(), "fake-token", "nonexistent")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected transport error: %v", err)
|
||||
}
|
||||
if result.Success {
|
||||
t.Fatal("expected success=false for business error")
|
||||
}
|
||||
t.Logf("✅ /cli/sendCliAuthApply business error: errorCode=%s, errorMsg=%s", result.ErrorCode, result.ErrorMsg)
|
||||
}
|
||||
@@ -26,8 +26,8 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/fatih/color"
|
||||
)
|
||||
|
||||
@@ -96,6 +96,23 @@ type serviceResult struct {
|
||||
}
|
||||
|
||||
func (p *DeviceFlowProvider) Login(ctx context.Context) (*TokenData, error) {
|
||||
// Ensure we have a valid client ID (fetch from MCP if not available)
|
||||
if p.clientID == "" {
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("client ID not configured, fetching from MCP server")
|
||||
}
|
||||
mcpClientID, mcpErr := FetchClientIDFromMCP(ctx)
|
||||
if mcpErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("获取 Client ID 失败"), mcpErr)
|
||||
}
|
||||
p.clientID = mcpClientID
|
||||
// Mark that clientID is from MCP
|
||||
SetClientIDFromMCP(mcpClientID)
|
||||
if p.logger != nil {
|
||||
p.logger.Debug("fetched client ID from MCP server", "clientID", mcpClientID)
|
||||
}
|
||||
}
|
||||
|
||||
const maxAttempts = 3
|
||||
for attempt := 1; attempt <= maxAttempts; attempt++ {
|
||||
tokenData, err := p.loginOnce(ctx, attempt)
|
||||
@@ -149,9 +166,54 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
|
||||
}
|
||||
|
||||
// Check if CLI auth is enabled for this organization (fail-closed: block on error)
|
||||
dfPrintStep(p.output(), 4, i18n.T("检查组织 CLI 授权状态..."), 0)
|
||||
authStatus, authErr := oauthProvider.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
|
||||
if authErr != nil {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 无法检查 CLI 数据访问权限状态")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请检查网络连接后重试。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("检查 CLI 授权状态失败"), authErr)
|
||||
} else if authStatus.Success && !authStatus.Result.CLIAuthEnabled {
|
||||
// CLI auth is disabled - show detailed error with admin info
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), dfRed(i18n.T("⚠️ 该组织尚未开启 CLI 数据访问权限")))
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
|
||||
// Try to get super admin list
|
||||
admins, adminErr := GetSuperAdmins(ctx, tokenData.AccessToken)
|
||||
if adminErr == nil && admins.Success && len(admins.Result) > 0 {
|
||||
// Show up to 3 admins
|
||||
maxAdmins := 3
|
||||
if len(admins.Result) < maxAdmins {
|
||||
maxAdmins = len(admins.Result)
|
||||
}
|
||||
var adminNames []string
|
||||
for i := 0; i < maxAdmins; i++ {
|
||||
adminNames = append(adminNames, admins.Result[i].Name)
|
||||
}
|
||||
_, _ = fmt.Fprintf(p.output(), " %s%s\n", i18n.T("组织主管理员:"), strings.Join(adminNames, "、"))
|
||||
}
|
||||
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 请联系组织主管理员开启后重新登录。"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T(" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings"))
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
return nil, errors.New(i18n.T("该组织尚未开启 CLI 数据访问权限,请联系管理员开启"))
|
||||
}
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := SaveTokenData(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
oauthProvider.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -46,6 +46,10 @@ func writeServiceResult(w http.ResponseWriter, success bool, result any, errCode
|
||||
func TestRequestDeviceCodeSuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Set a test client ID
|
||||
SetClientID("test-client-id")
|
||||
t.Cleanup(func() { SetClientID("") })
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
t.Fatalf("method = %s, want POST", r.Method)
|
||||
|
||||
+147
-5
@@ -15,6 +15,8 @@ package auth
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
@@ -58,16 +60,107 @@ const (
|
||||
|
||||
LogoutURL = "https://login.dingtalk.com/oauth2/logout"
|
||||
LogoutContinueURL = "https://login.dingtalk.com"
|
||||
|
||||
// MCP API endpoints for CLI authorization management.
|
||||
DefaultMCPBaseURL = "https://mcp.dingtalk.com"
|
||||
CLIAuthEnabledPath = "/cli/cliAuthEnabled"
|
||||
SuperAdminPath = "/cli/superAdmin"
|
||||
SendCliAuthApplyPath = "/cli/sendCliAuthApply"
|
||||
ClientIDPath = "/cli/clientId"
|
||||
|
||||
// MCP OAuth endpoints (used when clientId is fetched from MCP).
|
||||
MCPOAuthTokenPath = "/oauth2/getToken"
|
||||
MCPRefreshTokenPath = "/oauth2/refreshToken"
|
||||
MCPRevokeTokenPath = "/oauth2/revokeToken"
|
||||
)
|
||||
|
||||
// GetMCPBaseURL returns the MCP base URL with priority:
|
||||
// 1. ~/.dws/mcp_url file content (for pre-release environment)
|
||||
// 2. Default value (https://mcp.dingtalk.com)
|
||||
func GetMCPBaseURL() string {
|
||||
mcpURLPath := filepath.Join(getDefaultConfigDir(), "mcp_url")
|
||||
if data, err := os.ReadFile(mcpURLPath); err == nil {
|
||||
if url := strings.TrimSpace(string(data)); url != "" {
|
||||
return url
|
||||
}
|
||||
}
|
||||
return DefaultMCPBaseURL
|
||||
}
|
||||
|
||||
// Runtime overrides set via CLI flags (--client-id, --client-secret).
|
||||
// These take highest priority over environment variables and defaults.
|
||||
var (
|
||||
clientMu sync.RWMutex
|
||||
runtimeClientID string
|
||||
runtimeClientSecret string
|
||||
// clientIDFromMCP indicates whether the clientID was fetched from MCP server.
|
||||
// When true, MCP OAuth endpoints should be used instead of direct DingTalk API.
|
||||
clientIDFromMCP bool
|
||||
)
|
||||
|
||||
// SetClientIDFromMCP sets the clientID fetched from MCP server and marks it as MCP-sourced.
|
||||
func SetClientIDFromMCP(id string) {
|
||||
clientMu.Lock()
|
||||
defer clientMu.Unlock()
|
||||
runtimeClientID = id
|
||||
clientIDFromMCP = true
|
||||
}
|
||||
|
||||
// IsClientIDFromMCP returns true if the current clientID was fetched from MCP server.
|
||||
func IsClientIDFromMCP() bool {
|
||||
clientMu.RLock()
|
||||
defer clientMu.RUnlock()
|
||||
return clientIDFromMCP
|
||||
}
|
||||
|
||||
// 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 +177,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 +189,25 @@ func ClientID() string {
|
||||
if override != "" {
|
||||
return override
|
||||
}
|
||||
// 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 +215,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"
|
||||
|
||||
+1022
-12
File diff suppressed because it is too large
Load Diff
+358
-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,79 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
port := listener.Addr().(*net.TCPAddr).Port
|
||||
redirectURI := fmt.Sprintf("http://127.0.0.1:%d%s", port, CallbackPath)
|
||||
|
||||
codeCh := make(chan string, 1)
|
||||
// Channel to pass callback result (token data or error with CLI auth status)
|
||||
type callbackResult struct {
|
||||
token *TokenData
|
||||
err error
|
||||
cliAuthDisabled bool
|
||||
}
|
||||
resultCh := make(chan callbackResult, 1)
|
||||
errCh := make(chan error, 1)
|
||||
|
||||
// Shared state for API handlers (protected by mutex)
|
||||
var (
|
||||
callbackToken *TokenData
|
||||
callbackProcessedCode string // The auth code that has been successfully processed
|
||||
callbackAuthDisabled bool
|
||||
callbackApplySent bool // Whether apply request was sent
|
||||
callbackSelectedAdminId string // Selected admin ID for apply
|
||||
callbackCodeInProgress string // Code currently being processed (to prevent concurrent exchange)
|
||||
callbackTokenMu sync.Mutex
|
||||
)
|
||||
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc(CallbackPath, func(w http.ResponseWriter, r *http.Request) {
|
||||
// Get code first to check if this is a new authorization or page refresh
|
||||
code := r.URL.Query().Get("authCode")
|
||||
if code == "" {
|
||||
code = r.URL.Query().Get("code")
|
||||
}
|
||||
|
||||
// Check state and handle page refresh or concurrent requests
|
||||
callbackTokenMu.Lock()
|
||||
processedCode := callbackProcessedCode
|
||||
processedAuthDisabled := callbackAuthDisabled
|
||||
codeInProgress := callbackCodeInProgress
|
||||
hasToken := callbackToken != nil
|
||||
|
||||
// Case 1: This code was already successfully processed - show cached page
|
||||
if code != "" && code == processedCode {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if processedAuthDisabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Case 2: This code is being processed by another request - show wait page
|
||||
if code != "" && code == codeInProgress {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, `<html><head><meta http-equiv="refresh" content="1"></head><body><p>正在处理授权,请稍候...</p></body></html>`)
|
||||
return
|
||||
}
|
||||
|
||||
// Case 3: No code but we have a processed token - show cached page
|
||||
if code == "" && hasToken {
|
||||
callbackTokenMu.Unlock()
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if processedAuthDisabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Case 4: New code - mark as in-progress and process
|
||||
if code != "" {
|
||||
callbackCodeInProgress = code
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
if code == "" {
|
||||
select {
|
||||
case errCh <- errors.New(i18n.T("回调中未收到授权码")):
|
||||
@@ -117,14 +201,149 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
_, _ = fmt.Fprint(w, i18n.T("授权失败:未收到授权码"))
|
||||
return
|
||||
}
|
||||
select {
|
||||
case codeCh <- code:
|
||||
|
||||
// Exchange code for token
|
||||
tokenData, exchangeErr := p.exchangeCode(ctx, code)
|
||||
if exchangeErr != nil {
|
||||
// Clear in-progress state on error
|
||||
callbackTokenMu.Lock()
|
||||
if callbackCodeInProgress == code {
|
||||
callbackCodeInProgress = ""
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
default:
|
||||
// Select already exited (timeout/cancel); discard late callback.
|
||||
w.WriteHeader(http.StatusGone)
|
||||
_, _ = fmt.Fprintf(w, "<html><body><h1>授权失败</h1><p>%s</p></body></html>", exchangeErr.Error())
|
||||
select {
|
||||
case resultCh <- callbackResult{err: exchangeErr}:
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Mark as processed immediately after successful exchange
|
||||
callbackTokenMu.Lock()
|
||||
previouslyProcessed := callbackProcessedCode != ""
|
||||
callbackToken = tokenData
|
||||
callbackProcessedCode = code // Remember this code was successfully processed
|
||||
callbackCodeInProgress = "" // Clear in-progress state
|
||||
// Reset apply state for new authorization (user switched org)
|
||||
if previouslyProcessed {
|
||||
callbackApplySent = false
|
||||
callbackSelectedAdminId = ""
|
||||
}
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Check CLI auth enabled status (fail-closed: treat errors as disabled)
|
||||
authStatus, statusErr := p.CheckCLIAuthEnabled(ctx, tokenData.AccessToken)
|
||||
cliAuthEnabled := statusErr == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled
|
||||
|
||||
// Update CLI auth disabled state
|
||||
callbackTokenMu.Lock()
|
||||
callbackAuthDisabled = !cliAuthEnabled
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Display appropriate HTML based on CLI auth status
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if !cliAuthEnabled {
|
||||
_, _ = fmt.Fprint(w, notEnabledHTML)
|
||||
} else {
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
}
|
||||
// Ensure response is flushed to client
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
// Notify main goroutine with full result
|
||||
select {
|
||||
case resultCh <- callbackResult{token: tokenData, cliAuthDisabled: !cliAuthEnabled}:
|
||||
default:
|
||||
}
|
||||
})
|
||||
|
||||
// API endpoint: get super admins
|
||||
mux.HandleFunc("/api/superAdmin", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := GetSuperAdmins(ctx, token.AccessToken)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// API endpoint: send CLI auth apply
|
||||
mux.HandleFunc("/api/sendApply", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
adminStaffID := r.URL.Query().Get("adminStaffId")
|
||||
if adminStaffID == "" {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"缺少 adminStaffId 参数"}`))
|
||||
return
|
||||
}
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := SendCliAuthApply(ctx, token.AccessToken, adminStaffID)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
// Mark apply as sent and save selected admin on success
|
||||
if result.Success && result.Result {
|
||||
callbackTokenMu.Lock()
|
||||
callbackApplySent = true
|
||||
callbackSelectedAdminId = adminStaffID
|
||||
callbackTokenMu.Unlock()
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// API endpoint: get current status (clientId, applySent, selectedAdminId)
|
||||
mux.HandleFunc("/api/status", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
applySent := callbackApplySent
|
||||
selectedAdminId := callbackSelectedAdminId
|
||||
callbackTokenMu.Unlock()
|
||||
_, _ = fmt.Fprintf(w, `{"clientId":"%s","applySent":%t,"selectedAdminId":"%s"}`, p.clientID, applySent, selectedAdminId)
|
||||
})
|
||||
|
||||
// API endpoint: check CLI auth enabled status
|
||||
mux.HandleFunc("/api/cliAuthEnabled", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
callbackTokenMu.Lock()
|
||||
token := callbackToken
|
||||
callbackTokenMu.Unlock()
|
||||
if token == nil {
|
||||
_, _ = w.Write([]byte(`{"success":false,"errorMsg":"授权尚未完成"}`))
|
||||
return
|
||||
}
|
||||
result, err := p.CheckCLIAuthEnabled(ctx, token.AccessToken)
|
||||
if err != nil {
|
||||
_, _ = fmt.Fprintf(w, `{"success":false,"errorMsg":"%s"}`, err.Error())
|
||||
return
|
||||
}
|
||||
data, _ := json.Marshal(result)
|
||||
_, _ = w.Write(data)
|
||||
})
|
||||
|
||||
// Success page endpoint
|
||||
mux.HandleFunc("/success", func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprint(w, successHTML)
|
||||
})
|
||||
|
||||
server := &http.Server{Handler: mux}
|
||||
@@ -160,9 +379,9 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
timeout := time.NewTimer(5 * time.Minute)
|
||||
defer timeout.Stop()
|
||||
|
||||
var authCode string
|
||||
var result callbackResult
|
||||
select {
|
||||
case authCode = <-codeCh:
|
||||
case result = <-resultCh:
|
||||
case err := <-errCh:
|
||||
return nil, err
|
||||
case <-timeout.C:
|
||||
@@ -171,13 +390,82 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
|
||||
tokenData, err := p.exchangeCode(ctx, authCode)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
|
||||
// Handle callback errors
|
||||
if result.err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), result.err)
|
||||
}
|
||||
|
||||
// Handle CLI auth disabled - keep server running for user to apply
|
||||
if result.cliAuthDisabled {
|
||||
_, _ = fmt.Fprintln(p.output(), "")
|
||||
_, _ = fmt.Fprintln(p.output(), i18n.T("⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请..."))
|
||||
|
||||
// Poll for CLI auth status while waiting
|
||||
applyTimeout := time.NewTimer(10 * time.Minute)
|
||||
defer applyTimeout.Stop()
|
||||
pollTicker := time.NewTicker(5 * time.Second)
|
||||
defer pollTicker.Stop()
|
||||
|
||||
elapsedSeconds := 0
|
||||
for {
|
||||
select {
|
||||
case <-applyTimeout.C:
|
||||
return nil, errors.New(i18n.T("操作超时,请重新登录"))
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-pollTicker.C:
|
||||
elapsedSeconds += 5
|
||||
|
||||
// Get latest token and state (user may have switched org)
|
||||
callbackTokenMu.Lock()
|
||||
currentToken := callbackToken
|
||||
currentAuthDisabled := callbackAuthDisabled
|
||||
applySent := callbackApplySent
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
// Check if user switched to an org with CLI auth enabled
|
||||
if currentToken != nil && !currentAuthDisabled {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
|
||||
time.Sleep(2 * time.Second)
|
||||
result.token = currentToken
|
||||
result.cliAuthDisabled = false
|
||||
goto continueLogin
|
||||
}
|
||||
|
||||
// Check if CLI auth is now enabled (admin approved)
|
||||
if currentToken != nil {
|
||||
authStatus, err := p.CheckCLIAuthEnabled(ctx, currentToken.AccessToken)
|
||||
if err == nil && authStatus.Success && authStatus.Result.CLIAuthEnabled {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r%s\n", i18n.T("✅ 权限已开启,继续登录..."))
|
||||
time.Sleep(2 * time.Second)
|
||||
result.token = currentToken
|
||||
result.cliAuthDisabled = false
|
||||
goto continueLogin
|
||||
}
|
||||
}
|
||||
|
||||
// Show polling status based on apply state
|
||||
if applySent {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待管理员审批中"), elapsedSeconds)
|
||||
} else {
|
||||
_, _ = fmt.Fprintf(p.output(), "\r⏳ %s (%ds/600s) ", i18n.T("等待提交申请中"), elapsedSeconds)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
continueLogin:
|
||||
tokenData := result.token
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := SaveTokenData(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
|
||||
// Persist app credentials if using custom client credentials
|
||||
p.persistAppConfigIfNeeded()
|
||||
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
@@ -209,19 +497,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 +532,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 +547,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 +577,30 @@ func (p *OAuthProvider) Logout() error {
|
||||
func (p *OAuthProvider) Status() (*TokenData, error) {
|
||||
return LoadTokenData(p.configDir)
|
||||
}
|
||||
|
||||
// persistAppConfigIfNeeded saves app credentials if custom ones were used.
|
||||
// This ensures the client secret is available for future token refreshes.
|
||||
func (p *OAuthProvider) persistAppConfigIfNeeded() {
|
||||
// Check if custom credentials were provided via runtime flags
|
||||
clientID, clientSecret := getRuntimeCredentials()
|
||||
if clientID == "" || clientSecret == "" {
|
||||
return
|
||||
}
|
||||
|
||||
// Only persist if they differ from environment/default values
|
||||
envID := getEnvClientID()
|
||||
if clientID == envID || clientID == DefaultClientID {
|
||||
return
|
||||
}
|
||||
|
||||
// Save app config with secret stored in keychain
|
||||
config := &AppConfig{
|
||||
ClientID: clientID,
|
||||
ClientSecret: PlainSecret(clientSecret),
|
||||
}
|
||||
if err := SaveAppConfig(p.configDir, config); err != nil {
|
||||
if p.logger != nil {
|
||||
p.logger.Warn("failed to persist app credentials", "error", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,147 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
const (
|
||||
// secretKeyPrefix is the keychain account prefix for app secrets.
|
||||
secretKeyPrefix = "appsecret:"
|
||||
)
|
||||
|
||||
// SecretRef references a secret stored externally.
|
||||
type SecretRef struct {
|
||||
Source string `json:"source"` // "keychain" | "file"
|
||||
ID string `json:"id"` // keychain key or file path
|
||||
}
|
||||
|
||||
// SecretInput represents a secret value: either a plain string or a SecretRef object.
|
||||
type SecretInput struct {
|
||||
Plain string // non-empty for plain string values
|
||||
Ref *SecretRef // non-nil for SecretRef values
|
||||
}
|
||||
|
||||
// PlainSecret creates a SecretInput from a plain string.
|
||||
func PlainSecret(s string) SecretInput {
|
||||
return SecretInput{Plain: s}
|
||||
}
|
||||
|
||||
// IsZero returns true if the SecretInput has no value.
|
||||
func (s SecretInput) IsZero() bool {
|
||||
return s.Plain == "" && s.Ref == nil
|
||||
}
|
||||
|
||||
// IsSecretRef returns true if this is a SecretRef object.
|
||||
func (s SecretInput) IsSecretRef() bool {
|
||||
return s.Ref != nil
|
||||
}
|
||||
|
||||
// IsPlain returns true if this is a plain text string (not a SecretRef).
|
||||
func (s SecretInput) IsPlain() bool {
|
||||
return s.Ref == nil && s.Plain != ""
|
||||
}
|
||||
|
||||
// MarshalJSON serializes SecretInput: plain string → JSON string, SecretRef → JSON object.
|
||||
func (s SecretInput) MarshalJSON() ([]byte, error) {
|
||||
if s.Ref != nil {
|
||||
return json.Marshal(s.Ref)
|
||||
}
|
||||
return json.Marshal(s.Plain)
|
||||
}
|
||||
|
||||
// UnmarshalJSON deserializes SecretInput from either a JSON string or a SecretRef object.
|
||||
func (s *SecretInput) UnmarshalJSON(data []byte) error {
|
||||
// Try string first
|
||||
var plain string
|
||||
if err := json.Unmarshal(data, &plain); err == nil {
|
||||
s.Plain = plain
|
||||
s.Ref = nil
|
||||
return nil
|
||||
}
|
||||
// Try SecretRef object
|
||||
var ref SecretRef
|
||||
if err := json.Unmarshal(data, &ref); err == nil && isValidSource(ref.Source) && ref.ID != "" {
|
||||
s.Ref = &ref
|
||||
s.Plain = ""
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("clientSecret must be a string or {source, id} object")
|
||||
}
|
||||
|
||||
// ValidSecretSources is the set of recognized SecretRef sources.
|
||||
var ValidSecretSources = map[string]bool{
|
||||
"file": true, "keychain": true,
|
||||
}
|
||||
|
||||
func isValidSource(source string) bool {
|
||||
return ValidSecretSources[source]
|
||||
}
|
||||
|
||||
// secretAccountKey generates the keychain account key for an app's secret.
|
||||
func secretAccountKey(clientID string) string {
|
||||
return secretKeyPrefix + clientID
|
||||
}
|
||||
|
||||
// ResolveSecret resolves a SecretInput to a plain string.
|
||||
// SecretRef objects are resolved by source (file / keychain).
|
||||
func ResolveSecret(input SecretInput) (string, error) {
|
||||
if input.Ref == nil {
|
||||
return input.Plain, nil
|
||||
}
|
||||
switch input.Ref.Source {
|
||||
case "file":
|
||||
data, err := os.ReadFile(input.Ref.ID)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to read secret file %s: %w", input.Ref.ID, err)
|
||||
}
|
||||
return strings.TrimSpace(string(data)), nil
|
||||
case "keychain":
|
||||
val, err := keychain.Get(keychain.Service, input.Ref.ID)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("failed to get secret from keychain: %w", err)
|
||||
}
|
||||
return val, nil
|
||||
default:
|
||||
return "", fmt.Errorf("unknown secret source: %s", input.Ref.Source)
|
||||
}
|
||||
}
|
||||
|
||||
// StoreSecret stores a plain text secret in keychain and returns a SecretRef.
|
||||
// If the input is already a SecretRef, it is returned as-is.
|
||||
// Returns error if keychain is unavailable.
|
||||
func StoreSecret(clientID string, input SecretInput) (SecretInput, error) {
|
||||
if !input.IsPlain() {
|
||||
return input, nil // SecretRef → keep as-is
|
||||
}
|
||||
key := secretAccountKey(clientID)
|
||||
if err := keychain.Set(keychain.Service, key, input.Plain); err != nil {
|
||||
return SecretInput{}, fmt.Errorf("keychain unavailable: %w\nhint: use file reference in config to bypass keychain", err)
|
||||
}
|
||||
return SecretInput{Ref: &SecretRef{Source: "keychain", ID: key}}, nil
|
||||
}
|
||||
|
||||
// RemoveSecretStore cleans up keychain entries when an app is removed.
|
||||
// Errors are intentionally ignored — cleanup is best-effort.
|
||||
func RemoveSecretStore(input SecretInput) {
|
||||
if input.IsSecretRef() && input.Ref.Source == "keychain" {
|
||||
_ = keychain.Remove(keychain.Service, input.Ref.ID)
|
||||
}
|
||||
}
|
||||
@@ -21,8 +21,8 @@ import (
|
||||
"path/filepath"
|
||||
"sync"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/security"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
)
|
||||
|
||||
const secureDataFile = ".data"
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+87
-8
@@ -14,7 +14,9 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
@@ -32,6 +34,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 +61,60 @@ 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.
|
||||
// SaveTokenData saves TokenData to the platform keychain.
|
||||
// Uses the new keychain-based storage with random master key for better security.
|
||||
func SaveTokenData(configDir string, data *TokenData) error {
|
||||
return SaveSecureTokenData(configDir, data)
|
||||
return SaveTokenDataKeychain(data)
|
||||
}
|
||||
|
||||
// LoadTokenData reads TokenData from encrypted .data file.
|
||||
// LoadTokenData reads TokenData from the platform keychain.
|
||||
// On first call, it attempts to migrate legacy .data file if present.
|
||||
func LoadTokenData(configDir string) (*TokenData, error) {
|
||||
return LoadSecureTokenData(configDir)
|
||||
// Try loading from new keychain first
|
||||
if TokenDataExistsKeychain() {
|
||||
return LoadTokenDataKeychain()
|
||||
}
|
||||
|
||||
// Fallback: try legacy .data file and migrate
|
||||
data, err := LoadSecureTokenData(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Migrate to keychain for future use
|
||||
if err := SaveTokenDataKeychain(data); err == nil {
|
||||
// Successfully migrated, delete legacy file
|
||||
_ = DeleteSecureData(configDir)
|
||||
}
|
||||
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes encrypted .data file from configDir.
|
||||
// DeleteTokenData removes token data from both keychain and legacy storage.
|
||||
func DeleteTokenData(configDir string) error {
|
||||
return DeleteSecureData(configDir)
|
||||
// Delete from keychain
|
||||
keychainErr := DeleteTokenDataKeychain()
|
||||
|
||||
// Also clean up any legacy .data file
|
||||
legacyErr := DeleteSecureData(configDir)
|
||||
|
||||
// Return keychain error if any, otherwise legacy error
|
||||
if keychainErr != nil {
|
||||
return keychainErr
|
||||
}
|
||||
return legacyErr
|
||||
}
|
||||
|
||||
// RevokeTokenRemote calls the DingTalk logout endpoint to invalidate the access token.
|
||||
// RevokeTokenRemote calls the appropriate logout/revoke endpoint to invalidate the access token.
|
||||
// Uses MCP revoke endpoint when clientID is from MCP, otherwise uses DingTalk logout.
|
||||
// This should be called before deleting local token data.
|
||||
// The function is best-effort: errors are returned but callers may choose to ignore them.
|
||||
func RevokeTokenRemote(ctx context.Context) error {
|
||||
// Use MCP revoke endpoint when clientID is from MCP
|
||||
if IsClientIDFromMCP() {
|
||||
return revokeTokenViaMCP(ctx)
|
||||
}
|
||||
// Direct mode: use DingTalk logout endpoint
|
||||
logoutURL, err := url.Parse(LogoutURL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parsing logout URL: %w", err)
|
||||
@@ -114,3 +151,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 {
|
||||
|
||||
+167
-42
@@ -25,11 +25,12 @@ import (
|
||||
"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 +57,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,17 +95,27 @@ 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 {
|
||||
@@ -113,24 +124,20 @@ func NewSchemaCommand(loader CatalogLoader) *cobra.Command {
|
||||
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 +176,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 +211,12 @@ 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))
|
||||
}
|
||||
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 +267,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 +276,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 +310,125 @@ 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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
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 +488,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 +506,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 +585,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 +610,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 +655,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 {
|
||||
|
||||
+530
-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,524 @@ 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
|
||||
}
|
||||
|
||||
type errorLoader struct {
|
||||
err error
|
||||
}
|
||||
|
||||
+25
-1
@@ -22,11 +22,11 @@ import (
|
||||
"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"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -50,6 +50,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 +89,9 @@ 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
|
||||
}
|
||||
|
||||
type cachedCatalogState struct {
|
||||
@@ -109,6 +128,11 @@ func (l EnvironmentLoader) Load(ctx context.Context) (ir.Catalog, error) {
|
||||
}
|
||||
|
||||
transportClient := transport.NewClient(nil)
|
||||
if l.AuthTokenFunc != nil {
|
||||
if token := l.AuthTokenFunc(ctx); token != "" {
|
||||
transportClient = transportClient.WithAuth(token, nil)
|
||||
}
|
||||
}
|
||||
|
||||
// Use a bounded context so discovery doesn't hang in test or CI environments.
|
||||
timeout := defaultDiscoveryTimeout
|
||||
|
||||
@@ -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
|
||||
@@ -136,7 +141,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,10 +26,10 @@ 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"
|
||||
)
|
||||
|
||||
@@ -143,7 +145,20 @@ func NewDirectCommand(route Route, runner executor.Runner) *cobra.Command {
|
||||
}
|
||||
}
|
||||
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(
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
+100
-36
@@ -14,13 +14,12 @@
|
||||
package errors
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"bytes"
|
||||
)
|
||||
|
||||
// Category represents a stable error class with a documented exit code.
|
||||
@@ -36,18 +35,19 @@ const (
|
||||
|
||||
// Error is the structured repository-local error model for the Go rewrite.
|
||||
type Error struct {
|
||||
Category Category
|
||||
Message string
|
||||
Operation string
|
||||
ServerKey string
|
||||
Retryable bool
|
||||
Reason string
|
||||
Hint string
|
||||
Actions []string
|
||||
Snapshot string
|
||||
RPCCode int `json:"rpc_code,omitempty"`
|
||||
RPCData json.RawMessage `json:"rpc_data,omitempty"`
|
||||
Cause error `json:"-"`
|
||||
Category Category
|
||||
Message string
|
||||
Operation string
|
||||
ServerKey string
|
||||
Retryable bool
|
||||
Reason string
|
||||
Hint string
|
||||
Actions []string
|
||||
Snapshot string
|
||||
RPCCode int `json:"rpc_code,omitempty"`
|
||||
RPCData json.RawMessage `json:"rpc_data,omitempty"`
|
||||
ServerDiag ServerDiagnostics `json:"-"`
|
||||
Cause error `json:"-"`
|
||||
}
|
||||
|
||||
func (e *Error) Error() string {
|
||||
@@ -247,6 +247,23 @@ func PrintJSON(w io.Writer, err error) error {
|
||||
errorPayload["rpc_data"] = parsed
|
||||
}
|
||||
}
|
||||
if !typed.ServerDiag.IsEmpty() {
|
||||
if typed.ServerDiag.TraceID != "" {
|
||||
errorPayload["trace_id"] = typed.ServerDiag.TraceID
|
||||
}
|
||||
if typed.ServerDiag.ServerErrorCode != "" {
|
||||
errorPayload["server_error_code"] = typed.ServerDiag.ServerErrorCode
|
||||
// Add user-friendly hint for specific server error codes
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
errorPayload["friendly_hint"] = "该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。"
|
||||
errorPayload["action_url"] = "https://open-dev.dingtalk.com/fe/old#/developerSettings"
|
||||
}
|
||||
}
|
||||
if typed.ServerDiag.TechnicalDetail != "" {
|
||||
errorPayload["technical_detail"] = typed.ServerDiag.TechnicalDetail
|
||||
}
|
||||
}
|
||||
if typed.Cause != nil {
|
||||
errorPayload["cause"] = typed.Cause.Error()
|
||||
}
|
||||
@@ -263,8 +280,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 +309,23 @@ func PrintHuman(w io.Writer, err error) error {
|
||||
return writeErr
|
||||
}
|
||||
|
||||
// Line 1: Error summary
|
||||
lines := []string{
|
||||
fmt.Sprintf("Error: [%s] %s", strings.ToUpper(string(typed.Category)), typed.Message),
|
||||
}
|
||||
if typed.Reason != "" {
|
||||
lines = append(lines, fmt.Sprintf("Reason: %s", typed.Reason))
|
||||
}
|
||||
if typed.Operation != "" {
|
||||
lines = append(lines, fmt.Sprintf("Operation: %s", typed.Operation))
|
||||
}
|
||||
if typed.ServerKey != "" {
|
||||
lines = append(lines, fmt.Sprintf("Server: %s", typed.ServerKey))
|
||||
}
|
||||
|
||||
// Always shown: hint, actions, retryable
|
||||
if typed.Hint != "" {
|
||||
lines = append(lines, fmt.Sprintf("Hint: %s", typed.Hint))
|
||||
}
|
||||
|
||||
// Add user-friendly hint for specific server error codes
|
||||
switch typed.ServerDiag.ServerErrorCode {
|
||||
case "TOKEN_VERIFIED_FAILED", "CLI_ORG_NOT_AUTHORIZED":
|
||||
lines = append(lines, "Hint: 该组织尚未开启 CLI 数据访问权限,请联系组织主管理员开启。")
|
||||
lines = append(lines, "Action: 开启地址: https://open-dev.dingtalk.com/fe/old#/developerSettings")
|
||||
}
|
||||
|
||||
if len(typed.Actions) > 0 {
|
||||
for _, action := range typed.Actions {
|
||||
if strings.TrimSpace(action) == "" {
|
||||
@@ -298,22 +334,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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,6 +15,8 @@ package errors
|
||||
|
||||
import (
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"unicode"
|
||||
@@ -44,6 +46,57 @@ func ResourceName(name string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// isDangerousUnicode identifies Unicode code points used for visual spoofing attacks.
|
||||
// These characters are invisible or alter text direction, allowing attackers to make
|
||||
// "report.exe" display as "report.txt" (Bidi override) or insert hidden content
|
||||
// (zero-width characters).
|
||||
func isDangerousUnicode(r rune) bool {
|
||||
switch {
|
||||
case r >= 0x200B && r <= 0x200D: // zero-width space/non-joiner/joiner
|
||||
return true
|
||||
case r == 0xFEFF: // BOM / ZWNBSP
|
||||
return true
|
||||
case r >= 0x202A && r <= 0x202E: // Bidi: LRE/RLE/PDF/LRO/RLO
|
||||
return true
|
||||
case r >= 0x2028 && r <= 0x2029: // line/paragraph separator
|
||||
return true
|
||||
case r >= 0x2066 && r <= 0x2069: // Bidi isolates: LRI/RLI/FSI/PDI
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// rejectControlChars rejects control characters in a string.
|
||||
// Rejects C0 control characters (except \t and \n) and dangerous Unicode.
|
||||
// Tab and newline are allowed as they may appear in legitimate multi-line input.
|
||||
func rejectControlChars(s, fieldName string) error {
|
||||
for _, r := range s {
|
||||
// Allow tab (\t = 0x09) and newline (\n = 0x0A)
|
||||
if r != '\t' && r != '\n' && (r < 0x20 || r == 0x7f) {
|
||||
return fmt.Errorf("%s contains control characters", fieldName)
|
||||
}
|
||||
if isDangerousUnicode(r) {
|
||||
return fmt.Errorf("%s contains dangerous Unicode characters", fieldName)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RejectControlChars rejects C0 control characters (except \t and \n) and
|
||||
// dangerous Unicode characters from user input.
|
||||
//
|
||||
// Control characters cause subtle security issues:
|
||||
// - Null bytes truncate strings at the C layer
|
||||
// - \r\n enables HTTP header injection
|
||||
// - Unicode Bidi characters allow visual spoofing (e.g. making "report.exe" display as "report.txt")
|
||||
//
|
||||
// Tab and newline are allowed as they may appear in legitimate multi-line input.
|
||||
func RejectControlChars(value, flagName string) error {
|
||||
return rejectControlChars(value, flagName)
|
||||
}
|
||||
|
||||
// SafePath performs basic path validation checking for dangerous patterns.
|
||||
// For full security (symlink resolution, CWD containment), use SafeOutputPath or SafeInputPath.
|
||||
func SafePath(path string) error {
|
||||
if path == "" {
|
||||
return stderrors.New("path cannot be empty")
|
||||
@@ -57,6 +110,13 @@ func SafePath(path string) error {
|
||||
return ErrUnsafePath
|
||||
}
|
||||
|
||||
// Check for dangerous Unicode
|
||||
for _, r := range path {
|
||||
if isDangerousUnicode(r) {
|
||||
return ErrUnsafePath
|
||||
}
|
||||
}
|
||||
|
||||
lowerPath := strings.ToLower(path)
|
||||
for _, pattern := range []string{
|
||||
"..",
|
||||
@@ -77,3 +137,143 @@ func SafePath(path string) error {
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SafeOutputPath validates a download/export target path for --output flags.
|
||||
// It rejects absolute paths, resolves symlinks to their real location, and
|
||||
// verifies the canonical result is still under the current working directory.
|
||||
// This prevents an AI Agent from being tricked into writing files outside the
|
||||
// working directory (e.g. "../../.ssh/authorized_keys") or following symlinks
|
||||
// to sensitive locations.
|
||||
//
|
||||
// The returned absolute path MUST be used for all subsequent I/O to prevent
|
||||
// time-of-check-to-time-of-use (TOCTOU) race conditions.
|
||||
func SafeOutputPath(path string) (string, error) {
|
||||
return safePath(path, "--output")
|
||||
}
|
||||
|
||||
// SafeInputPath validates an upload/read source path for --file flags.
|
||||
// It applies the same rules as SafeOutputPath — rejecting absolute paths,
|
||||
// resolving symlinks, and enforcing working directory containment — to prevent
|
||||
// an AI Agent from being tricked into reading sensitive files like /etc/passwd.
|
||||
func SafeInputPath(path string) (string, error) {
|
||||
return safePath(path, "--file")
|
||||
}
|
||||
|
||||
// safePath is the shared implementation for SafeOutputPath and SafeInputPath.
|
||||
func safePath(raw, flagName string) (string, error) {
|
||||
if err := rejectControlChars(raw, flagName); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
path := filepath.Clean(raw)
|
||||
|
||||
// Reject absolute paths - force relative paths within CWD
|
||||
if filepath.IsAbs(path) {
|
||||
return "", fmt.Errorf("%s must be a relative path within the current directory, got %q (hint: cd to the target directory first, or use a relative path like ./filename)", flagName, raw)
|
||||
}
|
||||
|
||||
cwd, err := os.Getwd()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot determine working directory: %w", err)
|
||||
}
|
||||
resolved := filepath.Join(cwd, path)
|
||||
|
||||
// Resolve symlinks: for existing paths, follow to real location;
|
||||
// for non-existing paths, walk up to the nearest existing ancestor,
|
||||
// resolve its symlinks, and re-attach the remaining tail segments.
|
||||
// This prevents TOCTOU attacks where a non-existent intermediate
|
||||
// directory is replaced with a symlink between check and use.
|
||||
if _, err := os.Lstat(resolved); err == nil {
|
||||
resolved, err = filepath.EvalSymlinks(resolved)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot resolve symlinks: %w", err)
|
||||
}
|
||||
} else {
|
||||
resolved, err = resolveNearestAncestor(resolved)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("cannot resolve symlinks: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
canonicalCwd, _ := filepath.EvalSymlinks(cwd)
|
||||
if !isUnderDir(resolved, canonicalCwd) {
|
||||
return "", fmt.Errorf("%s %q resolves outside the current working directory (hint: the path must stay within the working directory after resolving .. and symlinks)", flagName, raw)
|
||||
}
|
||||
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
// resolveNearestAncestor walks up from path until it finds an existing
|
||||
// ancestor, resolves that ancestor's symlinks, and re-joins the tail.
|
||||
// This ensures even deeply nested non-existent paths are anchored to a
|
||||
// real filesystem location, closing the TOCTOU symlink gap.
|
||||
func resolveNearestAncestor(path string) (string, error) {
|
||||
var tail []string
|
||||
cur := path
|
||||
for {
|
||||
if _, err := os.Lstat(cur); err == nil {
|
||||
real, err := filepath.EvalSymlinks(cur)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
parts := append([]string{real}, tail...)
|
||||
return filepath.Join(parts...), nil
|
||||
}
|
||||
parent := filepath.Dir(cur)
|
||||
if parent == cur {
|
||||
// Reached filesystem root without finding an existing ancestor;
|
||||
// return path as-is and let the containment check reject it.
|
||||
parts := append([]string{cur}, tail...)
|
||||
return filepath.Join(parts...), nil
|
||||
}
|
||||
tail = append([]string{filepath.Base(cur)}, tail...)
|
||||
cur = parent
|
||||
}
|
||||
}
|
||||
|
||||
// isUnderDir checks whether child is under parent directory.
|
||||
func isUnderDir(child, parent string) bool {
|
||||
rel, err := filepath.Rel(parent, child)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return !strings.HasPrefix(rel, ".."+string(filepath.Separator)) && rel != ".."
|
||||
}
|
||||
|
||||
// SafeLocalFlagPath validates a flag value as a local file path.
|
||||
// Empty values and http/https URLs are returned unchanged without validation,
|
||||
// allowing the caller to handle non-path inputs (e.g. API keys, URLs) upstream.
|
||||
// For all other values, SafeInputPath rules apply.
|
||||
// The original relative path is returned unchanged (not resolved to absolute) so
|
||||
// upload helpers can re-validate at the actual I/O point via SafeUploadPath.
|
||||
func SafeLocalFlagPath(flagName, value string) (string, error) {
|
||||
if value == "" || strings.HasPrefix(value, "http://") || strings.HasPrefix(value, "https://") {
|
||||
return value, nil
|
||||
}
|
||||
if _, err := SafeInputPath(value); err != nil {
|
||||
return "", fmt.Errorf("%s: %v", flagName, err)
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
// RejectCRLF rejects strings containing carriage return (\r) or line feed (\n).
|
||||
// These characters enable MIME/HTTP header injection and must never appear in
|
||||
// header field names, values, Content-ID, or filename parameters.
|
||||
func RejectCRLF(value, fieldName string) error {
|
||||
if strings.ContainsAny(value, "\r\n") {
|
||||
return fmt.Errorf("%s contains invalid line break characters", fieldName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// StripQueryFragment removes any ?query or #fragment suffix from a URL path.
|
||||
// API parameters must go through structured --params flags, not embedded in
|
||||
// the path, to prevent parameter injection and behaviour confusion.
|
||||
func StripQueryFragment(path string) string {
|
||||
for i := 0; i < len(path); i++ {
|
||||
if path[i] == '?' || path[i] == '#' {
|
||||
return path[:i]
|
||||
}
|
||||
}
|
||||
return path
|
||||
}
|
||||
|
||||
@@ -73,3 +73,162 @@ func TestSafePath(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeLocalFlagPath(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
flagName string
|
||||
value string
|
||||
want string
|
||||
wantErr bool
|
||||
}{
|
||||
// URL pass-through
|
||||
{name: "http-url", flagName: "--source", value: "http://example.com/file", want: "http://example.com/file"},
|
||||
{name: "https-url", flagName: "--source", value: "https://api.example.com/data", want: "https://api.example.com/data"},
|
||||
// Empty pass-through
|
||||
{name: "empty", flagName: "--file", value: "", want: ""},
|
||||
// Valid relative paths (returns original relative path)
|
||||
{name: "relative-file", flagName: "--file", value: "data.json", want: "data.json"},
|
||||
{name: "relative-nested", flagName: "--file", value: "dir/file.txt", want: "dir/file.txt"},
|
||||
// Invalid paths
|
||||
{name: "absolute", flagName: "--file", value: "/etc/passwd", wantErr: true},
|
||||
{name: "traversal", flagName: "--file", value: "../secret", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
got, err := SafeLocalFlagPath(tc.flagName, tc.value)
|
||||
if tc.wantErr {
|
||||
if err == nil {
|
||||
t.Fatalf("SafeLocalFlagPath(%q, %q) error = nil, want failure", tc.flagName, tc.value)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("SafeLocalFlagPath(%q, %q) error = %v, want nil", tc.flagName, tc.value, err)
|
||||
}
|
||||
if got != tc.want {
|
||||
t.Fatalf("SafeLocalFlagPath(%q, %q) = %q, want %q", tc.flagName, tc.value, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRejectCRLF(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
value string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "clean", value: "normal text"},
|
||||
{name: "with-space", value: "text with spaces"},
|
||||
{name: "with-CR", value: "text\rwith CR", wantErr: true},
|
||||
{name: "with-LF", value: "text\nwith LF", wantErr: true},
|
||||
{name: "with-CRLF", value: "text\r\nwith CRLF", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := RejectCRLF(tc.value, "--header")
|
||||
if tc.wantErr && err == nil {
|
||||
t.Fatalf("RejectCRLF(%q) error = nil, want failure", tc.value)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Fatalf("RejectCRLF(%q) error = %v, want nil", tc.value, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripQueryFragment(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
input string
|
||||
want string
|
||||
}{
|
||||
{name: "no-query", input: "/api/v1/users", want: "/api/v1/users"},
|
||||
{name: "with-query", input: "/api/v1/users?page=1", want: "/api/v1/users"},
|
||||
{name: "with-fragment", input: "/docs#section", want: "/docs"},
|
||||
{name: "query-and-fragment", input: "/api?a=1#sec", want: "/api"},
|
||||
{name: "fragment-before-query", input: "/path#frag?notquery", want: "/path"},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
got := StripQueryFragment(tc.input)
|
||||
if got != tc.want {
|
||||
t.Fatalf("StripQueryFragment(%q) = %q, want %q", tc.input, got, tc.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRejectControlChars(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
input string
|
||||
wantErr bool
|
||||
}{
|
||||
// ── Normal text: allowed ──
|
||||
{name: "plain-text", input: "hello world"},
|
||||
{name: "with-tab", input: "hello\tworld"},
|
||||
{name: "with-newline", input: "hello\nworld"},
|
||||
{name: "unicode-text", input: "你好世界"},
|
||||
{name: "mixed-unicode", input: "Hello 世界 123"},
|
||||
{name: "emoji", input: "test 😀 emoji"},
|
||||
|
||||
// ── C0 control characters (except tab/newline): rejected ──
|
||||
{name: "null-byte", input: "bad\x00path", wantErr: true},
|
||||
{name: "bell", input: "alert\x07here", wantErr: true},
|
||||
{name: "backspace", input: "back\x08space", wantErr: true},
|
||||
{name: "form-feed", input: "form\x0cfeed", wantErr: true},
|
||||
{name: "carriage-return", input: "cr\rhere", wantErr: true},
|
||||
{name: "escape", input: "esc\x1bhere", wantErr: true},
|
||||
{name: "delete", input: "del\x7fete", wantErr: true},
|
||||
|
||||
// ── Dangerous Unicode: rejected ──
|
||||
{name: "zero-width-space", input: "foo\u200Bbar", wantErr: true},
|
||||
{name: "zero-width-non-joiner", input: "foo\u200Cbar", wantErr: true},
|
||||
{name: "zero-width-joiner", input: "foo\u200Dbar", wantErr: true},
|
||||
{name: "bom", input: "\uFEFFstart", wantErr: true},
|
||||
{name: "bidi-lre", input: "foo\u202Abar", wantErr: true},
|
||||
{name: "bidi-rle", input: "foo\u202Bbar", wantErr: true},
|
||||
{name: "bidi-pdf", input: "foo\u202Cbar", wantErr: true},
|
||||
{name: "bidi-lro", input: "foo\u202Dbar", wantErr: true},
|
||||
{name: "bidi-rlo", input: "foo\u202Ebar", wantErr: true},
|
||||
{name: "line-separator", input: "foo\u2028bar", wantErr: true},
|
||||
{name: "paragraph-separator", input: "foo\u2029bar", wantErr: true},
|
||||
{name: "bidi-lri", input: "foo\u2066bar", wantErr: true},
|
||||
{name: "bidi-rli", input: "foo\u2067bar", wantErr: true},
|
||||
{name: "bidi-fsi", input: "foo\u2068bar", wantErr: true},
|
||||
{name: "bidi-pdi", input: "foo\u2069bar", wantErr: true},
|
||||
}
|
||||
|
||||
for _, tc := range cases {
|
||||
tc := tc
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
err := RejectControlChars(tc.input, "--test")
|
||||
if tc.wantErr && err == nil {
|
||||
t.Fatalf("RejectControlChars(%q) error = nil, want failure", tc.input)
|
||||
}
|
||||
if !tc.wantErr && err != nil {
|
||||
t.Fatalf("RejectControlChars(%q) error = %v, want nil", tc.input, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,7 +41,7 @@ type frontmatterSpec struct {
|
||||
}
|
||||
|
||||
const (
|
||||
skillVersion = "1.0.0"
|
||||
skillVersion = "1.1.0"
|
||||
skillCategoryService = "service"
|
||||
skillCategoryHelper = "helper"
|
||||
skillCategoryPersona = "persona"
|
||||
@@ -123,7 +123,6 @@ var extended22CoverageTargets = []string{
|
||||
"finance",
|
||||
"law",
|
||||
"docparse",
|
||||
"credit",
|
||||
}
|
||||
|
||||
func Generate(catalog ir.Catalog) ([]Artifact, error) {
|
||||
@@ -395,7 +394,7 @@ func renderSharedSkill(catalog ir.Catalog) string {
|
||||
builder.WriteString("dws auth status\n")
|
||||
builder.WriteString("```\n\n")
|
||||
builder.WriteString("## Global Rules\n\n")
|
||||
builder.WriteString("- Always prefer `--format json` for agent-readable output.\n")
|
||||
builder.WriteString("- Output defaults to JSON. Use `--format table` for human-readable output.\n")
|
||||
builder.WriteString("- Confirm with user before any write/delete/revoke action.\n")
|
||||
builder.WriteString("- Never fabricate IDs; always extract from command output.\n")
|
||||
builder.WriteString("- For risky operations, run a read/list check before executing write operations.\n\n")
|
||||
|
||||
@@ -67,7 +67,6 @@ var knownRegistryProducts = map[string]struct{}{
|
||||
"chat": {},
|
||||
"conference": {},
|
||||
"contact": {},
|
||||
"credit": {},
|
||||
"devdoc": {},
|
||||
"ding": {},
|
||||
"doc": {},
|
||||
|
||||
@@ -26,9 +26,9 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/config"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -424,6 +424,9 @@ func newAITableUploadFileCommand(runner executor.Runner) *cobra.Command {
|
||||
}
|
||||
|
||||
resultMap := result.Response
|
||||
if content, ok := resultMap["content"].(map[string]any); ok && len(content) > 0 {
|
||||
resultMap = content
|
||||
}
|
||||
if resultMap == nil {
|
||||
return apperrors.NewValidation(i18n.T("prepare_attachment_upload 返回格式异常"))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
|
||||
type uploadFileRunner struct {
|
||||
last executor.Invocation
|
||||
result executor.Result
|
||||
}
|
||||
|
||||
func (r *uploadFileRunner) Run(_ context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
r.last = invocation
|
||||
return r.result, nil
|
||||
}
|
||||
|
||||
func TestAITableUploadFileUnwrapsRuntimeContent(t *testing.T) {
|
||||
filePath := filepath.Join(t.TempDir(), "report.txt")
|
||||
wantBody := []byte("hello from upload-file")
|
||||
if err := os.WriteFile(filePath, wantBody, 0o644); err != nil {
|
||||
t.Fatalf("WriteFile() error = %v", err)
|
||||
}
|
||||
|
||||
var gotBody []byte
|
||||
var gotContentType string
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
|
||||
if req.Method != http.MethodPut {
|
||||
t.Fatalf("method = %s, want PUT", req.Method)
|
||||
}
|
||||
gotContentType = req.Header.Get("Content-Type")
|
||||
body, err := io.ReadAll(req.Body)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadAll() error = %v", err)
|
||||
}
|
||||
gotBody = body
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
runner := &uploadFileRunner{
|
||||
result: executor.Result{
|
||||
Response: map[string]any{
|
||||
"content": map[string]any{
|
||||
"data": map[string]any{
|
||||
"uploadUrl": server.URL,
|
||||
"fileToken": "ft_test_123",
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
cmd := newAITableUploadFileCommand(runner)
|
||||
var out bytes.Buffer
|
||||
var errOut bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&errOut)
|
||||
cmd.SetArgs([]string{"--base-id", "BASE_001", "--file", filePath})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v\nstderr:\n%s", err, errOut.String())
|
||||
}
|
||||
|
||||
if runner.last.Tool != "prepare_attachment_upload" {
|
||||
t.Fatalf("tool = %q, want prepare_attachment_upload", runner.last.Tool)
|
||||
}
|
||||
if got := runner.last.Params["fileName"]; got != "report.txt" {
|
||||
t.Fatalf("fileName = %#v, want report.txt", got)
|
||||
}
|
||||
if string(gotBody) != string(wantBody) {
|
||||
t.Fatalf("uploaded body = %q, want %q", string(gotBody), string(wantBody))
|
||||
}
|
||||
if gotContentType != "text/plain; charset=utf-8" {
|
||||
t.Fatalf("content-type = %q, want text/plain; charset=utf-8", gotContentType)
|
||||
}
|
||||
|
||||
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 got := payload["fileToken"]; got != "ft_test_123" {
|
||||
t.Fatalf("fileToken = %#v, want ft_test_123", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
// AtomicWrite writes data to path atomically by creating a temp file in the
|
||||
// same directory, writing and fsyncing the data, then renaming over the target.
|
||||
// It replaces os.WriteFile for all config and download file writes.
|
||||
//
|
||||
// os.WriteFile truncates the target before writing, so a process kill (CI timeout,
|
||||
// OOM, Ctrl+C) between truncate and completion leaves the file empty or partial.
|
||||
// AtomicWrite avoids this: on any failure the temp file is cleaned up and the
|
||||
// original file remains untouched.
|
||||
func AtomicWrite(path string, data []byte, perm os.FileMode) error {
|
||||
return atomicWrite(path, perm, func(tmp *os.File) error {
|
||||
_, err := tmp.Write(data)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// AtomicWriteFromReader atomically copies reader contents into path.
|
||||
func AtomicWriteFromReader(path string, reader io.Reader, perm os.FileMode) (int64, error) {
|
||||
var copied int64
|
||||
err := atomicWrite(path, perm, func(tmp *os.File) error {
|
||||
n, err := io.Copy(tmp, reader)
|
||||
copied = n
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return copied, nil
|
||||
}
|
||||
|
||||
// AtomicWriteJSON is a convenience wrapper for writing JSON data atomically.
|
||||
// It uses 0600 permissions by default for sensitive data.
|
||||
func AtomicWriteJSON(path string, data []byte) error {
|
||||
return AtomicWrite(path, data, 0600)
|
||||
}
|
||||
|
||||
func atomicWrite(path string, perm os.FileMode, writeFn func(tmp *os.File) error) error {
|
||||
dir := filepath.Dir(path)
|
||||
|
||||
// Ensure directory exists with secure permissions
|
||||
if err := os.MkdirAll(dir, 0700); err != nil {
|
||||
return fmt.Errorf("create directory: %w", err)
|
||||
}
|
||||
|
||||
tmp, err := os.CreateTemp(dir, "."+filepath.Base(path)+".*.tmp")
|
||||
if err != nil {
|
||||
return fmt.Errorf("create temp file: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
|
||||
success := false
|
||||
defer func() {
|
||||
if !success {
|
||||
tmp.Close()
|
||||
os.Remove(tmpName)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := tmp.Chmod(perm); err != nil {
|
||||
return fmt.Errorf("set permissions: %w", err)
|
||||
}
|
||||
if err := writeFn(tmp); err != nil {
|
||||
return fmt.Errorf("write data: %w", err)
|
||||
}
|
||||
if err := tmp.Sync(); err != nil {
|
||||
return fmt.Errorf("sync to disk: %w", err)
|
||||
}
|
||||
if err := tmp.Close(); err != nil {
|
||||
return fmt.Errorf("close temp file: %w", err)
|
||||
}
|
||||
if err := os.Rename(tmpName, path); err != nil {
|
||||
return fmt.Errorf("rename to final: %w", err)
|
||||
}
|
||||
success = true
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAtomicWrite_Basic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.txt")
|
||||
data := []byte("hello world")
|
||||
|
||||
if err := AtomicWrite(path, data, 0600); err != nil {
|
||||
t.Fatalf("AtomicWrite() error = %v", err)
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("got %q, want %q", got, data)
|
||||
}
|
||||
|
||||
// Check permissions
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if perm := info.Mode().Perm(); perm != 0600 {
|
||||
t.Fatalf("permissions = %o, want 0600", perm)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtomicWrite_CreatesDirectory(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
base := t.TempDir()
|
||||
path := filepath.Join(base, "a", "b", "c", "test.txt")
|
||||
data := []byte("nested content")
|
||||
|
||||
if err := AtomicWrite(path, data, 0644); err != nil {
|
||||
t.Fatalf("AtomicWrite() error = %v", err)
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("got %q, want %q", got, data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtomicWrite_Overwrite(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.txt")
|
||||
|
||||
// Write initial content
|
||||
if err := AtomicWrite(path, []byte("initial"), 0600); err != nil {
|
||||
t.Fatalf("AtomicWrite() initial error = %v", err)
|
||||
}
|
||||
|
||||
// Overwrite
|
||||
newData := []byte("overwritten content")
|
||||
if err := AtomicWrite(path, newData, 0600); err != nil {
|
||||
t.Fatalf("AtomicWrite() overwrite error = %v", err)
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, newData) {
|
||||
t.Fatalf("got %q, want %q", got, newData)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtomicWrite_NoTempFileOnSuccess(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.txt")
|
||||
|
||||
if err := AtomicWrite(path, []byte("content"), 0600); err != nil {
|
||||
t.Fatalf("AtomicWrite() error = %v", err)
|
||||
}
|
||||
|
||||
// Check no .tmp files remain
|
||||
entries, err := os.ReadDir(dir)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadDir() error = %v", err)
|
||||
}
|
||||
for _, e := range entries {
|
||||
if strings.HasSuffix(e.Name(), ".tmp") {
|
||||
t.Fatalf("temp file remains: %s", e.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtomicWriteFromReader_Basic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "test.txt")
|
||||
data := "reader content"
|
||||
reader := strings.NewReader(data)
|
||||
|
||||
n, err := AtomicWriteFromReader(path, reader, 0600)
|
||||
if err != nil {
|
||||
t.Fatalf("AtomicWriteFromReader() error = %v", err)
|
||||
}
|
||||
if n != int64(len(data)) {
|
||||
t.Fatalf("written bytes = %d, want %d", n, len(data))
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if string(got) != data {
|
||||
t.Fatalf("got %q, want %q", got, data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAtomicWriteJSON_Basic(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "config.json")
|
||||
data := []byte(`{"key": "value"}`)
|
||||
|
||||
if err := AtomicWriteJSON(path, data); err != nil {
|
||||
t.Fatalf("AtomicWriteJSON() error = %v", err)
|
||||
}
|
||||
|
||||
got, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile() error = %v", err)
|
||||
}
|
||||
if !bytes.Equal(got, data) {
|
||||
t.Fatalf("got %q, want %q", got, data)
|
||||
}
|
||||
|
||||
// Check permissions are 0600
|
||||
info, err := os.Stat(path)
|
||||
if err != nil {
|
||||
t.Fatalf("Stat() error = %v", err)
|
||||
}
|
||||
if perm := info.Mode().Perm(); perm != 0600 {
|
||||
t.Fatalf("permissions = %o, want 0600", perm)
|
||||
}
|
||||
}
|
||||
+54
-169
@@ -15,9 +15,9 @@ package helpers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
@@ -60,8 +60,6 @@ func (chatHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
},
|
||||
}
|
||||
message.AddCommand(
|
||||
newChatMessageListCommand(runner),
|
||||
newChatMessageSendCommand(runner),
|
||||
newChatMessageSendByBotCommand(runner),
|
||||
newChatMessageRecallByBotCommand(runner),
|
||||
newChatMessageSendByWebhookCommand(runner),
|
||||
@@ -83,40 +81,6 @@ func (chatHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
return root
|
||||
}
|
||||
|
||||
func newChatMessageListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "拉取会话消息内容",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
params, tool, err := buildChatMessageListInvocation(cmd)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"chat",
|
||||
tool,
|
||||
params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("group", "", "群会话 openconversation_id (群聊必填)")
|
||||
cmd.Flags().Bool("forward", true, "true=正序拉取, false=倒序拉取")
|
||||
cmd.Flags().Int("limit", 0, "返回条数,不传为不限制")
|
||||
cmd.Flags().String("time", "", "起始时间,格式: yyyy-MM-dd HH:mm:ss")
|
||||
cmd.Flags().String("user", "", "单聊对方 userId (单聊必填)")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newChatMessageSendByBotCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "send-by-bot",
|
||||
@@ -129,12 +93,14 @@ func newChatMessageSendByBotCommand(runner executor.Runner) *cobra.Command {
|
||||
return err
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"chat",
|
||||
tool,
|
||||
params,
|
||||
))
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -266,7 +232,7 @@ func newChatGroupCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
}
|
||||
allMembers := prependOwner(currentUserID, memberUserIDs)
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"chat",
|
||||
"create_internal_group",
|
||||
@@ -274,7 +240,9 @@ func newChatGroupCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
"groupMembers": stringSliceToAny(allMembers),
|
||||
"groupName": name,
|
||||
},
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -289,48 +257,9 @@ func newChatGroupCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func buildChatMessageListInvocation(cmd *cobra.Command) (map[string]any, string, error) {
|
||||
group, err := cmd.Flags().GetString("group")
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --group")
|
||||
}
|
||||
user, err := cmd.Flags().GetString("user")
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --user")
|
||||
}
|
||||
timeValue, err := cmd.Flags().GetString("time")
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --time")
|
||||
}
|
||||
if strings.TrimSpace(timeValue) == "" {
|
||||
return nil, "", apperrors.NewValidation("--time is required")
|
||||
}
|
||||
if err := ensureExactlyOneTarget(group, user, "--group", "--user"); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"forward": cmd.Flag("forward").Value.String() == "true",
|
||||
"time": timeValue,
|
||||
}
|
||||
limit, err := cmd.Flags().GetInt("limit")
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --limit")
|
||||
}
|
||||
if limit > 0 {
|
||||
params["limit"] = limit
|
||||
}
|
||||
|
||||
if strings.TrimSpace(group) != "" {
|
||||
params["openconversation_id"] = group
|
||||
return params, "list_conversation_message_v2", nil
|
||||
}
|
||||
|
||||
params["userId"] = user
|
||||
return params, "list_individual_chat_message", nil
|
||||
}
|
||||
|
||||
func buildChatMessageSendByBotInvocation(cmd *cobra.Command) (map[string]any, string, error) {
|
||||
guard := cli.NewStdinGuard()
|
||||
|
||||
group, err := cmd.Flags().GetString("group")
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --group")
|
||||
@@ -343,13 +272,14 @@ func buildChatMessageSendByBotInvocation(cmd *cobra.Command) (map[string]any, st
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --robot-code")
|
||||
}
|
||||
title, err := cmd.Flags().GetString("title")
|
||||
title, err := resolveStringFlag(cmd, "title", guard, false)
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --title")
|
||||
return nil, "", err
|
||||
}
|
||||
text, err := cmd.Flags().GetString("text")
|
||||
// --text is the primary content flag: receives stdin pipe when empty.
|
||||
text, err := resolveStringFlag(cmd, "text", guard, true)
|
||||
if err != nil {
|
||||
return nil, "", apperrors.NewInternal("failed to read --text")
|
||||
return nil, "", err
|
||||
}
|
||||
|
||||
switch {
|
||||
@@ -382,17 +312,6 @@ func buildChatMessageSendByBotInvocation(cmd *cobra.Command) (map[string]any, st
|
||||
return params, "batch_send_robot_msg_to_users", nil
|
||||
}
|
||||
|
||||
func ensureExactlyOneTarget(left, right, leftFlag, rightFlag string) error {
|
||||
switch {
|
||||
case strings.TrimSpace(left) == "" && strings.TrimSpace(right) == "":
|
||||
return apperrors.NewValidation(fmt.Sprintf("either %s or %s is required", leftFlag, rightFlag))
|
||||
case strings.TrimSpace(left) != "" && strings.TrimSpace(right) != "":
|
||||
return apperrors.NewValidation(fmt.Sprintf("%s and %s are mutually exclusive", leftFlag, rightFlag))
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func splitCSV(raw string) []any {
|
||||
parts := strings.Split(raw, ",")
|
||||
values := make([]any, 0, len(parts))
|
||||
@@ -509,62 +428,6 @@ func helperResponseContent(result executor.Result) map[string]any {
|
||||
return content
|
||||
}
|
||||
|
||||
// ── message send (user identity) ───────────────────────────
|
||||
|
||||
func newChatMessageSendCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "send",
|
||||
Short: "以当前用户身份发送消息(--group 群聊 / --user 单聊)",
|
||||
Long: "--group 指定群会话 ID 发送群消息;--user 指定用户 ID 发送单聊消息。二者只能选其一。",
|
||||
Example: ` dws chat message send --group <openconversation_id> "hello"
|
||||
dws chat message send --user <userId> "请查收"`,
|
||||
Args: cobra.ExactArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
groupID, _ := cmd.Flags().GetString("group")
|
||||
userID, _ := cmd.Flags().GetString("user")
|
||||
if err := ensureExactlyOneTarget(groupID, userID, "--group", "--user"); err != nil {
|
||||
return err
|
||||
}
|
||||
text := args[0]
|
||||
title, _ := cmd.Flags().GetString("title")
|
||||
if strings.TrimSpace(groupID) != "" {
|
||||
params := map[string]any{
|
||||
"openConversation_id": groupID,
|
||||
"title": title,
|
||||
"text": text,
|
||||
"clawType": "default",
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "chat", "send_message_as_user", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
params := map[string]any{
|
||||
"receiverUserId": userID,
|
||||
"title": title,
|
||||
"text": text,
|
||||
"clawType": "default",
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "chat", "send_direct_message_as_user", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
cmd.Flags().String("group", "", "群会话 openconversation_id (群聊必填)")
|
||||
cmd.Flags().String("user", "", "接收者 userId (单聊必填)")
|
||||
cmd.Flags().String("title", "Message", "消息标题 (可选, 默认 'Message')")
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── message recall-by-bot ──────────────────────────────────
|
||||
|
||||
func newChatMessageRecallByBotCommand(runner executor.Runner) *cobra.Command {
|
||||
@@ -593,9 +456,11 @@ func newChatMessageRecallByBotCommand(runner executor.Runner) *cobra.Command {
|
||||
"openConversationId": groupID,
|
||||
"processQueryKeys": processQueryKeys,
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "bot", "recall_robot_group_message", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -605,9 +470,11 @@ func newChatMessageRecallByBotCommand(runner executor.Runner) *cobra.Command {
|
||||
"robotCode": robotCode,
|
||||
"processQueryKeys": processQueryKeys,
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "bot", "batch_recall_robot_users_msg", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -633,9 +500,17 @@ func newChatMessageSendByWebhookCommand(runner executor.Runner) *cobra.Command {
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
guard := cli.NewStdinGuard()
|
||||
token, _ := cmd.Flags().GetString("token")
|
||||
title, _ := cmd.Flags().GetString("title")
|
||||
text, _ := cmd.Flags().GetString("text")
|
||||
title, err := resolveStringFlag(cmd, "title", guard, false)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// --text is the primary content flag: receives stdin pipe when empty.
|
||||
text, err := resolveStringFlag(cmd, "text", guard, true)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.TrimSpace(token) == "" {
|
||||
return apperrors.NewValidation("--token is required")
|
||||
}
|
||||
@@ -659,9 +534,11 @@ func newChatMessageSendByWebhookCommand(runner executor.Runner) *cobra.Command {
|
||||
if v, _ := cmd.Flags().GetString("at-users"); v != "" {
|
||||
params["atUserIds"] = splitCSV(v)
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "bot", "send_message_by_custom_robot", params,
|
||||
))
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -724,9 +601,11 @@ func newChatGroupRenameCommand(runner executor.Runner) *cobra.Command {
|
||||
"openconversation_id": groupID,
|
||||
"group_name": name,
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "chat", "update_group_name", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -761,9 +640,11 @@ func newChatGroupMemberAddCommand(runner executor.Runner) *cobra.Command {
|
||||
"openconversation_id": groupID,
|
||||
"userId": splitCSV(usersStr),
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "chat", "add_group_member", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -798,9 +679,11 @@ func newChatGroupMemberRemoveCommand(runner executor.Runner) *cobra.Command {
|
||||
"openconversationId": groupID,
|
||||
"userIdList": splitCSV(usersStr),
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "chat", "remove_group_member", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -835,9 +718,11 @@ func newChatGroupMembersAddBotCommand(runner executor.Runner) *cobra.Command {
|
||||
"robotCode": robotCode,
|
||||
"openConversationId": groupID,
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
inv := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "bot", "add_robot_to_group", params,
|
||||
))
|
||||
)
|
||||
inv.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), inv)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,456 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
// inputCaptureRunner records invocations with a call counter.
|
||||
type inputCaptureRunner struct {
|
||||
last executor.Invocation
|
||||
called int
|
||||
}
|
||||
|
||||
func (r *inputCaptureRunner) Run(_ context.Context, inv executor.Invocation) (executor.Result, error) {
|
||||
r.last = inv
|
||||
r.called++
|
||||
return executor.Result{Invocation: inv}, nil
|
||||
}
|
||||
|
||||
func writeInputTestFile(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)
|
||||
}
|
||||
}
|
||||
|
||||
// newInputTestChatRoot builds the full chat command tree for testing.
|
||||
func newInputTestChatRoot(t *testing.T, runner executor.Runner) *cobra.Command {
|
||||
t.Helper()
|
||||
h := chatHandler{}
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.AddCommand(h.Command(runner))
|
||||
var out, errOut bytes.Buffer
|
||||
root.SetOut(&out)
|
||||
root.SetErr(&errOut)
|
||||
return root
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// send-by-bot: --text @file
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendByBotTextFromFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "msg.md")
|
||||
writeInputTestFile(t, filePath, []byte("# Weekly Report\n\nAll green."))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "周报",
|
||||
"--text", "@" + filePath,
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.called != 1 {
|
||||
t.Fatalf("runner called = %d, want 1", runner.called)
|
||||
}
|
||||
if runner.last.Params["markdown"] != "# Weekly Report\n\nAll green." {
|
||||
t.Errorf("params[markdown] = %q, want file content", runner.last.Params["markdown"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotTitleFromFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "title.txt")
|
||||
writeInputTestFile(t, filePath, []byte("Dynamic Title"))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "@" + filePath,
|
||||
"--text", "content here",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["title"] != "Dynamic Title" {
|
||||
t.Errorf("params[title] = %q, want %q", runner.last.Params["title"], "Dynamic Title")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotTextAndTitleBothFromFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
titlePath := filepath.Join(dir, "title.txt")
|
||||
textPath := filepath.Join(dir, "body.md")
|
||||
writeInputTestFile(t, titlePath, []byte("File Title"))
|
||||
writeInputTestFile(t, textPath, []byte("File Body"))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "@" + titlePath,
|
||||
"--text", "@" + textPath,
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["title"] != "File Title" {
|
||||
t.Errorf("params[title] = %q, want %q", runner.last.Params["title"], "File Title")
|
||||
}
|
||||
if runner.last.Params["markdown"] != "File Body" {
|
||||
t.Errorf("params[markdown] = %q, want %q", runner.last.Params["markdown"], "File Body")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotTextFromFileMissingReturnsError(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "test",
|
||||
"--text", "@/nonexistent/file.md",
|
||||
})
|
||||
|
||||
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 --text, got: %v", err)
|
||||
}
|
||||
if runner.called != 0 {
|
||||
t.Error("runner should not be called on @file error")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotTextUTF8Preserved(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "chinese.md")
|
||||
content := "你好世界 🌍\n第二行"
|
||||
writeInputTestFile(t, filePath, []byte(content))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "测试",
|
||||
"--text", "@" + filePath,
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["markdown"] != content {
|
||||
t.Errorf("params[markdown] = %q, want %q", runner.last.Params["markdown"], content)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// send-by-bot: backward compatibility (plain --text)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendByBotPlainTextStillWorks(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "test",
|
||||
"--text", "plain message",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["markdown"] != "plain message" {
|
||||
t.Errorf("params[markdown] = %q, want %q", runner.last.Params["markdown"], "plain message")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotSingleChatStillWorks(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--users", "u001,u002",
|
||||
"--title", "test",
|
||||
"--text", "hello",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Tool != "batch_send_robot_msg_to_users" {
|
||||
t.Errorf("tool = %q, want batch_send_robot_msg_to_users", runner.last.Tool)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// send-by-bot: validation still works
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestSendByBotEmptyTextStillErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--group", "G001",
|
||||
"--title", "test",
|
||||
// --text not provided
|
||||
})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() should fail when --text is empty")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--text") {
|
||||
t.Errorf("error should mention --text, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSendByBotMissingGroupAndUsersStillErrors(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-bot",
|
||||
"--robot-code", "BOT001",
|
||||
"--title", "test",
|
||||
"--text", "hello",
|
||||
})
|
||||
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatal("Execute() should fail when --group and --users both missing")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// send-by-webhook: --text @file and --title @file
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestWebhookTextFromFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "alert.md")
|
||||
writeInputTestFile(t, filePath, []byte("CPU > 90%"))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-webhook",
|
||||
"--token", "TOKEN001",
|
||||
"--title", "告警",
|
||||
"--text", "@" + filePath,
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["text"] != "CPU > 90%" {
|
||||
t.Errorf("params[text] = %q, want %q", runner.last.Params["text"], "CPU > 90%")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebhookTitleFromFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "title.txt")
|
||||
writeInputTestFile(t, filePath, []byte("Alert Title"))
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-webhook",
|
||||
"--token", "TOKEN001",
|
||||
"--title", "@" + filePath,
|
||||
"--text", "body content",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["title"] != "Alert Title" {
|
||||
t.Errorf("params[title] = %q, want %q", runner.last.Params["title"], "Alert Title")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWebhookPlainTextStillWorks(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
runner := &inputCaptureRunner{}
|
||||
cmd := newInputTestChatRoot(t, runner)
|
||||
cmd.SetArgs([]string{"chat", "message", "send-by-webhook",
|
||||
"--token", "TOKEN001",
|
||||
"--title", "test",
|
||||
"--text", "plain webhook message",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v", err)
|
||||
}
|
||||
if runner.last.Params["text"] != "plain webhook message" {
|
||||
t.Errorf("params[text] = %q, want %q", runner.last.Params["text"], "plain webhook message")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// resolveStringFlag unit tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestResolveStringFlagPlainValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("text", "", "")
|
||||
cmd.SetArgs([]string{"--text", "hello"})
|
||||
_ = cmd.Execute()
|
||||
|
||||
guard := cli.NewStdinGuard()
|
||||
val, err := resolveStringFlag(cmd, "text", guard, false)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "hello" {
|
||||
t.Errorf("got %q, want %q", val, "hello")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStringFlagAtFile(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
dir := t.TempDir()
|
||||
filePath := filepath.Join(dir, "data.txt")
|
||||
writeInputTestFile(t, filePath, []byte("file content"))
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("body", "", "")
|
||||
cmd.SetArgs([]string{"--body", "@" + filePath})
|
||||
_ = cmd.Execute()
|
||||
|
||||
guard := cli.NewStdinGuard()
|
||||
val, err := resolveStringFlag(cmd, "body", guard, false)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "file content" {
|
||||
t.Errorf("got %q, want %q", val, "file content")
|
||||
}
|
||||
if guard.Claimed() {
|
||||
t.Error("@file should not claim stdin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStringFlagAtFileMissing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("text", "", "")
|
||||
cmd.SetArgs([]string{"--text", "@/no/such/file"})
|
||||
_ = cmd.Execute()
|
||||
|
||||
guard := cli.NewStdinGuard()
|
||||
_, err := resolveStringFlag(cmd, "text", guard, false)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing @file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "--text") {
|
||||
t.Errorf("error should mention flag name, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStringFlagPrimaryContentNoStdinInTerminal(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// In go test context, stdin is a terminal — primary content fallback
|
||||
// should NOT read stdin (StdinIsPipe returns false).
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("text", "", "")
|
||||
cmd.SetArgs([]string{})
|
||||
_ = cmd.Execute()
|
||||
|
||||
guard := cli.NewStdinGuard()
|
||||
val, err := resolveStringFlag(cmd, "text", guard, true)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "" {
|
||||
t.Errorf("got %q, want empty (no stdin pipe in terminal)", val)
|
||||
}
|
||||
if guard.Claimed() {
|
||||
t.Error("guard should not be claimed in terminal context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveStringFlagExplicitValueBlocksStdinFallback(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// When --text has an explicit value, primaryContent stdin fallback
|
||||
// should NOT activate even if it's the primary flag.
|
||||
cmd := &cobra.Command{Use: "test"}
|
||||
cmd.Flags().String("text", "", "")
|
||||
cmd.SetArgs([]string{"--text", "explicit"})
|
||||
_ = cmd.Execute()
|
||||
|
||||
guard := cli.NewStdinGuard()
|
||||
val, err := resolveStringFlag(cmd, "text", guard, true)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if val != "explicit" {
|
||||
t.Errorf("got %q, want %q", val, "explicit")
|
||||
}
|
||||
if guard.Claimed() {
|
||||
t.Error("explicit value should not claim stdin")
|
||||
}
|
||||
}
|
||||
@@ -17,22 +17,30 @@ func (r *captureRunner) Run(_ context.Context, invocation executor.Invocation) (
|
||||
return executor.Result{Invocation: invocation}, nil
|
||||
}
|
||||
|
||||
func TestChatMessageSendIgnoresLegacyRealBuildModeEnv(t *testing.T) {
|
||||
func TestChatMessageSendByBotIgnoresLegacyRealBuildModeEnv(t *testing.T) {
|
||||
t.Setenv("DWS_"+"BUILD_MODE", "real")
|
||||
|
||||
runner := &captureRunner{}
|
||||
cmd := newChatMessageSendCommand(runner)
|
||||
cmd := newChatMessageSendByBotCommand(runner)
|
||||
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"--user", "user-001", "hello"})
|
||||
cmd.SetArgs([]string{
|
||||
"--users", "user-001",
|
||||
"--robot-code", "robot-001",
|
||||
"--title", "Greeting",
|
||||
"--text", "hello",
|
||||
})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
|
||||
if got := runner.last.Params["clawType"]; got != "default" {
|
||||
t.Fatalf("clawType = %#v, want default", got)
|
||||
if got := runner.last.Tool; got != "batch_send_robot_msg_to_users" {
|
||||
t.Fatalf("tool = %q, want batch_send_robot_msg_to_users", got)
|
||||
}
|
||||
if got := runner.last.Params["robotCode"]; got != "robot-001" {
|
||||
t.Fatalf("robotCode = %#v, want robot-001", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -15,6 +15,7 @@ package helpers
|
||||
|
||||
import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cli"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/spf13/cobra"
|
||||
"github.com/spf13/pflag"
|
||||
@@ -28,6 +29,37 @@ func preferLegacyLeaf(cmd *cobra.Command) {
|
||||
cli.SetOverridePriority(cmd, 100)
|
||||
}
|
||||
|
||||
// resolveStringFlag reads a string flag, resolves @file/@- input sources,
|
||||
// and falls back to stdin pipe when the flag is the designated primary
|
||||
// content flag and the user did not provide an explicit value.
|
||||
//
|
||||
// primaryContent indicates this flag is the default stdin receiver for the
|
||||
// command (e.g. --text for chat send). When true and the flag value is empty,
|
||||
// stdin pipe data is used automatically.
|
||||
func resolveStringFlag(cmd *cobra.Command, flagName string, guard *cli.StdinGuard, primaryContent bool) (string, error) {
|
||||
raw, err := cmd.Flags().GetString(flagName)
|
||||
if err != nil {
|
||||
return "", apperrors.NewInternal("failed to read --" + flagName)
|
||||
}
|
||||
|
||||
// Resolve @file / @- syntax.
|
||||
resolved, err := cli.ResolveInputSource(raw, flagName, guard)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
// Implicit stdin fallback: only for the primary content flag, only when
|
||||
// the user did not provide an explicit value and stdin is unclaimed.
|
||||
if resolved == "" && primaryContent && !guard.Claimed() && cli.StdinIsPipe() {
|
||||
if claimErr := guard.Claim("implicit stdin → --" + flagName); claimErr != nil {
|
||||
return "", claimErr
|
||||
}
|
||||
return cli.ReadStdin()
|
||||
}
|
||||
|
||||
return resolved, nil
|
||||
}
|
||||
|
||||
func commandDryRun(cmd *cobra.Command) bool {
|
||||
if cmd == nil {
|
||||
return false
|
||||
|
||||
@@ -1,217 +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 helpers
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
RegisterPublic(func() Handler {
|
||||
return creditHelper{}
|
||||
})
|
||||
}
|
||||
|
||||
type creditHelper struct{}
|
||||
|
||||
func (creditHelper) Name() string {
|
||||
return "credit"
|
||||
}
|
||||
|
||||
func (creditHelper) Command(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "credit",
|
||||
Short: "Enterprise credit search and risk helpers",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
cmd.AddCommand(
|
||||
newCreditSearchCommand(runner),
|
||||
newCreditRiskCommand(runner),
|
||||
newCreditEquityCommand(runner),
|
||||
)
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newCreditSearchCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "search",
|
||||
Short: "Enterprise name search",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name, err := cmd.Flags().GetString("name")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return fmt.Errorf("--name is required")
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"company_name": name,
|
||||
}
|
||||
if cmd.Flags().Changed("page") {
|
||||
page, err := cmd.Flags().GetInt("page")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["page"] = page
|
||||
}
|
||||
if cmd.Flags().Changed("size") {
|
||||
size, err := cmd.Flags().GetInt("size")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["size"] = size
|
||||
}
|
||||
|
||||
return runHelper(cmd, runner, "credit-ep", "ep_info_search_query", params)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("name", "", "Enterprise name keyword")
|
||||
cmd.Flags().Int("page", 0, "Page offset")
|
||||
cmd.Flags().Int("size", 0, "Page size")
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newCreditRiskCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "risk",
|
||||
Short: "Enterprise risk information",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
court := &cobra.Command{
|
||||
Use: "court",
|
||||
Short: "Court notice",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cert, err := cmd.Flags().GetString("cert")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cert = strings.TrimSpace(cert)
|
||||
if cert == "" {
|
||||
return fmt.Errorf("--cert is required")
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"ep_cert_no": cert,
|
||||
}
|
||||
if cmd.Flags().Changed("page") {
|
||||
page, err := cmd.Flags().GetInt("page")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["page"] = page
|
||||
}
|
||||
if cmd.Flags().Changed("size") {
|
||||
size, err := cmd.Flags().GetInt("size")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["size"] = size
|
||||
}
|
||||
|
||||
return runHelper(cmd, runner, "credit-risk", "ep_dossier_courtnotice_query", params)
|
||||
},
|
||||
}
|
||||
court.Flags().String("cert", "", "Enterprise registration number or credit code")
|
||||
court.Flags().Int("page", 0, "Page offset")
|
||||
court.Flags().Int("size", 0, "Page size")
|
||||
cmd.AddCommand(court)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newCreditEquityCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "equity",
|
||||
Short: "Enterprise equity information",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
shareholder := &cobra.Command{
|
||||
Use: "shareholder",
|
||||
Short: "Shareholder information",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
cert, err := cmd.Flags().GetString("cert")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
cert = strings.TrimSpace(cert)
|
||||
if cert == "" {
|
||||
return fmt.Errorf("--cert is required")
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"ep_cert_no": cert,
|
||||
}
|
||||
if cmd.Flags().Changed("page") {
|
||||
page, err := cmd.Flags().GetInt("page")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["page"] = page
|
||||
}
|
||||
if cmd.Flags().Changed("size") {
|
||||
size, err := cmd.Flags().GetInt("size")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["size"] = size
|
||||
}
|
||||
|
||||
return runHelper(cmd, runner, "credit-equity", "ep_dossier_shareholder_query", params)
|
||||
},
|
||||
}
|
||||
shareholder.Flags().String("cert", "", "Enterprise registration number or credit code")
|
||||
shareholder.Flags().Int("page", 0, "Page offset")
|
||||
shareholder.Flags().Int("size", 0, "Page size")
|
||||
cmd.AddCommand(shareholder)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runHelper(cmd *cobra.Command, runner executor.Runner, canonicalProduct, tool string, params map[string]any) error {
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(cobracmd.LegacyCommandPath(cmd), canonicalProduct, tool, params))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
@@ -1,379 +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 helpers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
type dingtalkVendorHandler struct {
|
||||
subtree string
|
||||
}
|
||||
|
||||
func init() {
|
||||
for _, subtree := range []string{"discovery", "oa-plus", "ai-sincere-hire"} {
|
||||
subtree := subtree
|
||||
RegisterHiddenDingTalk(func() Handler {
|
||||
return dingtalkVendorHandler{subtree: subtree}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func (d dingtalkVendorHandler) Name() string {
|
||||
return d.subtree
|
||||
}
|
||||
|
||||
func (d dingtalkVendorHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
switch d.subtree {
|
||||
case "discovery":
|
||||
return newDiscoveryCommand(runner)
|
||||
case "oa-plus":
|
||||
return newOAPlusCommand(runner)
|
||||
case "ai-sincere-hire":
|
||||
return newAISincereHireCommand(runner)
|
||||
default:
|
||||
return newHiddenGroup("dingtalk", "Hidden DingTalk vendor extensions")
|
||||
}
|
||||
}
|
||||
|
||||
func registerDingTalkFlags(cmd *cobra.Command) {
|
||||
cmd.PersistentFlags().String("json", "", "NewsFeedPushRequest JSON payload")
|
||||
cmd.PersistentFlags().String("source", "", "Crawl source ID")
|
||||
cmd.PersistentFlags().String("filenames", "", "Comma-separated file names")
|
||||
cmd.PersistentFlags().String("keyword", "", "Subscription keyword")
|
||||
cmd.PersistentFlags().String("instance-id", "", "Approval instance ID")
|
||||
cmd.PersistentFlags().String("process-code", "", "Approval process code")
|
||||
cmd.PersistentFlags().String("size", "20", "Result size")
|
||||
cmd.PersistentFlags().String("cursor", "", "Cursor")
|
||||
}
|
||||
|
||||
func newHiddenGroup(use, short string) *cobra.Command {
|
||||
return cobracmd.NewHiddenGroupCommand(use, short)
|
||||
}
|
||||
|
||||
func newVisibleGroup(use, short string) *cobra.Command {
|
||||
return cobracmd.NewGroupCommand(use, short)
|
||||
}
|
||||
|
||||
func newDiscoveryCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := newVisibleGroup("discovery", "Content discovery")
|
||||
registerDingTalkFlags(cmd)
|
||||
media := newVisibleGroup("media", "Media content")
|
||||
media.AddCommand(&cobra.Command{
|
||||
Use: "save",
|
||||
Short: "Save media content",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runDiscoveryMediaSave(cmd, runner)
|
||||
},
|
||||
})
|
||||
media.AddCommand(&cobra.Command{
|
||||
Use: "subscribe",
|
||||
Short: "Subscribe to media source",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runDiscoveryMediaSubscribe(cmd, runner)
|
||||
},
|
||||
})
|
||||
|
||||
oss := newVisibleGroup("oss", "Upload credentials")
|
||||
oss.AddCommand(&cobra.Command{
|
||||
Use: "get-upload-url",
|
||||
Short: "Get OSS upload URL",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runDiscoveryOSSGetUploadURL(cmd, runner)
|
||||
},
|
||||
})
|
||||
|
||||
subscribe := newVisibleGroup("subscribe", "Keyword subscription")
|
||||
subscribe.AddCommand(&cobra.Command{
|
||||
Use: "save",
|
||||
Short: i18n.T("保存Keyword subscription规则"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runDiscoverySubscribeSave(cmd, runner)
|
||||
},
|
||||
})
|
||||
|
||||
cmd.AddCommand(media, oss, subscribe)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runDiscoveryMediaSave(cmd *cobra.Command, runner executor.Runner) error {
|
||||
payload, err := cmd.Flags().GetString("json")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --json")
|
||||
}
|
||||
if strings.TrimSpace(payload) == "" {
|
||||
return errors.NewValidation("--json is required")
|
||||
}
|
||||
|
||||
var request any
|
||||
if err := json.Unmarshal([]byte(payload), &request); err != nil {
|
||||
return errors.NewValidation(fmt.Sprintf("--json must be valid JSON: %v", err))
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk discovery media save",
|
||||
"dingtalk-discovery",
|
||||
"save_video_and_image",
|
||||
map[string]any{
|
||||
"NewsFeedPushRequest": request,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runDiscoveryMediaSubscribe(cmd *cobra.Command, runner executor.Runner) error {
|
||||
source, err := cmd.Flags().GetString("source")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --source")
|
||||
}
|
||||
if strings.TrimSpace(source) == "" {
|
||||
return errors.NewValidation("--source is required")
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk discovery media subscribe",
|
||||
"dingtalk-discovery",
|
||||
"save_media_subscribe_rule",
|
||||
map[string]any{
|
||||
"skillCrawlId": source,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runDiscoveryOSSGetUploadURL(cmd *cobra.Command, runner executor.Runner) error {
|
||||
filenames, err := cmd.Flags().GetString("filenames")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --filenames")
|
||||
}
|
||||
if strings.TrimSpace(filenames) == "" {
|
||||
return errors.NewValidation("--filenames is required")
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk discovery oss get-upload-url",
|
||||
"dingtalk-discovery",
|
||||
"batch_get_oss_temp_upload_url",
|
||||
map[string]any{
|
||||
"filenames": filenames,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runDiscoverySubscribeSave(cmd *cobra.Command, runner executor.Runner) error {
|
||||
keyword, err := cmd.Flags().GetString("keyword")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --keyword")
|
||||
}
|
||||
if strings.TrimSpace(keyword) == "" {
|
||||
return errors.NewValidation("--keyword is required")
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk discovery subscribe save",
|
||||
"dingtalk-discovery",
|
||||
"save_keyword_subscribe_rule",
|
||||
map[string]any{
|
||||
"keyword": keyword,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func newOAPlusCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := newVisibleGroup("oa-plus", "OA approval enhanced")
|
||||
registerDingTalkFlags(cmd)
|
||||
approval := newVisibleGroup("approval", i18n.T("审批实例管理"))
|
||||
approval.AddCommand(&cobra.Command{
|
||||
Use: "get",
|
||||
Short: i18n.T("获取审批实例详情"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runOAPlusApprovalGet(cmd, runner)
|
||||
},
|
||||
})
|
||||
approval.AddCommand(&cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("分页查询审批实例"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runOAPlusApprovalList(cmd, runner)
|
||||
},
|
||||
})
|
||||
cmd.AddCommand(approval)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runOAPlusApprovalGet(cmd *cobra.Command, runner executor.Runner) error {
|
||||
instanceID, err := cmd.Flags().GetString("instance-id")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --instance-id")
|
||||
}
|
||||
if strings.TrimSpace(instanceID) == "" {
|
||||
return errors.NewValidation("--instance-id is required")
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk oa-plus approval get",
|
||||
"dingtalk-oa-plus",
|
||||
"get_approval_instance",
|
||||
map[string]any{
|
||||
"instanceId": instanceID,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runOAPlusApprovalList(cmd *cobra.Command, runner executor.Runner) error {
|
||||
processCode, err := cmd.Flags().GetString("process-code")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --process-code")
|
||||
}
|
||||
if strings.TrimSpace(processCode) == "" {
|
||||
return errors.NewValidation("--process-code is required")
|
||||
}
|
||||
|
||||
size, err := cmd.Flags().GetString("size")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --size")
|
||||
}
|
||||
cursor, err := cmd.Flags().GetString("cursor")
|
||||
if err != nil {
|
||||
return errors.NewInternal("failed to read --cursor")
|
||||
}
|
||||
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk oa-plus approval list",
|
||||
"dingtalk-oa-plus",
|
||||
"list_approval_instances",
|
||||
map[string]any{
|
||||
"processCode": processCode,
|
||||
"size": size,
|
||||
"cursor": cursor,
|
||||
},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func newAISincereHireCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := newVisibleGroup("ai-sincere-hire", i18n.T("AI诚聘"))
|
||||
registerDingTalkFlags(cmd)
|
||||
cmd.AddCommand(&cobra.Command{
|
||||
Use: "guide",
|
||||
Short: i18n.T("获取使用指引"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runAISincereHireGuide(cmd, runner)
|
||||
},
|
||||
})
|
||||
job := newVisibleGroup("job", i18n.T("岗位查询"))
|
||||
job.AddCommand(&cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("查询在招岗位"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runAISincereHireJobList(cmd, runner)
|
||||
},
|
||||
})
|
||||
talent := newVisibleGroup("talent", i18n.T("人才查询"))
|
||||
talent.AddCommand(&cobra.Command{
|
||||
Use: "list",
|
||||
Short: i18n.T("查询入职人才"),
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return runAISincereHireTalentList(cmd, runner)
|
||||
},
|
||||
})
|
||||
cmd.AddCommand(job, talent)
|
||||
return cmd
|
||||
}
|
||||
|
||||
func runAISincereHireGuide(cmd *cobra.Command, runner executor.Runner) error {
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk ai-sincere-hire guide",
|
||||
"dingtalk-ai-sincere-hire",
|
||||
"query_guide_url",
|
||||
map[string]any{},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runAISincereHireJobList(cmd *cobra.Command, runner executor.Runner) error {
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk ai-sincere-hire job list",
|
||||
"dingtalk-ai-sincere-hire",
|
||||
"query_opening_job_list",
|
||||
map[string]any{},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
|
||||
func runAISincereHireTalentList(cmd *cobra.Command, runner executor.Runner) error {
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
"dingtalk ai-sincere-hire talent list",
|
||||
"dingtalk-ai-sincere-hire",
|
||||
"query_success_talent_list",
|
||||
map[string]any{},
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
}
|
||||
@@ -50,15 +50,9 @@ const (
|
||||
nameMaxLen = 50
|
||||
)
|
||||
|
||||
type hiddenExtensionFactory struct {
|
||||
manifest Manifest
|
||||
factory Factory
|
||||
}
|
||||
|
||||
var (
|
||||
registryMu sync.Mutex
|
||||
publicFactories []Factory
|
||||
hiddenVendorFactories []hiddenExtensionFactory
|
||||
registryMu sync.Mutex
|
||||
publicFactories []Factory
|
||||
)
|
||||
|
||||
func RegisterPublic(factory Factory) {
|
||||
@@ -67,99 +61,10 @@ func RegisterPublic(factory Factory) {
|
||||
publicFactories = append(publicFactories, factory)
|
||||
}
|
||||
|
||||
func RegisterHiddenVendor(vendor string, factory Factory) {
|
||||
if factory == nil {
|
||||
panic("helpers: hidden vendor factory is nil")
|
||||
}
|
||||
handler := factory()
|
||||
if handler == nil {
|
||||
panic("helpers: hidden vendor handler is nil")
|
||||
}
|
||||
|
||||
manifest := Manifest{
|
||||
Vendor: strings.TrimSpace(vendor),
|
||||
Name: strings.TrimSpace(handler.Name()),
|
||||
}
|
||||
if err := ValidateNaming(manifest.Vendor, manifest.Name); err != nil {
|
||||
panic(fmt.Sprintf("helpers: invalid hidden vendor extension %s: %v", manifest.FullName(), err))
|
||||
}
|
||||
|
||||
registryMu.Lock()
|
||||
defer registryMu.Unlock()
|
||||
hiddenVendorFactories = append(hiddenVendorFactories, hiddenExtensionFactory{
|
||||
manifest: manifest,
|
||||
factory: factory,
|
||||
})
|
||||
}
|
||||
|
||||
func RegisterHiddenDingTalk(factory Factory) {
|
||||
RegisterHiddenVendor("dingtalk", factory)
|
||||
}
|
||||
|
||||
func NewPublicCommands(runner executor.Runner) []*cobra.Command {
|
||||
return buildCommands(publicFactories, runner)
|
||||
}
|
||||
|
||||
func NewHiddenVendorCommands(runner executor.Runner) []*cobra.Command {
|
||||
registryMu.Lock()
|
||||
factories := append([]hiddenExtensionFactory(nil), hiddenVendorFactories...)
|
||||
registryMu.Unlock()
|
||||
|
||||
if len(factories) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
byVendor := make(map[string][]*cobra.Command)
|
||||
for _, registered := range factories {
|
||||
handler := registered.factory()
|
||||
if handler == nil {
|
||||
continue
|
||||
}
|
||||
command := handler.Command(runner)
|
||||
if command == nil {
|
||||
continue
|
||||
}
|
||||
byVendor[registered.manifest.Vendor] = append(byVendor[registered.manifest.Vendor], command)
|
||||
}
|
||||
|
||||
vendors := make([]string, 0, len(byVendor))
|
||||
for vendor := range byVendor {
|
||||
vendors = append(vendors, vendor)
|
||||
}
|
||||
sort.Strings(vendors)
|
||||
|
||||
roots := make([]*cobra.Command, 0, len(vendors))
|
||||
for _, vendor := range vendors {
|
||||
commands := byVendor[vendor]
|
||||
sort.Slice(commands, func(i, j int) bool {
|
||||
return commands[i].Use < commands[j].Use
|
||||
})
|
||||
root := &cobra.Command{
|
||||
Use: vendor,
|
||||
Short: fmt.Sprintf("Hidden %s vendor extensions", vendor),
|
||||
Hidden: true,
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
root.AddCommand(commands...)
|
||||
roots = append(roots, root)
|
||||
}
|
||||
return roots
|
||||
}
|
||||
|
||||
func NewHiddenDingTalkCommand(runner executor.Runner) *cobra.Command {
|
||||
for _, root := range NewHiddenVendorCommands(runner) {
|
||||
if root != nil && root.Name() == "dingtalk" {
|
||||
return root
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func buildCommands(factories []Factory, runner executor.Runner) []*cobra.Command {
|
||||
registryMu.Lock()
|
||||
defer registryMu.Unlock()
|
||||
|
||||
@@ -13,11 +13,7 @@
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
import "testing"
|
||||
|
||||
func TestValidateNaming(t *testing.T) {
|
||||
t.Parallel()
|
||||
@@ -50,28 +46,3 @@ func TestValidateNaming(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewHiddenVendorCommandsIncludesDingTalk(t *testing.T) {
|
||||
roots := NewHiddenVendorCommands(executor.EchoRunner{})
|
||||
if len(roots) == 0 {
|
||||
t.Fatal("NewHiddenVendorCommands() = empty, want hidden vendor roots")
|
||||
}
|
||||
|
||||
var dingtalkFound bool
|
||||
for _, root := range roots {
|
||||
if root == nil || root.Name() != "dingtalk" {
|
||||
continue
|
||||
}
|
||||
dingtalkFound = true
|
||||
if !root.Hidden {
|
||||
t.Fatalf("dingtalk root hidden = false, want true")
|
||||
}
|
||||
if len(root.Commands()) == 0 {
|
||||
t.Fatalf("dingtalk root children = 0, want registered hidden extensions")
|
||||
}
|
||||
}
|
||||
|
||||
if !dingtalkFound {
|
||||
t.Fatal("NewHiddenVendorCommands() missing dingtalk root")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,515 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func init() {
|
||||
RegisterPublic(func() Handler {
|
||||
return reportHandler{}
|
||||
})
|
||||
}
|
||||
|
||||
type reportHandler struct{}
|
||||
|
||||
func (reportHandler) Name() string {
|
||||
return "report"
|
||||
}
|
||||
|
||||
func (reportHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
root := &cobra.Command{
|
||||
Use: "report",
|
||||
Aliases: []string{"log"},
|
||||
Short: "日志 / 模版 / 统计",
|
||||
Long: `钉钉日志:模版、创建、详情、列表、统计。
|
||||
|
||||
子命令:
|
||||
template 日志模版(list / detail)
|
||||
create 创建日志
|
||||
detail 获取日志详情
|
||||
list 查询收到的日志列表
|
||||
stats 获取日志统计数据
|
||||
sent 查询已发送的日志列表`,
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
|
||||
template := &cobra.Command{
|
||||
Use: "template",
|
||||
Short: "日志模版",
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
template.AddCommand(
|
||||
newReportTemplateListCommand(runner),
|
||||
newReportTemplateDetailCommand(runner),
|
||||
)
|
||||
|
||||
root.AddCommand(
|
||||
template,
|
||||
newReportCreateCommand(runner),
|
||||
newReportDetailCommand(runner),
|
||||
newReportListCommand(runner),
|
||||
newReportStatsCommand(runner),
|
||||
newReportSentCommand(runner),
|
||||
)
|
||||
return root
|
||||
}
|
||||
|
||||
// ── flexTimeLayouts: supported date formats, most specific first ──
|
||||
|
||||
var flexTimeLayouts = []string{
|
||||
time.RFC3339, // 2006-01-02T15:04:05+08:00
|
||||
"2006-01-02T15:04:05Z", // UTC Z suffix
|
||||
"2006-01-02T15:04:05-07:00", // with offset but no colon
|
||||
"2006-01-02T15:04:05", // no timezone
|
||||
"2006-01-02 15:04:05", // space-separated
|
||||
"2006-01-02T15:04", // no seconds
|
||||
"2006-01-02 15:04", // no seconds, space
|
||||
"2006-01-02", // date only
|
||||
"2006/01/02 15:04:05", // slash + time
|
||||
"2006/01/02", // slash date
|
||||
"20060102", // compact YYYYMMDD
|
||||
}
|
||||
|
||||
// parseFlexTimeToMillis parses a date string using multiple formats and returns Unix milliseconds.
|
||||
// Supports 11 formats for maximum compatibility with user input.
|
||||
func parseFlexTimeToMillis(flagName, value string) (int64, error) {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
return 0, apperrors.NewValidation(fmt.Sprintf(
|
||||
"--%s is required\n hint: example: 2026-03-10T14:00:00+08:00", flagName))
|
||||
}
|
||||
loc, _ := time.LoadLocation("Asia/Shanghai")
|
||||
if loc == nil {
|
||||
loc = time.Local
|
||||
}
|
||||
for _, layout := range flexTimeLayouts {
|
||||
t, err := time.ParseInLocation(layout, value, loc)
|
||||
if err == nil {
|
||||
return t.UnixMilli(), nil
|
||||
}
|
||||
}
|
||||
return 0, apperrors.NewValidation(fmt.Sprintf(
|
||||
"cannot parse time for --%s (input: %q)\n hint: supported formats: 2026-03-23T14:00:00+08:00, 2026-03-23 14:00:00, 2026-03-23",
|
||||
flagName, value))
|
||||
}
|
||||
|
||||
// validateTimeRange checks that endMs is strictly after startMs.
|
||||
func validateTimeRange(startMs, endMs int64) error {
|
||||
if endMs <= startMs {
|
||||
return apperrors.NewValidation("--end must be after --start\n hint: swap the values or adjust the time range")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ── template list ──────────────────────────────────────────
|
||||
|
||||
func newReportTemplateListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "获取当前用户可用的日志模版列表",
|
||||
Example: " dws report template list",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
params := map[string]any{}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_available_report_templates", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── template detail ────────────────────────────────────────
|
||||
|
||||
func newReportTemplateDetailCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "detail",
|
||||
Short: "获取日志模版详情",
|
||||
Example: " dws report template detail --name <templateName>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
name, _ := cmd.Flags().GetString("name")
|
||||
if name == "" {
|
||||
return apperrors.NewValidation("--name is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_template_name": name,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_template_details_by_name", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("name", "", "模版名称 (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── create ─────────────────────────────────────────────────
|
||||
|
||||
func newReportCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create",
|
||||
Short: "创建日志",
|
||||
Long: `按模版创建一条日志。--contents 为 JSON 数组,每项需含 key、sort、content、contentType、type,
|
||||
与远程 create_report 一致;可先通过 report template list / template detail 取得 templateId 与控件定义。`,
|
||||
Example: ` dws report create --template-id TPL_ID --contents '[{"content":"完成开发","sort":"0","key":"今日完成","contentType":"markdown","type":"1"}]'
|
||||
dws report create --template-id TPL_ID --contents '[...]' --to-chat --to-user-ids userId1,userId2`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
tplID, _ := cmd.Flags().GetString("template-id")
|
||||
if tplID == "" {
|
||||
return apperrors.NewValidation("--template-id is required")
|
||||
}
|
||||
contentsJSON, _ := cmd.Flags().GetString("contents")
|
||||
if contentsJSON == "" {
|
||||
return apperrors.NewValidation("--contents is required")
|
||||
}
|
||||
var contents []map[string]any
|
||||
if err := json.Unmarshal([]byte(contentsJSON), &contents); err != nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("--contents JSON parse failed: %v", err))
|
||||
}
|
||||
ddFrom, _ := cmd.Flags().GetString("dd-from")
|
||||
if ddFrom == "" {
|
||||
ddFrom = "dws"
|
||||
}
|
||||
toChat, _ := cmd.Flags().GetBool("to-chat")
|
||||
params := map[string]any{
|
||||
"templateId": tplID,
|
||||
"contents": contents,
|
||||
"ddFrom": ddFrom,
|
||||
"toChat": toChat,
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("to-user-ids"); v != "" {
|
||||
params["toUserIds"] = parseUserIDs(v)
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "create_report", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("template-id", "", "日志模版 ID (必填)")
|
||||
cmd.Flags().String("contents", "", "日志内容 JSON 数组 (必填),每项含 key/sort/content/contentType/type")
|
||||
cmd.Flags().String("dd-from", "dws", "创建来源标识")
|
||||
cmd.Flags().Bool("to-chat", false, "是否发送到日志接收人单聊")
|
||||
cmd.Flags().String("to-user-ids", "", "接收人 userId,逗号分隔 (可选)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── detail ─────────────────────────────────────────────────
|
||||
|
||||
func newReportDetailCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "detail",
|
||||
Short: "获取日志详情",
|
||||
Example: " dws report detail --report-id <reportId>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
reportID, _ := cmd.Flags().GetString("report-id")
|
||||
if reportID == "" {
|
||||
return apperrors.NewValidation("--report-id is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_id": reportID,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_entry_details", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("report-id", "", "日志 ID (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── list (received reports) ────────────────────────────────
|
||||
// Key fix: cursor defaults to 0, size defaults to 20, flexible date parsing
|
||||
|
||||
func newReportListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
Short: "查询当前人收到的日志列表",
|
||||
Example: ` dws report list --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
|
||||
dws report list --start "2026-03-10 00:00:00" --end "2026-03-10 23:59:59" --cursor 0 --size 20`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
startStr, _ := cmd.Flags().GetString("start")
|
||||
endStr, _ := cmd.Flags().GetString("end")
|
||||
|
||||
startMs, err := parseFlexTimeToMillis("start", startStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
endMs, err := parseFlexTimeToMillis("end", endStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateTimeRange(startMs, endMs); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// cursor defaults to 0, size defaults to 20
|
||||
cursor, _ := cmd.Flags().GetInt("cursor")
|
||||
size, _ := cmd.Flags().GetInt("size")
|
||||
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
|
||||
size = v
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"startTime": float64(startMs),
|
||||
"endTime": float64(endMs),
|
||||
"cursor": float64(cursor),
|
||||
"size": float64(size),
|
||||
}
|
||||
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_received_report_list", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("start", "", "开始时间 ISO-8601 (如 2026-03-10T00:00:00+08:00) (必填)")
|
||||
cmd.Flags().String("end", "", "结束时间 ISO-8601 (如 2026-03-10T23:59:59+08:00) (必填)")
|
||||
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
|
||||
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
|
||||
cmd.Flags().Int("limit", 0, "--size 的别名")
|
||||
_ = cmd.Flags().MarkHidden("limit")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── stats ──────────────────────────────────────────────────
|
||||
|
||||
func newReportStatsCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "stats",
|
||||
Short: "获取日志统计数据",
|
||||
Example: " dws report stats --report-id <reportId>",
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
reportID, _ := cmd.Flags().GetString("report-id")
|
||||
if reportID == "" {
|
||||
return apperrors.NewValidation("--report-id is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"report_id": reportID,
|
||||
}
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_report_statistics_by_id", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("report-id", "", "日志 ID (必填)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── sent (my created reports) ──────────────────────────────
|
||||
// Key fix: cursor defaults to 0, size defaults to 20, start/end default to last 30 days
|
||||
|
||||
func newReportSentCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "sent",
|
||||
Short: "查询当前人创建的日志列表",
|
||||
Example: ` dws report sent
|
||||
dws report sent --cursor 0 --size 20
|
||||
dws report sent --start "2026-03-10T00:00:00+08:00" --end "2026-03-10T23:59:59+08:00"
|
||||
dws report sent --template-name "日报"`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
// cursor defaults to 0, size defaults to 20
|
||||
cursor, _ := cmd.Flags().GetInt("cursor")
|
||||
size, _ := cmd.Flags().GetInt("size")
|
||||
if v, _ := cmd.Flags().GetInt("limit"); v > 0 && !cmd.Flags().Changed("size") {
|
||||
size = v
|
||||
}
|
||||
|
||||
params := map[string]any{
|
||||
"cursor": float64(cursor),
|
||||
"size": float64(size),
|
||||
}
|
||||
|
||||
// Default time range: last 30 days
|
||||
now := time.Now()
|
||||
startDefault := now.AddDate(0, 0, -30).Truncate(24 * time.Hour).Format(time.RFC3339)
|
||||
endDefault := time.Date(now.Year(), now.Month(), now.Day(), 23, 59, 59, 0, now.Location()).Format(time.RFC3339)
|
||||
|
||||
startStr, _ := cmd.Flags().GetString("start")
|
||||
if startStr == "" {
|
||||
startStr = startDefault
|
||||
}
|
||||
endStr, _ := cmd.Flags().GetString("end")
|
||||
if endStr == "" {
|
||||
endStr = endDefault
|
||||
}
|
||||
|
||||
startMs, err := parseFlexTimeToMillis("start", startStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["startTime"] = float64(startMs)
|
||||
|
||||
endMs, err := parseFlexTimeToMillis("end", endStr)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["endTime"] = float64(endMs)
|
||||
|
||||
if err := validateTimeRange(startMs, endMs); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Optional modified time filters
|
||||
if v, _ := cmd.Flags().GetString("modified-start"); v != "" {
|
||||
ms, err := parseFlexTimeToMillis("modified-start", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["modifiedStartTime"] = float64(ms)
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("modified-end"); v != "" {
|
||||
ms, err := parseFlexTimeToMillis("modified-end", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
params["modifiedEndTime"] = float64(ms)
|
||||
}
|
||||
|
||||
// Optional template name filter
|
||||
if v, _ := cmd.Flags().GetString("template-name"); v != "" {
|
||||
params["report_template_name"] = v
|
||||
}
|
||||
|
||||
if commandDryRun(cmd) {
|
||||
return writeCommandPayload(cmd, executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
|
||||
))
|
||||
}
|
||||
result, err := runner.Run(cmd.Context(), executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd), "report", "get_send_report_list", params,
|
||||
))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
cmd.Flags().Int("cursor", 0, "分页游标,首次传 0 (默认 0)")
|
||||
cmd.Flags().Int("size", 20, "每页条数,最大 20 (默认 20)")
|
||||
cmd.Flags().Int("limit", 0, "--size 的别名")
|
||||
_ = cmd.Flags().MarkHidden("limit")
|
||||
cmd.Flags().String("start", "", "创建开始时间 ISO-8601 (默认最近 30 天)")
|
||||
cmd.Flags().String("end", "", "创建结束时间 ISO-8601 (默认最近 30 天)")
|
||||
cmd.Flags().String("modified-start", "", "修改开始时间 ISO-8601 (可选)")
|
||||
cmd.Flags().String("modified-end", "", "修改结束时间 ISO-8601 (可选)")
|
||||
cmd.Flags().String("template-name", "", "日志模板名称 (可选,不传查全部)")
|
||||
preferLegacyLeaf(cmd)
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── helpers ────────────────────────────────────────────────
|
||||
|
||||
func parseUserIDs(s string) []string {
|
||||
parts := strings.Split(s, ",")
|
||||
out := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p != "" {
|
||||
out = append(out, p)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package helpers
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestParseFlexTimeToMillis_RFC3339(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-10T00:00:00+08:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_SpaceSeparated(t *testing.T) {
|
||||
t.Parallel()
|
||||
// This is the format that was causing the HTTP 400 error
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-01 00:00:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_DateOnly(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-01")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_NoTimezone(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026-03-10T14:00:00")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_SlashFormat(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "2026/03/10")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_CompactFormat(t *testing.T) {
|
||||
t.Parallel()
|
||||
ms, err := parseFlexTimeToMillis("start", "20260310")
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if ms <= 0 {
|
||||
t.Errorf("expected positive milliseconds, got %d", ms)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_Empty(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := parseFlexTimeToMillis("start", "")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for empty value")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseFlexTimeToMillis_Invalid(t *testing.T) {
|
||||
t.Parallel()
|
||||
_, err := parseFlexTimeToMillis("start", "not-a-date")
|
||||
if err == nil {
|
||||
t.Fatal("expected error for invalid date")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Valid(t *testing.T) {
|
||||
t.Parallel()
|
||||
start := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
end := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(start, end); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Invalid(t *testing.T) {
|
||||
t.Parallel()
|
||||
start := time.Date(2026, 3, 10, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
end := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(start, end); err == nil {
|
||||
t.Fatal("expected error when end is before start")
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateTimeRange_Equal(t *testing.T) {
|
||||
t.Parallel()
|
||||
ts := time.Date(2026, 3, 1, 0, 0, 0, 0, time.UTC).UnixMilli()
|
||||
if err := validateTimeRange(ts, ts); err == nil {
|
||||
t.Fatal("expected error when start equals end")
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseUserIDs(t *testing.T) {
|
||||
t.Parallel()
|
||||
tests := []struct {
|
||||
input string
|
||||
want int
|
||||
}{
|
||||
{"user1,user2,user3", 3},
|
||||
{"user1", 1},
|
||||
{"user1, user2, user3", 3},
|
||||
{"user1,,user2", 2},
|
||||
{"", 0},
|
||||
{" , , ", 0},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
got := parseUserIDs(tt.input)
|
||||
if len(got) != tt.want {
|
||||
t.Errorf("parseUserIDs(%q) = %d items, want %d", tt.input, len(got), tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
+307
-4
@@ -18,9 +18,10 @@ import (
|
||||
"strings"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -41,7 +42,8 @@ func (todoHandler) Name() string {
|
||||
func (todoHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
root := &cobra.Command{
|
||||
Use: "todo",
|
||||
Short: "Todo helper overrides",
|
||||
Short: i18n.T("待办任务管理"),
|
||||
Long: i18n.T("管理钉钉个人待办:创建、查询列表、查看详情、修改、标记完成、删除。"),
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
@@ -60,11 +62,94 @@ func (todoHandler) Command(runner executor.Runner) *cobra.Command {
|
||||
return cmd.Help()
|
||||
},
|
||||
}
|
||||
task.AddCommand(newTodoTaskListCommand(runner))
|
||||
task.AddCommand(
|
||||
newTodoTaskCreateCommand(runner),
|
||||
newTodoTaskListCommand(runner),
|
||||
newTodoTaskUpdateCommand(runner),
|
||||
newTodoTaskDoneCommand(runner),
|
||||
newTodoTaskGetCommand(runner),
|
||||
newTodoTaskDeleteCommand(runner),
|
||||
)
|
||||
root.AddCommand(task)
|
||||
return root
|
||||
}
|
||||
|
||||
// ── create ─────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskCreateCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "create",
|
||||
Short: i18n.T("创建待办"),
|
||||
Example: ` dws todo task create --title "修复线上Bug" --executors userId1,userId2 --priority 40
|
||||
dws todo task create --title "提交报告" --executors userId1 --due "2026-03-10T18:00:00+08:00"
|
||||
|
||||
# 查询 userId: dws contact user search --keyword "姓名"`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
title := cmdutil.FlagOrFallback(cmd, "title", "subject", "content")
|
||||
if strings.TrimSpace(title) == "" {
|
||||
return apperrors.NewValidation("--title is required")
|
||||
}
|
||||
executorsStr, _ := cmd.Flags().GetString("executors")
|
||||
if strings.TrimSpace(executorsStr) == "" {
|
||||
return apperrors.NewValidation("--executors is required")
|
||||
}
|
||||
executorIds := parseExecutorIds(executorsStr)
|
||||
|
||||
vo := map[string]any{
|
||||
"subject": title,
|
||||
"executorIds": executorIds,
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("due"); v != "" {
|
||||
ms, err := cmdutil.ParseISOTimeToMillis("due", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
vo["dueTime"] = ms
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("priority"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil {
|
||||
vo["priority"] = n
|
||||
}
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("recurrence"); v != "" {
|
||||
vo["recurrence"] = v
|
||||
}
|
||||
params := map[string]any{"PersonalTodoCreateVO": vo}
|
||||
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"todo",
|
||||
"create_personal_todo",
|
||||
params,
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("title", "", i18n.T("待办标题 (必填)"))
|
||||
cmd.Flags().String("executors", "", i18n.T("执行者 userId 列表 (必填)"))
|
||||
cmd.Flags().String("due", "", i18n.T("截止时间 ISO-8601 (如 2026-03-10T18:00:00+08:00)"))
|
||||
cmd.Flags().String("priority", "", i18n.T("优先级: 10低/20普通/30较高/40紧急"))
|
||||
cmd.Flags().String("recurrence", "", i18n.T("循环待办 (需先设置 --due); 格式: DTSTART:...\\nRRULE:FREQ=DAILY;INTERVAL=1"))
|
||||
|
||||
cmd.Flags().String("subject", "", i18n.T("--title 的别名"))
|
||||
cmd.Flags().String("content", "", i18n.T("--title 的别名"))
|
||||
_ = cmd.Flags().MarkHidden("subject")
|
||||
_ = cmd.Flags().MarkHidden("content")
|
||||
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── list ───────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskListCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "list",
|
||||
@@ -88,6 +173,7 @@ func newTodoTaskListCommand(runner executor.Runner) *cobra.Command {
|
||||
)
|
||||
|
||||
if size <= todoListPageSizeMax {
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -158,6 +244,223 @@ func newTodoTaskListCommand(runner executor.Runner) *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── update ─────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskUpdateCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "update",
|
||||
Short: i18n.T("修改待办任务"),
|
||||
Example: ` dws todo task update --task-id <taskId> --title "新标题"
|
||||
dws todo task update --task-id <taskId> --priority 40 --due "2026-03-10T18:00:00+08:00"
|
||||
dws todo task update --task-id <taskId> --done true
|
||||
|
||||
# 查询 taskId: dws todo task list`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
taskID, _ := cmd.Flags().GetString("task-id")
|
||||
if strings.TrimSpace(taskID) == "" {
|
||||
return apperrors.NewValidation("--task-id is required")
|
||||
}
|
||||
inner := map[string]any{
|
||||
"taskId": taskID,
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("title"); v != "" {
|
||||
inner["subject"] = v
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("due"); v != "" {
|
||||
ms, err := cmdutil.ParseISOTimeToMillis("due", v)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
inner["dueTime"] = ms
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("priority"); v != "" {
|
||||
if n, err := strconv.Atoi(v); err == nil {
|
||||
inner["priority"] = n
|
||||
}
|
||||
}
|
||||
if v, _ := cmd.Flags().GetString("done"); v != "" {
|
||||
inner["isDone"] = v == "true"
|
||||
}
|
||||
params := map[string]any{"TodoUpdateRequest": inner}
|
||||
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"todo",
|
||||
"update_todo_task",
|
||||
params,
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("task-id", "", i18n.T("待办任务 ID (必填)"))
|
||||
cmd.Flags().String("title", "", i18n.T("新标题"))
|
||||
cmd.Flags().String("due", "", i18n.T("截止时间 ISO-8601 (如 2026-03-10T18:00:00+08:00)"))
|
||||
cmd.Flags().String("priority", "", i18n.T("优先级: 10低/20普通/30较高/40紧急"))
|
||||
cmd.Flags().String("done", "", i18n.T("完成状态: true/false"))
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── done ───────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskDoneCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "done",
|
||||
Short: i18n.T("修改执行者的待办完成状态"),
|
||||
Example: ` dws todo task done --task-id <taskId> --status true
|
||||
dws todo task done --task-id <taskId> --status false
|
||||
|
||||
# 查询 taskId: dws todo task list`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
taskID, _ := cmd.Flags().GetString("task-id")
|
||||
if strings.TrimSpace(taskID) == "" {
|
||||
return apperrors.NewValidation("--task-id is required")
|
||||
}
|
||||
status, _ := cmd.Flags().GetString("status")
|
||||
if strings.TrimSpace(status) == "" {
|
||||
return apperrors.NewValidation("--status is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"taskId": taskID,
|
||||
"isDone": status,
|
||||
}
|
||||
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"todo",
|
||||
"update_todo_done_status",
|
||||
params,
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("task-id", "", i18n.T("待办任务 ID (必填)"))
|
||||
cmd.Flags().String("status", "", i18n.T("完成状态: true=已完成, false=未完成 (必填)"))
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── get ────────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskGetCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "get",
|
||||
Short: i18n.T("待办详情"),
|
||||
Example: ` dws todo task get --task-id <taskId>
|
||||
|
||||
# 查询 taskId: dws todo task list`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
taskID, _ := cmd.Flags().GetString("task-id")
|
||||
if strings.TrimSpace(taskID) == "" {
|
||||
return apperrors.NewValidation("--task-id is required")
|
||||
}
|
||||
params := map[string]any{
|
||||
"taskId": taskID,
|
||||
}
|
||||
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"todo",
|
||||
"query_todo_detail",
|
||||
params,
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("task-id", "", i18n.T("待办任务 ID (必填)"))
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── delete ─────────────────────────────────────────────────
|
||||
|
||||
func newTodoTaskDeleteCommand(runner executor.Runner) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "delete",
|
||||
Short: i18n.T("删除待办"),
|
||||
Example: ` dws todo task delete --task-id <taskId>
|
||||
dws todo task delete --task-id <taskId> --yes
|
||||
|
||||
# 查询 taskId: dws todo task list`,
|
||||
Args: cobra.NoArgs,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
taskID, _ := cmd.Flags().GetString("task-id")
|
||||
if strings.TrimSpace(taskID) == "" {
|
||||
return apperrors.NewValidation("--task-id is required")
|
||||
}
|
||||
if !confirmDeletePrompt(cmd, i18n.T("待办"), taskID) {
|
||||
return nil
|
||||
}
|
||||
params := map[string]any{
|
||||
"taskId": taskID,
|
||||
}
|
||||
|
||||
invocation := executor.NewHelperInvocation(
|
||||
cobracmd.LegacyCommandPath(cmd),
|
||||
"todo",
|
||||
"delete_todo",
|
||||
params,
|
||||
)
|
||||
invocation.DryRun = commandDryRun(cmd)
|
||||
result, err := runner.Run(cmd.Context(), invocation)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return writeCommandPayload(cmd, result)
|
||||
},
|
||||
}
|
||||
preferLegacyLeaf(cmd)
|
||||
|
||||
cmd.Flags().String("task-id", "", i18n.T("待办任务 ID (必填)"))
|
||||
cmd.Flags().Bool("yes", false, i18n.T("跳过确认直接删除"))
|
||||
return cmd
|
||||
}
|
||||
|
||||
// ── helpers ────────────────────────────────────────────────
|
||||
|
||||
// parseExecutorIds splits "id1,id2" into []string for the MCP executorIds array.
|
||||
func parseExecutorIds(s string) []string {
|
||||
s = strings.TrimSpace(s)
|
||||
if s == "" {
|
||||
return nil
|
||||
}
|
||||
parts := strings.Split(s, ",")
|
||||
ids := make([]string, 0, len(parts))
|
||||
for _, p := range parts {
|
||||
if id := strings.TrimSpace(p); id != "" {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
}
|
||||
return ids
|
||||
}
|
||||
|
||||
// ── list pagination helpers ────────────────────────────────
|
||||
|
||||
func normalizePage(raw string) string {
|
||||
if trimmed := strings.TrimSpace(raw); trimmed != "" {
|
||||
return trimmed
|
||||
|
||||
@@ -9,49 +9,51 @@
|
||||
"AI诚聘": "AI Recruitment",
|
||||
"Base ID (必填)": "Base ID (required)",
|
||||
"MCP 服务返回了无法解析的协议响应;检查服务版本或上游代理。": "MCP service returned an unparseable protocol response; check service version or upstream proxy.",
|
||||
"OSS 上传失败 HTTP %d: %s": "OSS uploadfailed HTTP %d: %s",
|
||||
"OSS 上传失败 HTTP %d: %s": "OSS upload failed HTTP %d: %s",
|
||||
"[%d] 轮询中... (%ds)": "[%d] polling ... (%ds)",
|
||||
"prepare_attachment_upload 返回格式异常": "prepare_attachment_upload returned abnormal format",
|
||||
"refresh_token 刷新失败": "refresh_token refreshfailed",
|
||||
"refresh_token 刷新失败": "refresh_token refresh failed",
|
||||
"refresh_token 刷新失败,将尝试扫码登录": "refresh_token refresh failed, will attempt QR code login",
|
||||
"true=已完成, false=未完成": "true=alreadycompleted, false=un-completed",
|
||||
"true=已完成, false=未完成": "true=completed, false=not completed",
|
||||
"⏳ 等待授权中...": "⏳ Waiting for authorization...",
|
||||
"⚠️ 即将删除 %s: %s\\\n": "⚠️ About to delete %s: %s\\\n",
|
||||
"上传失败: %w": "uploadfailed: %w",
|
||||
"上传失败: %w": "upload failed: %w",
|
||||
"上游服务异常;可稍后重试,若持续失败请查看 recovery snapshot。": "Upstream service error; retry later. If persistent, check recovery snapshot.",
|
||||
"下载失败 (HTTP %d): %s": "downloadfailed (HTTP %d): %s",
|
||||
"下载失败 (HTTP %d): %s": "download failed (HTTP %d): %s",
|
||||
"下载失败: %w": "Download failed: %w",
|
||||
"不是文件: ": "notisfile:",
|
||||
"不是文件: ": "not a file: ",
|
||||
"人才查询": "Talent query",
|
||||
"使用授权码换取 Access Token...": "Exchanging authorization code for Access Token...",
|
||||
"保存 token 失败": "Failed to save token",
|
||||
"保存Keyword subscription规则": "saveKeyword subscriptionrule",
|
||||
"写入文件失败: %w": "writefilefailed: %w",
|
||||
"准备上传失败: %w": "prepareuploadfailed: %w",
|
||||
"分页查询审批实例": "paginated queryapproval instance",
|
||||
"创建 / 查询 / 更新 / 删除待办": "create / query / update / deletetodo",
|
||||
"创建文件失败: %w": "createfilefailed: %w",
|
||||
"保存Keyword subscription规则": "Save keyword subscription rule",
|
||||
"写入文件失败: %w": "write file failed: %w",
|
||||
"准备上传失败: %w": "prepare upload failed: %w",
|
||||
"分页查询审批实例": "Paginated query approval instances",
|
||||
"创建 / 查询 / 更新 / 删除待办": "Create / query / update / delete todo",
|
||||
"创建待办": "Create todo",
|
||||
"创建文件失败: %w": "create file failed: %w",
|
||||
"创建请求失败": "Failed to create request",
|
||||
"删除 AI 表格": "delete AI table",
|
||||
"删除原因(可选)": "deletereason(optional)",
|
||||
"删除字段": "deletefield",
|
||||
"删除待办": "Delete todo",
|
||||
"删除 AI 表格": "Delete AI table",
|
||||
"删除原因(可选)": "Delete reason (optional)",
|
||||
"删除字段": "Delete field",
|
||||
"删除指定 Base(高风险、不可逆)。使用 --yes 跳过确认。": "Delete specified Base (high risk, irreversible). Use --yes to skip confirmation.",
|
||||
"删除指定字段(高风险、不可逆)。使用 --yes 跳过确认。": "Delete specified field (high risk, irreversible). Use --yes to skip confirmation.",
|
||||
"删除指定数据表(高风险、不可逆)。使用 --yes 跳过确认。": "Delete specified data table (high risk, irreversible). Use --yes to skip confirmation.",
|
||||
"删除数据表": "deletedata table",
|
||||
"删除数据表": "Delete data table",
|
||||
"删除行记录": "Delete row records",
|
||||
"参数不符合工具输入 schema;请检查 --json/--params/flags。": "Parameters do not match tool input schema; check --json/--params/flags.",
|
||||
"发送请求失败": "Failed to send request",
|
||||
"回调中未收到授权码": "callback un-receivetoauthorization code",
|
||||
"回调中未收到授权码": "Callback did not receive authorization code",
|
||||
"如果浏览器未自动打开,请手动访问:\n %s\n\n": "If the browser did not open automatically, please visit:\n %s\n\n",
|
||||
"字段": "field",
|
||||
"字段 ID (必填)": "field ID (required)",
|
||||
"字段管理": "fieldmanagement",
|
||||
"审批实例管理": "approval instancemanagement",
|
||||
"字段管理": "Field management",
|
||||
"审批实例管理": "Approval instance management",
|
||||
"岗位查询": "Position query",
|
||||
"工具协议不兼容;请检查服务版本、工具名或刷新发现缓存。": "Tool protocol incompatible; check service version, tool name, or refresh discovery cache.",
|
||||
"工具调用失败;请检查参数和上游服务状态。": "Tool invocation failed; check parameters and upstream service status.",
|
||||
"已取消操作": "alreadycanceloperations",
|
||||
"已取消操作": "Operation cancelled",
|
||||
"当前平台 %s 没有可用的预编译二进制": "No pre-built binary available for platform %s",
|
||||
"或者直接打开以下链接:": "Or open the following link:",
|
||||
"所有凭证已失效,请运行 dws auth login 重新登录": "All credentials have expired, please run dws auth login to re-authenticate",
|
||||
@@ -65,9 +67,9 @@
|
||||
"授权超时(5分钟),请重试": "Authorization timeout (5 minutes), please retry",
|
||||
"数据表": "data table",
|
||||
"数据表 ID (必填)": "data table ID (required)",
|
||||
"数据表管理": "data tablemanagement",
|
||||
"文件不存在: ": "filenot found: ",
|
||||
"文件为空": "fileis empty",
|
||||
"数据表管理": "Data table management",
|
||||
"文件不存在: ": "File not found: ",
|
||||
"文件为空": "File is empty",
|
||||
"文件过大 (%d 字节,限制 %d 字节)": "File too large (%d bytes, limit %d bytes)",
|
||||
"无法打开文件: %w": "Cannot open file: %w",
|
||||
"无法自动打开浏览器": "Cannot automatically open browser",
|
||||
@@ -82,7 +84,7 @@
|
||||
"服务返回了空结果;请稍后重试,必要时查看 recovery snapshot。": "Service returned empty results; retry later. Check recovery snapshot if persistent.",
|
||||
"未找到 MCP Server URL": "MCP Server URL not found",
|
||||
"未找到认证信息,请运行 dws auth login": "No credentials found, please run dws auth login",
|
||||
"未登录,请运行 dws auth login": "not logged in,please run dws auth login",
|
||||
"未登录,请运行 dws auth login": "Not logged in, please run dws auth login",
|
||||
"未知错误: %s": "Unknown error: %s",
|
||||
"本地文件一键上传到 AITable 附件字段": "One-click upload local file to AITable attachment field",
|
||||
"本地文件路径 (必填)": "Local file path (required)",
|
||||
@@ -90,36 +92,51 @@
|
||||
"构建上传请求失败: %w": "Failed to build upload request: %w",
|
||||
"查询入职人才": "Query onboarding candidates",
|
||||
"查询在招岗位": "Query open positions",
|
||||
"修改待办任务": "Update todo task",
|
||||
"修改执行者的待办完成状态": "Update executor todo done status",
|
||||
"优先级: 10低/20普通/30较高/40紧急": "Priority: 10=low/20=normal/30=high/40=urgent",
|
||||
"查询待办列表": "List todos",
|
||||
"待办": "todo",
|
||||
"待办任务 ID (必填)": "Todo task ID (required)",
|
||||
"待办任务管理": "Todo task management",
|
||||
"待办标题 (必填)": "Todo title (required)",
|
||||
"待办详情": "Todo detail",
|
||||
"循环待办 (需先设置 --due); 格式: DTSTART:...\\nRRULE:FREQ=DAILY;INTERVAL=1": "Recurring todo (requires --due); format: DTSTART:...\\nRRULE:FREQ=DAILY;INTERVAL=1",
|
||||
"截止时间 ISO-8601 (如 2026-03-10T18:00:00+08:00)": "Due time ISO-8601 (e.g. 2026-03-10T18:00:00+08:00)",
|
||||
"执行者 userId 列表 (必填)": "Executor userId list (required)",
|
||||
"检查登录状态后重试": "Check login status and retry",
|
||||
"检查服务连通性后重试;如持续失败,请确认 MCP 服务响应正常。": "Check service connectivity and retry; if persistent, verify MCP service is responding.",
|
||||
"检查服务连通性和协议版本后重试": "Check service connectivity and protocol version, then retry",
|
||||
"检查认证、权限和参数后重试原命令": "Check authentication, permissions, and parameters, then retry",
|
||||
"新标题": "New title",
|
||||
"管理钉钉个人待办:创建、查询列表、查看详情、修改、标记完成、删除。": "Manage DingTalk personal todos: create, list, view detail, update, mark done, delete.",
|
||||
"步骤 1/3: 准备上传 %s (%d 字节, %s)...\\\n": "Step 1/3: Preparing upload %s (%d bytes, %s)...\\\n",
|
||||
"步骤 2/3: 上传文件到 OSS...": "Step 2/3: Uploading file to OSS...",
|
||||
"步骤 3/3: 上传完成!": "Step 3/3: Upload complete!",
|
||||
"用户拒绝了授权请求": "User rejected the authorization request",
|
||||
"确认删除? (yes/no): ": "confirmdelete? (yes/no): ",
|
||||
"完成状态: true/false": "Done status: true/false",
|
||||
"完成状态: true=已完成, false=未完成 (必填)": "Done status: true=completed, false=not completed (required)",
|
||||
"确认删除? (yes/no): ": "Confirm delete? (yes/no): ",
|
||||
"等待用户授权...": "Waiting for user authorization...",
|
||||
"结果格式与客户端预期不一致;请检查服务协议变更或回退到最近可用版本。": "Result format does not match client expectations; check service protocol changes or rollback to latest working version.",
|
||||
"网络错误,继续重试...": "Network error, retrying...",
|
||||
"获取使用指引": "Get usage guide",
|
||||
"获取审批实例详情": "getapproval instancedetails",
|
||||
"获取审批实例详情": "Get approval instance details",
|
||||
"获取数量,超过 20 自动分页 (默认 20)": "Fetch count, auto-paginate if over 20 (default 20)",
|
||||
"获取版本信息失败 (HTTP %d)": "getversioninfofailed (HTTP %d)",
|
||||
"获取版本信息失败 (HTTP %d)": "Get version info failed (HTTP %d)",
|
||||
"解析令牌数据失败": "Failed to parse token data",
|
||||
"解析响应失败": "Failed to parse response",
|
||||
"解析版本信息失败: %w": "parseversioninfofailed: %w",
|
||||
"解析版本信息失败: %w": "Parse version info failed: %w",
|
||||
"解析设备授权数据失败": "Failed to parse device authorization data",
|
||||
"认证信息已失效,请重新执行上一条命令(最多重试两次)": "Credentials expired, re-execute the last command (max 2 retries)",
|
||||
"认证失败;请检查登录状态或产品 URL 覆盖。": "Authentication failed; check login status or product URL override.",
|
||||
"记录": "record",
|
||||
"记录 ID 列表,逗号分隔 (必填)": "record ID list,comma-separated (required)",
|
||||
"记录管理": "recordmanagement",
|
||||
"记录 ID 列表,逗号分隔 (必填)": "Record ID list, comma-separated (required)",
|
||||
"记录管理": "Record management",
|
||||
"设备授权流程失败(已重试 %d 次)": "Device authorization flow failed (retried %d times)",
|
||||
"设备授权码已过期": "Device authorization code has expired",
|
||||
"设备授权码已过期(%d 秒),请重试": "Device authorization code expired (%d seconds), please retry",
|
||||
"设置权限失败: %w": "settingspermissionfailed: %w",
|
||||
"设置权限失败: %w": "Set permission failed: %w",
|
||||
"请在浏览器中完成扫码授权。": "Please complete QR code authorization in the browser.",
|
||||
"请在浏览器中打开以下链接,并输入授权码:": "Please open the following link in your browser and enter the authorization code:",
|
||||
"请检查服务 endpoint 是否为空或格式不合法。": "Please check if the service endpoint is empty or has invalid format.",
|
||||
@@ -130,11 +147,28 @@
|
||||
"请求设备授权码失败": "Failed to request device authorization code",
|
||||
"读取 zip 条目失败: %w": "Failed to read zip entry: %w",
|
||||
"读取响应失败": "Failed to read response",
|
||||
"读取版本信息失败: %w": "readversioninfofailed: %w",
|
||||
"读取版本信息失败: %w": "Read version info failed: %w",
|
||||
"跳过确认直接删除": "Skip confirmation and delete directly",
|
||||
"--title 的别名": "Alias for --title",
|
||||
"调用被拒绝;请检查认证状态、租户身份或访问权限。": "Call rejected; check authentication status, tenant identity, or access permissions.",
|
||||
"轮询过快,间隔增加至 %ds": "Polling too fast, interval increased to %ds",
|
||||
"返回数据缺少 uploadUrl 或 fileToken": "Response data missing uploadUrl or fileToken",
|
||||
"附件工作流": "attachmentworkflow",
|
||||
"附件工作流": "Attachment workflow",
|
||||
"页码 (必填)": "page number (required)",
|
||||
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ Unable to verify CLI data access permission status",
|
||||
" 请检查网络连接后重试。": " Please check your network connection and retry.",
|
||||
"检查 CLI 授权状态失败": "Failed to check CLI auth status",
|
||||
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ CLI data access is not enabled for this organization",
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " The organization admin has not enabled \"Allow members to access their personal data via CLI\".",
|
||||
" 组织主管理员:": " Organization super admins: ",
|
||||
" 请联系组织主管理员开启后重新登录。": " Please contact the organization super admin to enable it and re-login.",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " Admin settings: https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "CLI data access is not enabled for this organization, please contact admin to enable it",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ CLI data access is not enabled for this organization, please submit an authorization request in the browser...",
|
||||
"✅ 权限已开启,继续登录...": "✅ Permission enabled, continuing login...",
|
||||
"等待管理员审批中": "Waiting for admin approval",
|
||||
"等待提交申请中": "Waiting to submit request",
|
||||
"操作超时,请重新登录": "Operation timed out, please re-login",
|
||||
"检查组织 CLI 授权状态...": "Checking organization CLI auth status...",
|
||||
"🔐 登录钉钉": "🔐 Login to DingTalk"
|
||||
}
|
||||
|
||||
@@ -30,8 +30,10 @@
|
||||
"准备上传失败: %w": "准备上传失败: %w",
|
||||
"分页查询审批实例": "分页查询审批实例",
|
||||
"创建 / 查询 / 更新 / 删除待办": "创建 / 查询 / 更新 / 删除待办",
|
||||
"创建待办": "创建待办",
|
||||
"创建文件失败: %w": "创建文件失败: %w",
|
||||
"创建请求失败": "创建请求失败",
|
||||
"删除待办": "删除待办",
|
||||
"删除 AI 表格": "删除 AI 表格",
|
||||
"删除原因(可选)": "删除原因(可选)",
|
||||
"删除字段": "删除字段",
|
||||
@@ -90,15 +92,30 @@
|
||||
"构建上传请求失败: %w": "构建上传请求失败: %w",
|
||||
"查询入职人才": "查询入职人才",
|
||||
"查询在招岗位": "查询在招岗位",
|
||||
"修改待办任务": "修改待办任务",
|
||||
"修改执行者的待办完成状态": "修改执行者的待办完成状态",
|
||||
"优先级: 10低/20普通/30较高/40紧急": "优先级: 10低/20普通/30较高/40紧急",
|
||||
"查询待办列表": "查询待办列表",
|
||||
"待办": "待办",
|
||||
"待办任务 ID (必填)": "待办任务 ID (必填)",
|
||||
"待办任务管理": "待办任务管理",
|
||||
"待办标题 (必填)": "待办标题 (必填)",
|
||||
"待办详情": "待办详情",
|
||||
"循环待办 (需先设置 --due); 格式: DTSTART:...\\nRRULE:FREQ=DAILY;INTERVAL=1": "循环待办 (需先设置 --due); 格式: DTSTART:...\\nRRULE:FREQ=DAILY;INTERVAL=1",
|
||||
"截止时间 ISO-8601 (如 2026-03-10T18:00:00+08:00)": "截止时间 ISO-8601 (如 2026-03-10T18:00:00+08:00)",
|
||||
"执行者 userId 列表 (必填)": "执行者 userId 列表 (必填)",
|
||||
"检查登录状态后重试": "检查登录状态后重试",
|
||||
"检查服务连通性后重试;如持续失败,请确认 MCP 服务响应正常。": "检查服务连通性后重试;如持续失败,请确认 MCP 服务响应正常。",
|
||||
"检查服务连通性和协议版本后重试": "检查服务连通性和协议版本后重试",
|
||||
"检查认证、权限和参数后重试原命令": "检查认证、权限和参数后重试原命令",
|
||||
"新标题": "新标题",
|
||||
"管理钉钉个人待办:创建、查询列表、查看详情、修改、标记完成、删除。": "管理钉钉个人待办:创建、查询列表、查看详情、修改、标记完成、删除。",
|
||||
"步骤 1/3: 准备上传 %s (%d 字节, %s)...\\\n": "步骤 1/3: 准备上传 %s (%d 字节, %s)...\\\n",
|
||||
"步骤 2/3: 上传文件到 OSS...": "步骤 2/3: 上传文件到 OSS...",
|
||||
"步骤 3/3: 上传完成!": "步骤 3/3: 上传完成!",
|
||||
"用户拒绝了授权请求": "用户拒绝了授权请求",
|
||||
"完成状态: true/false": "完成状态: true/false",
|
||||
"完成状态: true=已完成, false=未完成 (必填)": "完成状态: true=已完成, false=未完成 (必填)",
|
||||
"确认删除? (yes/no): ": "确认删除? (yes/no): ",
|
||||
"等待用户授权...": "等待用户授权...",
|
||||
"结果格式与客户端预期不一致;请检查服务协议变更或回退到最近可用版本。": "结果格式与客户端预期不一致;请检查服务协议变更或回退到最近可用版本。",
|
||||
@@ -131,10 +148,27 @@
|
||||
"读取 zip 条目失败: %w": "读取 zip 条目失败: %w",
|
||||
"读取响应失败": "读取响应失败",
|
||||
"读取版本信息失败: %w": "读取版本信息失败: %w",
|
||||
"跳过确认直接删除": "跳过确认直接删除",
|
||||
"--title 的别名": "--title 的别名",
|
||||
"调用被拒绝;请检查认证状态、租户身份或访问权限。": "调用被拒绝;请检查认证状态、租户身份或访问权限。",
|
||||
"轮询过快,间隔增加至 %ds": "轮询过快,间隔增加至 %ds",
|
||||
"返回数据缺少 uploadUrl 或 fileToken": "返回数据缺少 uploadUrl 或 fileToken",
|
||||
"附件工作流": "附件工作流",
|
||||
"页码 (必填)": "页码 (必填)",
|
||||
"⚠️ 无法检查 CLI 数据访问权限状态": "⚠️ 无法检查 CLI 数据访问权限状态",
|
||||
" 请检查网络连接后重试。": " 请检查网络连接后重试。",
|
||||
"检查 CLI 授权状态失败": "检查 CLI 授权状态失败",
|
||||
"⚠️ 该组织尚未开启 CLI 数据访问权限": "⚠️ 该组织尚未开启 CLI 数据访问权限",
|
||||
" 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。": " 你所选择的组织管理员尚未开启「允许成员通过 CLI 访问其个人数据」的权限。",
|
||||
" 组织主管理员:": " 组织主管理员:",
|
||||
" 请联系组织主管理员开启后重新登录。": " 请联系组织主管理员开启后重新登录。",
|
||||
" 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings": " 管理员操作入口:https://open-dev.dingtalk.com/fe/old#/developerSettings",
|
||||
"该组织尚未开启 CLI 数据访问权限,请联系管理员开启": "该组织尚未开启 CLI 数据访问权限,请联系管理员开启",
|
||||
"⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...": "⏳ 该组织尚未开启 CLI 数据访问权限,请在浏览器中提交授权申请...",
|
||||
"✅ 权限已开启,继续登录...": "✅ 权限已开启,继续登录...",
|
||||
"等待管理员审批中": "等待管理员审批中",
|
||||
"等待提交申请中": "等待提交申请中",
|
||||
"操作超时,请重新登录": "操作超时,请重新登录",
|
||||
"检查组织 CLI 授权状态...": "检查组织 CLI 授权状态...",
|
||||
"🔐 登录钉钉": "🔐 登录钉钉"
|
||||
}
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
// Package keychain provides cross-platform secure storage for secrets.
|
||||
// - macOS: System Keychain stores DEK (Data Encryption Key), data encrypted with AES-256-GCM
|
||||
// - Linux: File-based DEK storage with AES-256-GCM encryption
|
||||
// - Windows: DPAPI + Registry storage
|
||||
package keychain
|
||||
|
||||
const (
|
||||
// Service is the unified keychain service name for all secrets.
|
||||
Service = "dws-cli"
|
||||
|
||||
// AccountToken is the account key for storing auth token data.
|
||||
AccountToken = "auth-token"
|
||||
)
|
||||
|
||||
// KeychainAccess abstracts keychain Get/Set/Remove for dependency injection.
|
||||
type KeychainAccess interface {
|
||||
Get(service, account string) (string, error)
|
||||
Set(service, account, value string) error
|
||||
Remove(service, account string) error
|
||||
}
|
||||
|
||||
// Get retrieves a value from the keychain.
|
||||
// Returns empty string and nil error if the entry does not exist.
|
||||
func Get(service, account string) (string, error) {
|
||||
return platformGet(service, account)
|
||||
}
|
||||
|
||||
// Set stores a value in the keychain, overwriting any existing entry.
|
||||
func Set(service, account, data string) error {
|
||||
return platformSet(service, account, data)
|
||||
}
|
||||
|
||||
// Remove deletes an entry from the keychain.
|
||||
// Returns nil if the entry does not exist.
|
||||
func Remove(service, account string) error {
|
||||
return platformRemove(service, account)
|
||||
}
|
||||
|
||||
// Exists checks if an entry exists in the keychain.
|
||||
func Exists(service, account string) bool {
|
||||
val, err := Get(service, account)
|
||||
return err == nil && val != ""
|
||||
}
|
||||
@@ -0,0 +1,200 @@
|
||||
// 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.
|
||||
|
||||
//go:build darwin
|
||||
|
||||
package keychain
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
"github.com/zalando/go-keyring"
|
||||
)
|
||||
|
||||
const (
|
||||
keychainTimeout = 5 * time.Second
|
||||
dekBytes = 32 // DEK = Data Encryption Key (AES-256)
|
||||
ivBytes = 12
|
||||
tagBytes = 16
|
||||
)
|
||||
|
||||
// StorageDir returns the storage directory for a given service name on macOS.
|
||||
// Uses ~/Library/Application Support/<service> following Apple conventions.
|
||||
func StorageDir(service string) string {
|
||||
home, err := os.UserHomeDir()
|
||||
if err != nil || home == "" {
|
||||
return filepath.Join(".dws", "keychain", service)
|
||||
}
|
||||
return filepath.Join(home, "Library", "Application Support", service)
|
||||
}
|
||||
|
||||
var safeFileNameRe = regexp.MustCompile(`[^a-zA-Z0-9._-]`)
|
||||
|
||||
func safeFileName(account string) string {
|
||||
return safeFileNameRe.ReplaceAllString(account, "_") + ".enc"
|
||||
}
|
||||
|
||||
// getDEK retrieves or generates the Data Encryption Key from system Keychain.
|
||||
func getDEK(service string) ([]byte, error) {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), keychainTimeout)
|
||||
defer cancel()
|
||||
|
||||
type result struct {
|
||||
key []byte
|
||||
err error
|
||||
}
|
||||
resCh := make(chan result, 1)
|
||||
|
||||
go func() {
|
||||
defer func() { recover() }()
|
||||
|
||||
// Try to get existing DEK from system Keychain
|
||||
encodedKey, err := keyring.Get(service, "dek")
|
||||
if err == nil {
|
||||
key, decodeErr := base64.StdEncoding.DecodeString(encodedKey)
|
||||
if decodeErr == nil && len(key) == dekBytes {
|
||||
resCh <- result{key: key, err: nil}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Generate new DEK if not found or invalid
|
||||
key := make([]byte, dekBytes)
|
||||
if _, randErr := rand.Read(key); randErr != nil {
|
||||
resCh <- result{key: nil, err: randErr}
|
||||
return
|
||||
}
|
||||
|
||||
// Store in system Keychain
|
||||
encodedKey = base64.StdEncoding.EncodeToString(key)
|
||||
setErr := keyring.Set(service, "dek", encodedKey)
|
||||
resCh <- result{key: key, err: setErr}
|
||||
}()
|
||||
|
||||
select {
|
||||
case res := <-resCh:
|
||||
return res.key, res.err
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func encryptData(plaintext string, key []byte) ([]byte, error) {
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
aesGCM, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
iv := make([]byte, ivBytes)
|
||||
if _, err := rand.Read(iv); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ciphertext := aesGCM.Seal(nil, iv, []byte(plaintext), nil)
|
||||
result := make([]byte, 0, ivBytes+len(ciphertext))
|
||||
result = append(result, iv...)
|
||||
result = append(result, ciphertext...)
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func decryptData(data []byte, key []byte) (string, error) {
|
||||
if len(data) < ivBytes+tagBytes {
|
||||
return "", fmt.Errorf("ciphertext too short")
|
||||
}
|
||||
block, err := aes.NewCipher(key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
aesGCM, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
iv := data[:ivBytes]
|
||||
ciphertext := data[ivBytes:]
|
||||
plaintext, err := aesGCM.Open(nil, iv, ciphertext, nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decryption failed: %w", err)
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
func platformGet(service, account string) (string, error) {
|
||||
key, err := getDEK(service)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
data, err := os.ReadFile(filepath.Join(StorageDir(service), safeFileName(account)))
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", nil // Not found is not an error
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
plaintext, err := decryptData(data, key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return plaintext, nil
|
||||
}
|
||||
|
||||
func platformSet(service, account, data string) error {
|
||||
key, err := getDEK(service)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
dir := StorageDir(service)
|
||||
if err := os.MkdirAll(dir, 0700); err != nil {
|
||||
return err
|
||||
}
|
||||
encrypted, err := encryptData(data, key)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
targetPath := filepath.Join(dir, safeFileName(account))
|
||||
tmpPath := filepath.Join(dir, safeFileName(account)+"."+uuid.New().String()+".tmp")
|
||||
defer os.Remove(tmpPath)
|
||||
|
||||
if err := os.WriteFile(tmpPath, encrypted, 0600); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Atomic rename to prevent file corruption during multi-process writes
|
||||
if err := os.Rename(tmpPath, targetPath); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func platformRemove(service, account string) error {
|
||||
err := os.Remove(filepath.Join(StorageDir(service), safeFileName(account)))
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user