Compare commits
150
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
129e8a10ef | ||
|
|
02fba09c1e | ||
|
|
94b64f74ac | ||
|
|
441289cdfe | ||
|
|
d03823d772 | ||
|
|
c16a377863 | ||
|
|
31c3acc94b | ||
|
|
f7e702df8d | ||
|
|
7fd40ea19a | ||
|
|
bee246e62c | ||
|
|
089caa92a8 | ||
|
|
2f0f32f56f | ||
|
|
068d9ff2f5 | ||
|
|
21cd3f8bc4 | ||
|
|
418928b9a5 | ||
|
|
b047b2c3c9 | ||
|
|
a516e5f54a | ||
|
|
70e4e75c66 | ||
|
|
1116916b24 | ||
|
|
c0c81b4d70 | ||
|
|
eedc41ac54 | ||
|
|
c14e24569c | ||
|
|
8add2c00cf | ||
|
|
29dceec5ce | ||
|
|
b78a0dee47 | ||
|
|
807191396e | ||
|
|
cce9b798d5 | ||
|
|
bfa3a1bf33 | ||
|
|
e154b4ecde | ||
|
|
f83c305749 | ||
|
|
b898f5c987 | ||
|
|
20750df20b | ||
|
|
749149b94a | ||
|
|
e5c8ff9acd | ||
|
|
706535b41e | ||
|
|
e15a2c4efb | ||
|
|
9e88116a2d | ||
|
|
05a306148a | ||
|
|
0dcc796f4c | ||
|
|
3e792b1c86 | ||
|
|
b9c822d49d | ||
|
|
f9b9b83f48 | ||
|
|
aa9e67e7c8 | ||
|
|
6c0cf3438b | ||
|
|
e3f30420fb | ||
|
|
16ff02903a | ||
|
|
65d3f2959c | ||
|
|
cb3087ba9b | ||
|
|
5068cfdab8 | ||
|
|
3c81e5d47d | ||
|
|
d93925a892 | ||
|
|
faab9e0282 | ||
|
|
76d301268d | ||
|
|
7fddace8df | ||
|
|
99e5a3cceb | ||
|
|
7e31043875 | ||
|
|
55d7fbf59a | ||
|
|
c0f4d21c05 | ||
|
|
660c908585 | ||
|
|
377ebc5e85 | ||
|
|
076d77da8e | ||
|
|
2eca203e74 | ||
|
|
6e070a7e24 | ||
|
|
e867abd03c | ||
|
|
9d89965de9 | ||
|
|
41088bb965 | ||
|
|
8259116f15 | ||
|
|
9ec1fa0638 | ||
|
|
9afd3be79b | ||
|
|
67da5019e3 | ||
|
|
b0ded7deb8 | ||
|
|
22905fc41e | ||
|
|
ec9ff653fc | ||
|
|
876cf8e958 | ||
|
|
1c5ed6646e | ||
|
|
80549a80e0 | ||
|
|
883d416d83 | ||
|
|
25c70aeb24 | ||
|
|
6cfa9e3afb | ||
|
|
544a91e994 | ||
|
|
024d487a22 | ||
|
|
6b50cc41c5 | ||
|
|
11e50662f9 | ||
|
|
29abdb6e79 | ||
|
|
bc587ddd91 | ||
|
|
6196e2565e | ||
|
|
609d56305e | ||
|
|
e69a1084a7 | ||
|
|
978ee6e636 | ||
|
|
987c63d99c | ||
|
|
e565746fb7 | ||
|
|
4f76d7cb4c | ||
|
|
e94f230236 | ||
|
|
5242a0e1b1 | ||
|
|
c3e57b874b | ||
|
|
fa558372d2 | ||
|
|
ec9ae33a43 | ||
|
|
2b49a2f365 | ||
|
|
ae9b14e536 | ||
|
|
996c4ab250 | ||
|
|
cf36ccb46e | ||
|
|
361115956f | ||
|
|
e82574cdde | ||
|
|
b2f917aa47 | ||
|
|
7cb0398bae | ||
|
|
91bd7c7802 | ||
|
|
5a8376ac0f | ||
|
|
31c984e18c | ||
|
|
69c0eb1a49 | ||
|
|
82e98d98a3 | ||
|
|
ef509ecdeb | ||
|
|
8a7e1c7be7 | ||
|
|
f59be6c19a | ||
|
|
fbc2575c93 | ||
|
|
58cb4789cd | ||
|
|
63b5fe3143 | ||
|
|
41a65f268f | ||
|
|
03796388c3 | ||
|
|
69b31da4ba | ||
|
|
833d0cc05e | ||
|
|
b7cbef1c6f | ||
|
|
bdc480cf49 | ||
|
|
43eaadcf07 | ||
|
|
f8e1be5970 | ||
|
|
8d1ccd1b98 | ||
|
|
c84b5d05f4 | ||
|
|
5224d9c527 | ||
|
|
e9360fe11b | ||
|
|
2d143589f8 | ||
|
|
93854178e3 | ||
|
|
c7c9a6f926 | ||
|
|
669518682c | ||
|
|
642e676f79 | ||
|
|
a0224e1cbd | ||
|
|
0d81f061d8 | ||
|
|
cd22cfb530 | ||
|
|
adc528c206 | ||
|
|
a2f1e79603 | ||
|
|
1adb4bc681 | ||
|
|
723c577484 | ||
|
|
766930f6e7 | ||
|
|
2e3311c955 | ||
|
|
1b4bb6b498 | ||
|
|
368e439280 | ||
|
|
8965fd2707 | ||
|
|
c0a7ad88a4 | ||
|
|
b62b1848aa | ||
|
|
5dd7f9abd3 | ||
|
|
eefee3c063 | ||
|
|
6f5d17335b |
@@ -5,6 +5,13 @@
|
||||
|
||||
## Verification
|
||||
|
||||
For an exact in-place `CHANGELOG.md`-only pull request, the full-suite checks
|
||||
may be marked `N/A`, but the targeted CHANGELOG check is required. For every
|
||||
other pull request, mark the targeted check `N/A` and complete the applicable
|
||||
full-suite checks.
|
||||
|
||||
- [ ] Exact `CHANGELOG.md`-only check (otherwise `N/A`):
|
||||
`./scripts/policy/check-changelog-pr.sh --fast-path "$(git merge-base HEAD origin/main)" HEAD`
|
||||
- [ ] `make build`
|
||||
- [ ] `make lint`
|
||||
- [ ] `make test`
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
name: AI Behavior Check
|
||||
name: Code Admission — AI Behavior
|
||||
|
||||
on:
|
||||
pull_request_target:
|
||||
types: [opened, synchronize, reopened, labeled, unlabeled]
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -11,7 +14,7 @@ permissions:
|
||||
|
||||
jobs:
|
||||
ai-behavior-check:
|
||||
name: AI Behavior Policy Evaluator
|
||||
name: AI Behavior
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
@@ -21,75 +24,108 @@ jobs:
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const sha = context.payload.pull_request.head.sha;
|
||||
const pullRequest = context.payload.pull_request;
|
||||
const sha = context.eventName === 'push' ? context.sha : pullRequest.head.sha;
|
||||
const setStatus = (state, description) =>
|
||||
github.rest.repos.createCommitStatus({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
sha,
|
||||
state,
|
||||
context: 'AI Behavior Check',
|
||||
context: 'AI Behavior',
|
||||
description,
|
||||
});
|
||||
|
||||
await setStatus('pending', 'Evaluating AI-generated PR boundaries');
|
||||
|
||||
const labels = context.payload.pull_request.labels.map(({ name }) => name);
|
||||
if (!labels.includes('ai-generated')) {
|
||||
await setStatus('success', 'Not labeled ai-generated');
|
||||
core.notice('Not an ai-generated PR; no AI-only policy applied.');
|
||||
if (context.eventName === 'push') {
|
||||
await setStatus('success', 'Not applicable to the protected main push');
|
||||
core.notice('AI Behavior is a PR policy; the main push context is sealed.');
|
||||
return;
|
||||
}
|
||||
|
||||
const files = await github.paginate(github.rest.pulls.listFiles, {
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.issue.number,
|
||||
per_page: 100,
|
||||
});
|
||||
try {
|
||||
const expectedHead = pullRequest.head.sha;
|
||||
const expectedBase = pullRequest.base.sha;
|
||||
const currentPull = async (phase) => {
|
||||
const { data: pull } = await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.issue.number,
|
||||
});
|
||||
if (pull.head.sha !== expectedHead || pull.base.sha !== expectedBase) {
|
||||
throw new Error(
|
||||
`Pull request revision changed during ${phase}: ` +
|
||||
`expected base/head ${expectedBase}/${expectedHead}, ` +
|
||||
`got ${pull.base.sha}/${pull.head.sha}`
|
||||
);
|
||||
}
|
||||
return pull;
|
||||
};
|
||||
|
||||
const before = await currentPull('pre-policy check');
|
||||
const labels = before.labels.map(({ name }) => name);
|
||||
if (!labels.includes('ai-generated')) {
|
||||
await setStatus('success', 'Not labeled ai-generated');
|
||||
core.notice('Not an ai-generated PR; no AI-only policy applied.');
|
||||
return;
|
||||
}
|
||||
|
||||
const files = await github.paginate(github.rest.pulls.listFiles, {
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.issue.number,
|
||||
per_page: 100,
|
||||
});
|
||||
await currentPull('post-policy check');
|
||||
|
||||
const maxChangedFiles = 30;
|
||||
if (files.length > maxChangedFiles) {
|
||||
await setStatus(
|
||||
'failure',
|
||||
`Changes ${files.length} files; limit is ${maxChangedFiles}`
|
||||
);
|
||||
core.setFailed(
|
||||
`AI-generated PR changes ${files.length} files; limit is ${maxChangedFiles}.`
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
const isProtectedPath = (filename) =>
|
||||
typeof filename === 'string' &&
|
||||
(
|
||||
filename.startsWith('.github/workflows/') ||
|
||||
filename.startsWith('scripts/ci/') ||
|
||||
filename.startsWith('scripts/policy/') ||
|
||||
filename.startsWith('scripts/release/') ||
|
||||
filename === 'test/fixtures/cli-interface-baseline.txt' ||
|
||||
filename === '.goreleaser.yaml' ||
|
||||
filename === 'Makefile'
|
||||
);
|
||||
const protectedPaths = [...new Set(
|
||||
files
|
||||
.flatMap(({ filename, previous_filename }) => [filename, previous_filename])
|
||||
.filter(isProtectedPath)
|
||||
)];
|
||||
|
||||
if (protectedPaths.length > 0) {
|
||||
await setStatus('failure', 'Modifies protected release/CI infrastructure');
|
||||
core.setFailed(
|
||||
'AI-generated PR modifies protected release/CI infrastructure:\n' +
|
||||
protectedPaths.map((filename) => ` - ${filename}`).join('\n') +
|
||||
'\nSplit these changes into a human-owned PR with explicit review.'
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
const maxChangedFiles = 30;
|
||||
if (files.length > maxChangedFiles) {
|
||||
await setStatus(
|
||||
'failure',
|
||||
`Changes ${files.length} files; limit is ${maxChangedFiles}`
|
||||
'success',
|
||||
`Passed with ${files.length} changed files (limit ${maxChangedFiles})`
|
||||
);
|
||||
core.setFailed(
|
||||
`AI-generated PR changes ${files.length} files; limit is ${maxChangedFiles}.`
|
||||
core.notice(
|
||||
`AI behavior check passed (${files.length} changed files; limit ${maxChangedFiles}).`
|
||||
);
|
||||
return;
|
||||
} catch (error) {
|
||||
await setStatus('error', 'Could not evaluate the exact pull request revision');
|
||||
throw error;
|
||||
}
|
||||
|
||||
const isProtectedPath = (filename) =>
|
||||
typeof filename === 'string' &&
|
||||
(
|
||||
filename.startsWith('.github/workflows/') ||
|
||||
filename.startsWith('scripts/policy/') ||
|
||||
filename.startsWith('scripts/release/') ||
|
||||
filename === 'test/fixtures/cli-interface-baseline.txt' ||
|
||||
filename === '.goreleaser.yaml' ||
|
||||
filename === 'Makefile'
|
||||
);
|
||||
const protectedPaths = [...new Set(
|
||||
files
|
||||
.flatMap(({ filename, previous_filename }) => [filename, previous_filename])
|
||||
.filter(isProtectedPath)
|
||||
)];
|
||||
|
||||
if (protectedPaths.length > 0) {
|
||||
await setStatus('failure', 'Modifies protected release/CI infrastructure');
|
||||
core.setFailed(
|
||||
'AI-generated PR modifies protected release/CI infrastructure:\n' +
|
||||
protectedPaths.map((filename) => ` - ${filename}`).join('\n') +
|
||||
'\nSplit these changes into a human-owned PR with explicit review.'
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
await setStatus(
|
||||
'success',
|
||||
`Passed with ${files.length} changed files (limit ${maxChangedFiles})`
|
||||
);
|
||||
core.notice(
|
||||
`AI behavior check passed (${files.length} changed files; limit ${maxChangedFiles}).`
|
||||
);
|
||||
|
||||
+372
-142
@@ -18,48 +18,138 @@ jobs:
|
||||
name: Lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
outputs:
|
||||
changelog_only: ${{ steps.classify.outputs.changelog_only }}
|
||||
changelog_changed: ${{ steps.classify.outputs.changelog_changed }}
|
||||
platform_sensitive: ${{ steps.classify.outputs.platform_sensitive }}
|
||||
steps:
|
||||
- name: Classify pull request scope
|
||||
id: classify
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
let changelogOnly = false;
|
||||
let changelogChanged = false;
|
||||
let platformSensitive = context.eventName === 'push';
|
||||
let files = [];
|
||||
|
||||
if (context.eventName === 'pull_request') {
|
||||
const expectedHead = context.payload.pull_request.head.sha;
|
||||
const expectedBase = context.payload.pull_request.base.sha;
|
||||
const assertCurrentRevision = async (phase) => {
|
||||
const { data: pull } = await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.issue.number,
|
||||
});
|
||||
if (pull.head.sha !== expectedHead || pull.base.sha !== expectedBase) {
|
||||
throw new Error(
|
||||
`Pull request revision changed during ${phase}: ` +
|
||||
`expected base/head ${expectedBase}/${expectedHead}, ` +
|
||||
`got ${pull.base.sha}/${pull.head.sha}`
|
||||
);
|
||||
}
|
||||
return pull;
|
||||
};
|
||||
|
||||
const before = await assertCurrentRevision('pre-classification');
|
||||
files = await github.paginate(github.rest.pulls.listFiles, {
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.issue.number,
|
||||
per_page: 100,
|
||||
});
|
||||
const after = await assertCurrentRevision('post-classification');
|
||||
if (
|
||||
before.changed_files !== files.length ||
|
||||
after.changed_files !== files.length
|
||||
) {
|
||||
throw new Error(
|
||||
`Pull request file list is incomplete: API reports ` +
|
||||
`${after.changed_files} changed files, pagination returned ${files.length}`
|
||||
);
|
||||
}
|
||||
|
||||
changelogOnly =
|
||||
files.length === 1 &&
|
||||
files[0].filename === 'CHANGELOG.md' &&
|
||||
files[0].status === 'modified' &&
|
||||
!files[0].previous_filename;
|
||||
changelogChanged = files.some(
|
||||
({ filename, previous_filename }) =>
|
||||
filename === 'CHANGELOG.md' ||
|
||||
previous_filename === 'CHANGELOG.md'
|
||||
);
|
||||
|
||||
const isPlatformSensitive = (filename) =>
|
||||
typeof filename === 'string' &&
|
||||
(
|
||||
filename.startsWith('internal/auth/') ||
|
||||
filename.startsWith('internal/keychain/') ||
|
||||
/_(darwin|windows|linux|unix)\.go$/.test(filename) ||
|
||||
filename.startsWith('scripts/release/') ||
|
||||
filename.startsWith('scripts/install') ||
|
||||
filename.startsWith('Formula/') ||
|
||||
filename.startsWith('build/npm/') ||
|
||||
filename === '.goreleaser.yaml' ||
|
||||
filename === '.github/workflows/release.yml'
|
||||
);
|
||||
platformSensitive = files.some(
|
||||
({ filename, previous_filename }) =>
|
||||
isPlatformSensitive(filename) ||
|
||||
isPlatformSensitive(previous_filename)
|
||||
);
|
||||
}
|
||||
|
||||
core.setOutput('changelog_only', String(changelogOnly));
|
||||
core.setOutput('changelog_changed', String(changelogChanged));
|
||||
core.setOutput('platform_sensitive', String(platformSensitive));
|
||||
await core.summary
|
||||
.addHeading('Code Admission scope')
|
||||
.addRaw(`- Event: \`${context.eventName}\`\n`)
|
||||
.addRaw(`- Exact modified CHANGELOG only: \`${changelogOnly}\`\n`)
|
||||
.addRaw(`- CHANGELOG touched: \`${changelogChanged}\`\n`)
|
||||
.addRaw(`- Native-platform risk paths touched: \`${platformSensitive}\`\n`)
|
||||
.addRaw(`- Changed files: \`${files.length}\`\n`)
|
||||
.write();
|
||||
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: steps.classify.outputs.changelog_only == 'true'
|
||||
run: echo "Lint is satisfied by the trusted CHANGELOG-only Policy path." >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
- name: Check out repository
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Go
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Verify Test Package Plan
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
run: make test-plan
|
||||
|
||||
- name: Format Check
|
||||
run: |
|
||||
unformatted="$(find cmd internal test scripts/policy -name '*.go' -print0 | xargs -0r gofmt -l)"
|
||||
test -z "$unformatted" || (printf '%s\n' "$unformatted" && exit 1)
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
run: make format-check
|
||||
|
||||
- name: Go Vet
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
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: ./...
|
||||
|
||||
actionlint:
|
||||
name: Actionlint
|
||||
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: Check GitHub Actions workflows
|
||||
if: steps.classify.outputs.changelog_only != 'true'
|
||||
run: go run github.com/rhysd/actionlint/cmd/actionlint@v1.7.12
|
||||
|
||||
test-race:
|
||||
name: "Test (race: ${{ matrix.shard }})"
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
strategy:
|
||||
@@ -79,45 +169,29 @@ jobs:
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Install archive tooling
|
||||
run: sudo apt-get update && sudo apt-get install -y zip unzip
|
||||
|
||||
- name: Build
|
||||
if: ${{ matrix.shard == 'remaining' }}
|
||||
run: make build
|
||||
|
||||
- name: Test shard with Race Detection
|
||||
shell: bash
|
||||
env:
|
||||
DWS_PACKAGE_VERSION: 0.0.0-test
|
||||
TEST_SHARD: ${{ matrix.shard }}
|
||||
run: |
|
||||
case "$TEST_SHARD" in
|
||||
app)
|
||||
packages=(./internal/app/...)
|
||||
;;
|
||||
generators)
|
||||
packages=(./internal/generator/...)
|
||||
;;
|
||||
helpers)
|
||||
packages=(./internal/helpers/...)
|
||||
;;
|
||||
remaining)
|
||||
mapfile -t packages < <(
|
||||
go list ./cmd/... ./internal/... |
|
||||
grep -Ev '/internal/(app|generator|helpers)(/|$)'
|
||||
)
|
||||
;;
|
||||
*)
|
||||
printf 'unknown test shard: %s\n' "$TEST_SHARD" >&2
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
set -euo pipefail
|
||||
package_output="$(./scripts/ci/test-packages.sh list "$TEST_SHARD")"
|
||||
test -n "$package_output"
|
||||
mapfile -t packages <<< "$package_output"
|
||||
test "${#packages[@]}" -gt 0
|
||||
go test -v -race -count=1 -timeout=10m "${packages[@]}"
|
||||
|
||||
test-release-scripts:
|
||||
name: Test (release scripts)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
@@ -131,27 +205,96 @@ jobs:
|
||||
run: sudo apt-get update && sudo apt-get install -y zip unzip
|
||||
|
||||
- name: Test release scripts
|
||||
run: go test -v -count=1 -timeout=5m ./test/scripts
|
||||
shell: bash
|
||||
env:
|
||||
DWS_PACKAGE_VERSION: 0.0.0-test
|
||||
run: |
|
||||
set -euo pipefail
|
||||
package_output="$(./scripts/ci/test-packages.sh list release-scripts)"
|
||||
test -n "$package_output"
|
||||
mapfile -t packages <<< "$package_output"
|
||||
test "${#packages[@]}" -gt 0
|
||||
go test -v -count=1 -timeout=10m "${packages[@]}"
|
||||
|
||||
test-cross-platform:
|
||||
name: Test (cross-platform compile)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
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: Compile supported operating systems
|
||||
shell: bash
|
||||
run: |
|
||||
set -eu
|
||||
for target in darwin/amd64 darwin/arm64 windows/amd64 windows/arm64; do
|
||||
goos="${target%/*}"
|
||||
goarch="${target#*/}"
|
||||
output="$RUNNER_TEMP/dws-${goos}-${goarch}"
|
||||
if [ "$goos" = windows ]; then
|
||||
output="${output}.exe"
|
||||
fi
|
||||
printf 'compile %s/%s\n' "$goos" "$goarch"
|
||||
CGO_ENABLED=0 GOOS="$goos" GOARCH="$goarch" \
|
||||
go build -o "$output" ./cmd
|
||||
done
|
||||
|
||||
test:
|
||||
name: Test
|
||||
needs:
|
||||
- lint
|
||||
- test-race
|
||||
- test-release-scripts
|
||||
if: ${{ always() }}
|
||||
- test-cross-platform
|
||||
- test-darwin
|
||||
- test-windows
|
||||
if: ${{ always() && needs.lint.result == 'success' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions: {}
|
||||
steps:
|
||||
- name: Verify test shards
|
||||
env:
|
||||
CHANGELOG_ONLY: ${{ needs.lint.outputs.changelog_only }}
|
||||
PLATFORM_SENSITIVE: ${{ needs.lint.outputs.platform_sensitive }}
|
||||
RACE_RESULT: ${{ needs.test-race.result }}
|
||||
RELEASE_SCRIPTS_RESULT: ${{ needs.test-release-scripts.result }}
|
||||
CROSS_PLATFORM_RESULT: ${{ needs.test-cross-platform.result }}
|
||||
DARWIN_RESULT: ${{ needs.test-darwin.result }}
|
||||
WINDOWS_RESULT: ${{ needs.test-windows.result }}
|
||||
run: |
|
||||
failed=0
|
||||
if [ "$CHANGELOG_ONLY" = true ]; then
|
||||
for shard in \
|
||||
"race shards:$RACE_RESULT" \
|
||||
"release scripts:$RELEASE_SCRIPTS_RESULT" \
|
||||
"cross-platform compile:$CROSS_PLATFORM_RESULT" \
|
||||
"macOS native:$DARWIN_RESULT" \
|
||||
"Windows native:$WINDOWS_RESULT"
|
||||
do
|
||||
name="${shard%%:*}"
|
||||
result="${shard#*:}"
|
||||
printf '%s: %s\n' "$name" "$result"
|
||||
if [ "$result" != skipped ]; then
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
test "$failed" -eq 0
|
||||
exit
|
||||
fi
|
||||
|
||||
for shard in \
|
||||
"race shards:$RACE_RESULT" \
|
||||
"release scripts:$RELEASE_SCRIPTS_RESULT"
|
||||
"release scripts:$RELEASE_SCRIPTS_RESULT" \
|
||||
"cross-platform compile:$CROSS_PLATFORM_RESULT"
|
||||
do
|
||||
name="${shard%%:*}"
|
||||
result="${shard#*:}"
|
||||
@@ -160,10 +303,28 @@ jobs:
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
|
||||
native_expected=skipped
|
||||
if [ "$PLATFORM_SENSITIVE" = true ]; then
|
||||
native_expected=success
|
||||
fi
|
||||
for native in \
|
||||
"macOS native:$DARWIN_RESULT" \
|
||||
"Windows native:$WINDOWS_RESULT"
|
||||
do
|
||||
name="${native%%:*}"
|
||||
result="${native#*:}"
|
||||
printf '%s: %s\n' "$name" "$result"
|
||||
if [ "$result" != "$native_expected" ]; then
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
test "$failed" -eq 0
|
||||
|
||||
test-darwin:
|
||||
name: Test (macOS auth/keychain)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' && needs.lint.outputs.platform_sensitive == 'true' }}
|
||||
runs-on: macos-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
@@ -180,6 +341,8 @@ jobs:
|
||||
|
||||
test-windows:
|
||||
name: Test (Windows)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' && needs.lint.outputs.platform_sensitive == 'true' }}
|
||||
runs-on: windows-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
@@ -202,6 +365,8 @@ jobs:
|
||||
|
||||
coverage-darwin:
|
||||
name: Coverage (macOS)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' && needs.lint.outputs.platform_sensitive == 'true' }}
|
||||
runs-on: macos-latest
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
@@ -244,6 +409,8 @@ jobs:
|
||||
|
||||
coverage-windows:
|
||||
name: Coverage (Windows)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' && needs.lint.outputs.platform_sensitive == 'true' }}
|
||||
runs-on: windows-latest
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
@@ -288,6 +455,8 @@ jobs:
|
||||
|
||||
coverage-current:
|
||||
name: Coverage (current)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
@@ -321,6 +490,8 @@ jobs:
|
||||
|
||||
coverage-supporting:
|
||||
name: Coverage (supporting)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
@@ -358,6 +529,8 @@ jobs:
|
||||
|
||||
coverage-baseline:
|
||||
name: Coverage (baseline)
|
||||
needs: lint
|
||||
if: ${{ needs.lint.outputs.changelog_only != 'true' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 20
|
||||
steps:
|
||||
@@ -421,20 +594,35 @@ jobs:
|
||||
coverage:
|
||||
name: Coverage
|
||||
needs:
|
||||
- lint
|
||||
- coverage-current
|
||||
- coverage-supporting
|
||||
- coverage-baseline
|
||||
if: ${{ always() }}
|
||||
- coverage-darwin
|
||||
- coverage-windows
|
||||
if: ${{ always() && needs.lint.result == 'success' }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Verify coverage profile jobs
|
||||
env:
|
||||
CHANGELOG_ONLY: ${{ needs.lint.outputs.changelog_only }}
|
||||
PLATFORM_SENSITIVE: ${{ needs.lint.outputs.platform_sensitive }}
|
||||
CURRENT_RESULT: ${{ needs.coverage-current.result }}
|
||||
SUPPORTING_RESULT: ${{ needs.coverage-supporting.result }}
|
||||
BASELINE_RESULT: ${{ needs.coverage-baseline.result }}
|
||||
DARWIN_RESULT: ${{ needs.coverage-darwin.result }}
|
||||
WINDOWS_RESULT: ${{ needs.coverage-windows.result }}
|
||||
run: |
|
||||
failed=0
|
||||
expected=success
|
||||
native_expected=skipped
|
||||
if [ "$CHANGELOG_ONLY" = true ]; then
|
||||
expected=skipped
|
||||
elif [ "$PLATFORM_SENSITIVE" = true ]; then
|
||||
native_expected=success
|
||||
fi
|
||||
|
||||
for profile in \
|
||||
"current:$CURRENT_RESULT" \
|
||||
"supporting:$SUPPORTING_RESULT" \
|
||||
@@ -443,24 +631,42 @@ jobs:
|
||||
name="${profile%%:*}"
|
||||
result="${profile#*:}"
|
||||
printf '%s: %s\n' "$name" "$result"
|
||||
if [ "$result" != "success" ]; then
|
||||
if [ "$result" != "$expected" ]; then
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
|
||||
for native in \
|
||||
"macOS native:$DARWIN_RESULT" \
|
||||
"Windows native:$WINDOWS_RESULT"
|
||||
do
|
||||
name="${native%%:*}"
|
||||
result="${native#*:}"
|
||||
printf '%s: %s\n' "$name" "$result"
|
||||
if [ "$CHANGELOG_ONLY" = true ]; then
|
||||
native_expected=skipped
|
||||
fi
|
||||
if [ "$result" != "$native_expected" ]; then
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
test "$failed" -eq 0
|
||||
|
||||
- name: Check out repository
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
ref: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Resolve authoritative coverage base
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
env:
|
||||
PR_HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
@@ -478,33 +684,40 @@ jobs:
|
||||
echo "COVERAGE_BASE_REF=$base_ref" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Download current coverage profile
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: coverage-current-profile
|
||||
path: .
|
||||
|
||||
- name: Download supporting coverage profiles
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: coverage-supporting-profiles
|
||||
path: .
|
||||
|
||||
- name: Download baseline coverage profile
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/download-artifact@v4
|
||||
with:
|
||||
name: coverage-baseline-profile
|
||||
path: .
|
||||
|
||||
- name: Enforce coverage gate
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
env:
|
||||
COVERAGE_TARGET: "80"
|
||||
COVERAGE_TARGET: "100"
|
||||
COVERAGE_ENFORCE_OVERALL: "false"
|
||||
run: COVERAGE_ADDITIONAL_PROFILE=coverage-shortcut.txt make coverage-gate BASE_REF="$COVERAGE_BASE_REF"
|
||||
COVERAGE_OVERALL_TOLERANCE: "0"
|
||||
run: COVERAGE_ADDITIONAL_DIFF_PROFILE=coverage-shortcut.txt make coverage-gate BASE_REF="$COVERAGE_BASE_REF"
|
||||
|
||||
- name: Generate coverage report
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: go tool cover -html=coverage.txt -o coverage.html
|
||||
|
||||
- name: Upload coverage artifact
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: coverage-report
|
||||
@@ -516,44 +729,114 @@ jobs:
|
||||
coverage.html
|
||||
|
||||
policy:
|
||||
name: Policy Check
|
||||
name: Policy
|
||||
needs: lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Verify pull request merge revision
|
||||
if: github.event_name == 'pull_request'
|
||||
env:
|
||||
PR_BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
PR_HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
run: |
|
||||
set -eu
|
||||
test "$(git rev-parse HEAD^1)" = "$PR_BASE_SHA" || {
|
||||
echo "checked-out merge first parent does not match event base" >&2
|
||||
exit 1
|
||||
}
|
||||
test "$(git rev-parse HEAD^2)" = "$PR_HEAD_SHA" || {
|
||||
echo "checked-out merge second parent does not match event head" >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
- name: Validate changed CHANGELOG content
|
||||
if: github.event_name == 'pull_request'
|
||||
env:
|
||||
CLASSIFIED_CHANGELOG_CHANGED: ${{ needs.lint.outputs.changelog_changed }}
|
||||
CHANGELOG_ONLY: ${{ needs.lint.outputs.changelog_only }}
|
||||
PR_BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
run: |
|
||||
set -eu
|
||||
merge_changelog_changed=false
|
||||
if git diff --no-ext-diff --find-renames --name-status \
|
||||
"$PR_BASE_SHA" HEAD |
|
||||
awk -F '\t' '
|
||||
{
|
||||
for (field = 2; field <= NF; field++) {
|
||||
if ($field == "CHANGELOG.md") found = 1
|
||||
}
|
||||
}
|
||||
END { exit !found }
|
||||
'
|
||||
then
|
||||
merge_changelog_changed=true
|
||||
fi
|
||||
test "$merge_changelog_changed" = "$CLASSIFIED_CHANGELOG_CHANGED" || {
|
||||
echo "Files API and synthetic merge tree disagree on CHANGELOG scope" >&2
|
||||
exit 1
|
||||
}
|
||||
if [ "$merge_changelog_changed" != true ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
mode=--content-only
|
||||
if [ "$CHANGELOG_ONLY" = true ]; then
|
||||
mode=--fast-path
|
||||
fi
|
||||
./scripts/policy/check-changelog-pr.sh \
|
||||
"$mode" "$PR_BASE_SHA" HEAD
|
||||
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: needs.lint.outputs.changelog_only == 'true'
|
||||
run: |
|
||||
echo "Only the base-equivalent CHANGELOG validator ran; full Policy resumes on main." \
|
||||
>> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
- name: Build
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make build
|
||||
|
||||
- name: Policy
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make policy
|
||||
|
||||
interface-integrity:
|
||||
name: Interface Integrity
|
||||
needs: lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
steps:
|
||||
- name: Check out repository
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 0
|
||||
ref: ${{ github.event_name == 'pull_request' && github.event.pull_request.head.sha || github.sha }}
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Build
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make build
|
||||
|
||||
- name: Resolve authoritative compatibility merge-base
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
env:
|
||||
PR_HEAD_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BASE_SHA: ${{ github.event.pull_request.base.sha }}
|
||||
@@ -578,6 +861,7 @@ jobs:
|
||||
echo "COMPATIBILITY_STABLE_REF=$stable_ref" >> "$GITHUB_ENV"
|
||||
|
||||
- name: Check historical commands and help compatibility
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: |
|
||||
make authoritative-interface-integrity \
|
||||
BASE_REF="$COMPATIBILITY_BASE_REF"
|
||||
@@ -587,143 +871,89 @@ jobs:
|
||||
fi
|
||||
|
||||
- name: Check complete Schema compatibility
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make schema-compatibility BASE_REF="$COMPATIBILITY_BASE_REF"
|
||||
|
||||
- name: Check skill command references
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make skill-command-integrity
|
||||
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: needs.lint.outputs.changelog_only == 'true'
|
||||
run: echo "Interface Integrity is unaffected by an exact CHANGELOG-only diff." >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
cli-smoke:
|
||||
name: CLI Smoke
|
||||
needs: lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Check out repository
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Build
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make build
|
||||
|
||||
- name: Check public top-level commands
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make cli-smoke
|
||||
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: needs.lint.outputs.changelog_only == 'true'
|
||||
run: echo "CLI Smoke is unaffected by an exact CHANGELOG-only diff." >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
mock-mcp-smoke:
|
||||
name: Mock MCP Smoke
|
||||
name: Mock MCP
|
||||
needs: lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Check out repository
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Check HTTP and stdio MCP transport
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: make mock-mcp-smoke
|
||||
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: needs.lint.outputs.changelog_only == 'true'
|
||||
run: echo "Mock MCP is unaffected by an exact CHANGELOG-only diff." >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
edition-tests:
|
||||
name: Edition Contract Tests
|
||||
name: Edition
|
||||
needs: lint
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 10
|
||||
steps:
|
||||
- name: Check out repository
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Go
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
uses: actions/setup-go@v5
|
||||
with:
|
||||
go-version-file: go.mod
|
||||
|
||||
- name: Run edition contract tests
|
||||
if: needs.lint.outputs.changelog_only != 'true'
|
||||
run: go test -v -count=1 ./pkg/editiontest/...
|
||||
|
||||
ci-gate:
|
||||
name: CI Gate
|
||||
needs:
|
||||
- lint
|
||||
- actionlint
|
||||
- test
|
||||
- test-darwin
|
||||
- test-windows
|
||||
- coverage
|
||||
- coverage-darwin
|
||||
- coverage-windows
|
||||
- policy
|
||||
- interface-integrity
|
||||
- cli-smoke
|
||||
- mock-mcp-smoke
|
||||
- edition-tests
|
||||
if: ${{ always() }}
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions: {}
|
||||
steps:
|
||||
- name: Verify required checks
|
||||
env:
|
||||
LINT_RESULT: ${{ needs.lint.result }}
|
||||
ACTIONLINT_RESULT: ${{ needs.actionlint.result }}
|
||||
TEST_RESULT: ${{ needs.test.result }}
|
||||
TEST_DARWIN_RESULT: ${{ needs.test-darwin.result }}
|
||||
TEST_WINDOWS_RESULT: ${{ needs.test-windows.result }}
|
||||
COVERAGE_RESULT: ${{ needs.coverage.result }}
|
||||
COVERAGE_DARWIN_RESULT: ${{ needs.coverage-darwin.result }}
|
||||
COVERAGE_WINDOWS_RESULT: ${{ needs.coverage-windows.result }}
|
||||
POLICY_RESULT: ${{ needs.policy.result }}
|
||||
INTERFACE_INTEGRITY_RESULT: ${{ needs.interface-integrity.result }}
|
||||
CLI_SMOKE_RESULT: ${{ needs.cli-smoke.result }}
|
||||
MOCK_MCP_SMOKE_RESULT: ${{ needs.mock-mcp-smoke.result }}
|
||||
EDITION_TESTS_RESULT: ${{ needs.edition-tests.result }}
|
||||
run: |
|
||||
failed=0
|
||||
for check in \
|
||||
"Lint:$LINT_RESULT" \
|
||||
"Actionlint:$ACTIONLINT_RESULT" \
|
||||
"Test:$TEST_RESULT" \
|
||||
"Test (macOS auth/keychain):$TEST_DARWIN_RESULT" \
|
||||
"Test (Windows):$TEST_WINDOWS_RESULT" \
|
||||
"Coverage:$COVERAGE_RESULT" \
|
||||
"Coverage (macOS):$COVERAGE_DARWIN_RESULT" \
|
||||
"Coverage (Windows):$COVERAGE_WINDOWS_RESULT" \
|
||||
"Policy Check:$POLICY_RESULT" \
|
||||
"Interface Integrity:$INTERFACE_INTEGRITY_RESULT" \
|
||||
"CLI Smoke:$CLI_SMOKE_RESULT" \
|
||||
"Mock MCP Smoke:$MOCK_MCP_SMOKE_RESULT" \
|
||||
"Edition Contract Tests:$EDITION_TESTS_RESULT"
|
||||
do
|
||||
name="${check%%:*}"
|
||||
result="${check#*:}"
|
||||
printf '%s: %s\n' "$name" "$result"
|
||||
if [ "$result" != "success" ]; then
|
||||
failed=1
|
||||
fi
|
||||
done
|
||||
test "$failed" -eq 0
|
||||
|
||||
notify-downstream:
|
||||
name: Notify Wukong Overlay
|
||||
needs: [ci-gate]
|
||||
runs-on: ubuntu-latest
|
||||
if: github.ref == 'refs/heads/main' && github.event_name == 'push'
|
||||
permissions: {}
|
||||
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
|
||||
- name: Record CHANGELOG-only fast path
|
||||
if: needs.lint.outputs.changelog_only == 'true'
|
||||
run: echo "Edition is unaffected by an exact CHANGELOG-only diff." >> "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
name: Multi Profile E2E
|
||||
name: Main Integration — 主干集成
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
@@ -14,7 +15,7 @@ concurrency:
|
||||
|
||||
jobs:
|
||||
multi-profile-e2e:
|
||||
name: Multi Profile E2E
|
||||
name: Multi-profile E2E
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
env:
|
||||
@@ -36,7 +37,7 @@ jobs:
|
||||
mkdir -p .tmp-bin
|
||||
bash scripts/dev/test-multi-profile-e2e.sh --keep-workdir | tee "$MULTI_PROFILE_E2E_LOG"
|
||||
{
|
||||
echo "### Multi Profile E2E"
|
||||
echo "### Multi-profile E2E"
|
||||
echo "- Command: \`bash scripts/dev/test-multi-profile-e2e.sh --keep-workdir\`"
|
||||
echo "- Scope: isolated auth/profile storage, profile switch/use, one-shot profile override, CSV multi-profile aggregation, legacy migration"
|
||||
echo "- Result: passed"
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
name: Main Integration — Wukong Overlay
|
||||
|
||||
on:
|
||||
workflow_run:
|
||||
workflows:
|
||||
- CI
|
||||
types:
|
||||
- completed
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
notify-downstream:
|
||||
name: Notify Wukong Overlay
|
||||
if: >-
|
||||
github.event.workflow_run.conclusion == 'success' &&
|
||||
github.event.workflow_run.event == 'push' &&
|
||||
github.event.workflow_run.head_branch == 'main'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
steps:
|
||||
- name: Trigger downstream CI
|
||||
env:
|
||||
UPSTREAM_SHA: ${{ github.event.workflow_run.head_sha }}
|
||||
WUKONG_TRIGGER_TOKEN: ${{ secrets.WUKONG_TRIGGER_TOKEN }}
|
||||
WUKONG_TRIGGER_URL: ${{ secrets.WUKONG_TRIGGER_URL }}
|
||||
run: |
|
||||
if [ -n "$WUKONG_TRIGGER_TOKEN" ]; then
|
||||
curl --fail --silent --show-error \
|
||||
-X POST \
|
||||
-F "token=$WUKONG_TRIGGER_TOKEN" \
|
||||
-F "ref=main" \
|
||||
-F "variables[UPSTREAM_SHA]=$UPSTREAM_SHA" \
|
||||
"$WUKONG_TRIGGER_URL"
|
||||
echo "Downstream CI triggered."
|
||||
else
|
||||
echo "No WUKONG_TRIGGER_TOKEN configured, skipping downstream notification."
|
||||
fi
|
||||
+1955
-209
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,148 @@
|
||||
name: Withdraw release
|
||||
run-name: Withdraw ${{ inputs.version }}
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
version:
|
||||
description: "Exact published version to withdraw (vX.Y.Z or vX.Y.Z-beta.N)"
|
||||
required: true
|
||||
type: string
|
||||
reason:
|
||||
description: "Public, single-line withdrawal reason (8-300 characters)"
|
||||
required: true
|
||||
type: string
|
||||
confirmation:
|
||||
description: "Type WITHDRAW followed by a space and the exact version"
|
||||
required: true
|
||||
type: string
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
# Share the publication lock with release.yml. A withdrawal and a publication
|
||||
# must never mutate channel pointers concurrently.
|
||||
concurrency:
|
||||
group: dws-release-publication
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
withdraw:
|
||||
name: Withdraw release from every distribution channel
|
||||
environment: release-withdrawal
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 180
|
||||
permissions:
|
||||
actions: read
|
||||
contents: write
|
||||
steps:
|
||||
- name: Verify withdrawal environment protection
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const { owner, repo } = context.repo;
|
||||
const response = await github.request(
|
||||
"GET /repos/{owner}/{repo}/environments/{environment_name}",
|
||||
{ owner, repo, environment_name: "release-withdrawal" },
|
||||
);
|
||||
const reviewerRule = response.data.protection_rules.find(
|
||||
(rule) => rule.type === "required_reviewers",
|
||||
);
|
||||
if (
|
||||
!reviewerRule ||
|
||||
reviewerRule.prevent_self_review !== true ||
|
||||
!Array.isArray(reviewerRule.reviewers) ||
|
||||
reviewerRule.reviewers.length === 0
|
||||
) {
|
||||
core.setFailed("release-withdrawal must require a reviewer and prevent self-review");
|
||||
return;
|
||||
}
|
||||
if (response.data.deployment_branch_policy?.protected_branches !== true) {
|
||||
core.setFailed("release-withdrawal must allow only protected branches");
|
||||
}
|
||||
if (response.data.can_admins_bypass !== false) {
|
||||
core.setFailed("release-withdrawal must not allow administrator bypass");
|
||||
}
|
||||
|
||||
- name: Require the exact current official default-branch commit
|
||||
uses: actions/github-script@v7
|
||||
with:
|
||||
script: |
|
||||
const expectedRepository = "DingTalk-Real-AI/dingtalk-workspace-cli";
|
||||
const defaultBranch = context.payload.repository.default_branch;
|
||||
if (context.eventName !== "workflow_dispatch") {
|
||||
core.setFailed("release withdrawal accepts workflow_dispatch only");
|
||||
return;
|
||||
}
|
||||
if (`${context.repo.owner}/${context.repo.repo}` !== expectedRepository) {
|
||||
core.setFailed(`release withdrawal is restricted to ${expectedRepository}`);
|
||||
return;
|
||||
}
|
||||
if (context.ref !== `refs/heads/${defaultBranch}`) {
|
||||
core.setFailed(`release withdrawal must be dispatched from ${defaultBranch}`);
|
||||
return;
|
||||
}
|
||||
const branch = await github.rest.git.getRef({
|
||||
...context.repo,
|
||||
ref: `heads/${defaultBranch}`,
|
||||
});
|
||||
if (branch.data.object.sha !== context.sha) {
|
||||
core.setFailed(
|
||||
`default branch advanced to ${branch.data.object.sha}; re-dispatch from the new head`,
|
||||
);
|
||||
}
|
||||
|
||||
- name: Check out trusted withdrawal tooling
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
ref: ${{ github.sha }}
|
||||
fetch-depth: 0
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Node.js for npm channel withdrawal
|
||||
uses: actions/setup-node@v4
|
||||
with:
|
||||
node-version: "22"
|
||||
registry-url: "https://registry.npmjs.org"
|
||||
|
||||
- name: Withdraw immutable release and roll back channels
|
||||
id: withdrawal
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}
|
||||
GITHUB_EVENT_DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
|
||||
WITHDRAW_VERSION: ${{ inputs.version }}
|
||||
WITHDRAW_REASON: ${{ inputs.reason }}
|
||||
WITHDRAW_CONFIRMATION: ${{ inputs.confirmation }}
|
||||
OSS_ACCESS_KEY_ID: ${{ secrets.OSS_ACCESS_KEY_ID }}
|
||||
OSS_ACCESS_KEY_SECRET: ${{ secrets.OSS_ACCESS_KEY_SECRET }}
|
||||
OSS_ENDPOINT: ${{ secrets.OSS_ENDPOINT }}
|
||||
OSS_BUCKET: ${{ secrets.OSS_BUCKET }}
|
||||
OSS_PREFIX: ${{ secrets.OSS_PREFIX }}
|
||||
GITEE_TOKEN: ${{ secrets.GITEE_TOKEN }}
|
||||
GITEE_USER: ${{ secrets.GITEE_USER }}
|
||||
GITEE_REPO: ${{ secrets.GITEE_REPO }}
|
||||
DWS_GITEE_ENABLED: ${{ vars.ENABLE_GITEE_UPLOAD_FALLBACK == 'true' && 'true' || 'false' }}
|
||||
HOMEBREW_PR_TOKEN: ${{ secrets.HOMEBREW_PR_TOKEN }}
|
||||
run: |
|
||||
./scripts/release/withdraw-release.sh \
|
||||
"$WITHDRAW_VERSION" \
|
||||
"$WITHDRAW_REASON" \
|
||||
"$WITHDRAW_CONFIRMATION"
|
||||
|
||||
- name: Report withdrawal boundary
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
VERSION: ${{ inputs.version }}
|
||||
RESULT: ${{ steps.withdrawal.outcome }}
|
||||
run: |
|
||||
{
|
||||
echo "### Release withdrawal: ${VERSION}"
|
||||
echo
|
||||
echo "- Workflow result: ${RESULT}"
|
||||
echo "- Success means every configured channel was verified and the permanent withdrawn/${VERSION} tombstone remains as the version-reuse barrier."
|
||||
echo "- Failure may occur before or after the tombstone/channel mutations; inspect the failed step and rerun the exact same inputs after fixing the cause."
|
||||
echo "- The problem GitHub Release and original tag are removed after npm and every tag-enabled/configured mirror are rolled back, so GitHub installers stop resolving the bad version while the Homebrew rollback PR is reviewed."
|
||||
echo "- npm is deprecated rather than unpublished; already-installed clients cannot be remotely downgraded."
|
||||
echo "- If a Homebrew rollback PR was opened, this run remains failed until that PR is independently reviewed, merged, and the workflow is rerun."
|
||||
} >> "$GITHUB_STEP_SUMMARY"
|
||||
@@ -54,3 +54,11 @@ test/dev_functional/results.jsonl
|
||||
/coverage-policy.txt
|
||||
/coverage.html
|
||||
dwsbin
|
||||
|
||||
# Local shortcut eval / real-backend capture artifacts — may contain real PII
|
||||
# (employee names/emails, userIds, conversation & message IDs). Never commit.
|
||||
/docs/shortcut-real-read-results.json
|
||||
/docs/shortcut-real-write-results.json
|
||||
/docs/shortcut-comparison.html
|
||||
/docs/shortcut-gsb-eval.*
|
||||
/scripts/run_shortcut_real_read_matrix.py
|
||||
|
||||
+96
-1
@@ -6,6 +6,102 @@ The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and th
|
||||
|
||||
## [Unreleased]
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Schema CLI path compatibility** — user-facing Schema lookups once again accept space-, dot-, and slash-separated CLI paths without weakening strict canonical identity resolution.
|
||||
- **Plugin CLI overlays** — installed plugins register their manifest-authored command trees again for HTTP and stdio servers, and a plugin may now replace a hidden compatibility fallback (for example `conference`) instead of being skipped as a distribution conflict.
|
||||
|
||||
## [1.0.53] - 2026-07-21
|
||||
|
||||
This release promotes the validated `v1.0.53-beta.7` baseline to stable. It adds enterprise onboarding, declarative shortcuts, Sheet/Aitable writes, multi-account profiles, and broader personal IM events, while hardening authentication and the guarded release path.
|
||||
|
||||
### Added
|
||||
|
||||
- **Enterprise and office command coverage** — adds enterprise creation, employee invitation, and account provisioning commands; 366 declarative service shortcuts; Sheet import commands; and Aitable workflow create/update support with reviewed Schema contracts.
|
||||
- **Multiple accounts in one DingTalk organization** — profiles can distinguish accounts by organization and user, select them explicitly, and log out one account or an entire organization without overwriting another account's credentials.
|
||||
- **Expanded personal IM event subscriptions** (#651) — adds read-receipt, recall, and reaction events for one-to-one and group chats, plus specified-sender subscriptions by staff ID or OpenDingTalk ID.
|
||||
- **Official multi-platform Homebrew channel** — ships separate stable and keg-only beta Formulae for macOS and Linux across amd64 and arm64, with isolated update PRs.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Personal event output contract** (#651) — `event consume` now emits event-specific top-level structured fields; scripts that consumed the former transport envelope must use the flat fields or select `-f raw`, while `--debug-raw-events` retains the diagnostic envelope.
|
||||
- **Guarded release lifecycle** — beta/stable publication now uses explicit promotion, immutable delivery proofs, protected recovery, and tag-bound optional OSS policy; an unprovisioned OSS mirror is sealed as `deferred` so GitHub, npm, and Homebrew are not blocked.
|
||||
- **Relaxed stable promotion contract** (#729) — a stable release still requires a delivered, non-withdrawn beta baseline in its commit history, but no longer requires a byte-identical tree with that beta; reviewed commits merged to `main` after the beta can now ship in the stable release. Local releases now accept any sealed commit contained in `main` history and push only the release tag, so `main` is never frozen during the beta-to-stable window.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Authentication and credential reliability** — organization-policy denials stop before mutation or polling, long-running clients reload and refresh access tokens consistently, concurrent credential writes are atomic, and Windows portable-auth commands fail before reading or writing unsupported credential bundles.
|
||||
- **Command validation and compatibility** — invalid Sheet/task targets fail locally, IM shortcuts preserve AI-tag and alias compatibility, and Aitable import uploads require and forward a positive file size.
|
||||
- **Release publication reliability** — GitHub draft publication is bound to one verified release ID and exact assets, preflight uses isolated installer worktrees, guarded local tags remain compatible, cloud planning fingerprints the actual allocated release refs, and npm channel verification waits for bounded registry propagation without moving tags.
|
||||
- **Package-manager version verification** (#735) — npm-vendored, Homebrew-installed, and packaged release binaries are now verified by searching their raw bytes for the injected version marker, so a correctly versioned stable binary is no longer rejected when the short version marker coalesces with adjacent printable linker metadata; incorrect or missing markers still fail closed.
|
||||
|
||||
## [1.0.53-beta.7] - 2026-07-21
|
||||
|
||||
This beta validates bounded npm channel verification after registry publication.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **npm dist-tag eventual consistency** — Release delivery now tolerates a briefly stale `latest` or `beta` read after publishing by retrying only when npm reports a valid older version. Registry errors, invalid or incomparable tags, and channels that never converge still fail closed without moving any tag during verification.
|
||||
|
||||
## [1.0.53-beta.6] - 2026-07-21
|
||||
|
||||
This beta validates guarded local release compatibility and tag-bound OSS deferral so an unprovisioned mirror cannot block the primary release channels.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Tag-bound optional OSS release mirror** — Official cloud Release runs no longer block GitHub, npm, and Homebrew delivery when an OSS bucket has not been provisioned. Cloud tags immutably record `OSS-Mirror: enabled|deferred`; publication, repair, and withdrawal consume that sealed policy instead of the current repository variable. Enabled releases remain fail-closed, while deferred releases skip the nonexistent channel and cannot be backfilled without a future audited repair proof.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Guarded local release compatibility** — The tag-push Release workflow now accepts the `Channel`-only annotated tags created by the guarded local release entry while continuing to reject any partial cloud-only seal metadata.
|
||||
- **Cloud release tag allocation fingerprint** — Release planning now fingerprints the actual `v*` and `withdrawn/v*` refs fetched from GitHub, matching the seal job's API view instead of hashing an empty non-wildcard ref prefix and rejecting every publish before tag creation.
|
||||
|
||||
## [1.0.53-beta.5] - 2026-07-21
|
||||
|
||||
This beta validates long-running access-token recovery and the faster, recoverable guarded release path introduced after v1.0.53-beta.4.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Fast guarded beta and stable releases** — successful local release checks now leave a six-hour proof bound to the exact version, commit, repository identity, remote `main`, and stable baseline, so the subsequent guarded `--publish` invocation revalidates authority without repeating tests and packaging. A default-branch governance smoke uses the same dedicated immutable-release credential as the tag workflow before any tag is allocated.
|
||||
- **Protected existing-tag recovery** — `dws-release recover <version>` can resume a failed, unpublished annotated tag through the normal contract, build, Developer ID signing, immutable GitHub Release, Homebrew, npm, and OSS jobs. Recovery requires the exact tag object, peeled commit, failed tag-push run, typed version confirmation, and the protected `release-recovery` environment; successful runs are accepted as future beta/stable delivery evidence.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Long-running event authentication recovery** — personal and portal event streams resolve the current access token for every ticket request, refresh a server-rejected token with compare-and-refresh semantics, and reconnect with backoff when refresh is temporarily blocked by network failures, rate limits, or 5xx responses.
|
||||
- **Consistent access-token caching and errors** — runtime, recovery, Skill, PAT polling, and personal/portal event clients now resolve user access tokens through one expiry- and publication-aware manager, so long-running processes reload rotated credentials while keychain, refresh, parse, permission, and cancellation failures remain observable instead of being collapsed into “not authenticated.”
|
||||
- **Tag-push GitHub Release publication** — Draft publication now locks one GitHub Release database ID, verifies its exact tag, channel, notes, recovery marker, asset set, and uploaded bytes, then publishes and rechecks that same ID as immutable. Recovery runs use the trusted default-branch release helpers instead of the sealed tag's historical scripts, fixing the Draft-only `GET /releases/tags/{tag}` 404 without allowing the release identity to drift during recovery.
|
||||
- **Release preflight reliability** — source-mode installer tests now use isolated temporary checkouts and HOME directories instead of overwriting and deleting the real repository `dws` binary, release preflight explicitly rebuilds before policy checks, and the full-suite runner gives the growing script package a non-flaky five-minute per-suite budget.
|
||||
|
||||
## [1.0.53-beta.4] - 2026-07-17
|
||||
|
||||
This beta validates the expanded personal IM event subscriptions and the flattened `event consume` structured output introduced after v1.0.53-beta.3.
|
||||
|
||||
### Added
|
||||
|
||||
- **Expanded personal IM event subscriptions** (#651) — adds one-to-one and group events for message read receipts, recalls, and reactions; publishes the specified-sender receive event; and lets one-to-one/sender subscriptions target either a staff `--user` or an `--open-dingtalk-id`. Event Schema now exposes these alternatives through machine-readable parameter constraints.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Personal event structured output is now flat** (#651) — `event consume` projects NDJSON/JSON/pretty/compact output into event-specific top-level DTOs, so consumers read fields such as `content`, `sender`, and `conversation_id` directly instead of parsing `.data | fromjson`. This is a breaking change for scripts using the former transport envelope; the original server payload remains available through `-f raw`, while `--debug-raw-events` preserves the full diagnostic envelope.
|
||||
|
||||
## [1.0.53-beta.3] - 2026-07-17
|
||||
|
||||
This beta validates multi-account profile support and the post-v1.0.53-beta.2 compatibility fixes for Windows portable authentication, IM shortcuts, and Aitable import uploads.
|
||||
|
||||
### Added
|
||||
|
||||
- **Multiple accounts in one DingTalk organization** — profiles are keyed by `corpId:userId`, `--profile` accepts organization IDs/names plus user IDs/names, and organization-only selection uses its explicitly remembered current account or asks for an exact account when ambiguous.
|
||||
|
||||
### Changed
|
||||
|
||||
- **Profile-scoped logout and consistent token storage** — `dws auth logout --profile` can remove one account or every account in an organization, while identity token slots remain the source of truth and legacy organization/global mirrors stay compatible without overwriting newer account credentials.
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Windows portable-auth contract** — `dws auth export` and `dws auth import` now fail early without reading credentials, bundles, or writing files instead of claiming portable-bundle support for DPAPI-protected HKCU Registry credentials.
|
||||
- **IM shortcut message tags and compatibility aliases** (#646) — IM send shortcuts now add the same AI-sent marker as `chat message send` by default, support `--ai-tag=false` to opt out, and preserve compatible search, conversation-ID, and page-size aliases.
|
||||
- **Aitable import upload file-size validation** (#654) — `dws aitable import upload` and `dws aitable +import-upload` now require a positive `--file-size` and always send it to the upload-preparation API, preventing invalid requests without the actual file size.
|
||||
|
||||
## [1.0.53-beta.2] - 2026-07-16
|
||||
|
||||
This beta validates the accumulated post-v1.0.52 command surface, release automation, and runtime hardening changes, including enterprise contact onboarding, declarative shortcuts, Sheet/Aitable writes, multi-platform Homebrew formulas, and credential and target-validation fixes.
|
||||
@@ -26,7 +122,6 @@ This beta validates the accumulated post-v1.0.52 command surface, release automa
|
||||
|
||||
- **PAT organization-policy denials stop immediately** — `PAT_ORG_POLICY_DENIED` now remains terminal even if a backend also returns `flowId`, authorization URLs, or client credentials; the CLI does not mutate process credentials, open a browser, poll, or retry until an organization administrator changes the policy.
|
||||
- **Sheet and task invalid-target failures** — `sheet range read/get` now rejects a null cell-info response instead of printing `null` and exiting successfully, while task completion and attachment listing verify that a task exists before calling lenient backend endpoints. Attachment listing is also published through Runtime Schema for schema-first Agent discovery.
|
||||
- **Windows portable-auth contract** — `dws auth export` and `dws auth import` now fail early without reading credentials, bundles, or writing files instead of claiming portable-bundle support for DPAPI-protected HKCU Registry credentials.
|
||||
- **Concurrent credential writes and reentrant CLI execution** — secure-token writers now use isolated, exclusive temporary files before atomic replacement so concurrent processes cannot remove each other's in-flight data, and repeated in-process CLI runs close the previous file logger before replacing it instead of retaining the prior log-file handle.
|
||||
|
||||
## [1.0.52] - 2026-07-14
|
||||
|
||||
+1
-1
@@ -36,8 +36,8 @@ Common repository checks already used here include:
|
||||
./scripts/policy/check-open-source-assets.sh
|
||||
go test ./...
|
||||
make test
|
||||
make test-plan
|
||||
make lint
|
||||
bash test/scripts/run_all_tests.sh --jobs 8
|
||||
./scripts/policy/check-generated-drift.sh
|
||||
./scripts/policy/check-command-surface.sh --strict
|
||||
./scripts/release/verify-package-managers.sh
|
||||
|
||||
@@ -1,33 +1,33 @@
|
||||
class DingtalkWorkspaceCliBeta < Formula
|
||||
desc "Automate DingTalk workspace tasks from the terminal (beta channel)"
|
||||
homepage "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli"
|
||||
version "1.0.53-beta.1"
|
||||
version "1.0.53-beta.4"
|
||||
license "Apache-2.0"
|
||||
keg_only "it is the beta channel and conflicts with dingtalk-workspace-cli"
|
||||
|
||||
on_macos do
|
||||
if Hardware::CPU.arm?
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.1/dws-darwin-arm64.tar.gz"
|
||||
sha256 "7fef5add684189bfef16a1731fab4883fa096115dcd645c79bf9b3918758dbcb"
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.4/dws-darwin-arm64.tar.gz"
|
||||
sha256 "32a442d5b42dfed8512a695a7cef513722db6912f3c3168954cfaefb68e0b075"
|
||||
else
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.1/dws-darwin-amd64.tar.gz"
|
||||
sha256 "8cde0fd62849dcb52bc053022361f983953bd5a4bac1edfc96076a84636debbe"
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.4/dws-darwin-amd64.tar.gz"
|
||||
sha256 "00c694677b9ce2e1a711535740681defe0a5f83d6b72140f305483501628faf8"
|
||||
end
|
||||
end
|
||||
|
||||
on_linux do
|
||||
if Hardware::CPU.arm?
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.1/dws-linux-arm64.tar.gz"
|
||||
sha256 "17989f94b422d6490249c5e095049a9efc1c3861176fc249d5e7eea283f5e675"
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.4/dws-linux-arm64.tar.gz"
|
||||
sha256 "98516620e861e516cf846cf418f1a0ad5c5eb9a8681bd066b1fd29f4847aca6a"
|
||||
else
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.1/dws-linux-amd64.tar.gz"
|
||||
sha256 "439d6a1ffc9572ed461b77a8bb27f6ddb4d4c9ce9d58474c0180955fc8a1d8d7"
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.4/dws-linux-amd64.tar.gz"
|
||||
sha256 "266df80e8a989971789a157dd134685a7c9eda01dd1d082719ec9e09d5a07bf0"
|
||||
end
|
||||
end
|
||||
|
||||
resource "skills" do
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.1/dws-skills.zip"
|
||||
sha256 "a93ba8a73f2e319038a7bcb32038b9e7bd93d1e9af012a430ebf994498c140d4"
|
||||
url "https://github.com/DingTalk-Real-AI/dingtalk-workspace-cli/releases/download/v1.0.53-beta.4/dws-skills.zip"
|
||||
sha256 "6770511ab9b04b4d97da1858069ff694830ba91d47c1e8564fd85856f95b016e"
|
||||
end
|
||||
|
||||
def install
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
GO ?= go
|
||||
DWS_PACKAGE_VERSION ?= 0.0.0-test
|
||||
REMOTE ?=
|
||||
PUBLISH ?= 0
|
||||
YES ?= 0
|
||||
DWS_POLICY_TMPDIR ?= $(CURDIR)/.worktrees/policy-tmp
|
||||
POLICY_GOTMPDIR ?= $(DWS_POLICY_TMPDIR)/go
|
||||
POLICY_ENV = DWS_POLICY_TMPDIR="$(DWS_POLICY_TMPDIR)" GOTMPDIR="$(POLICY_GOTMPDIR)"
|
||||
GO_SOURCE_LIST = git ls-files -z --cached --others --exclude-standard -- '*.go'
|
||||
|
||||
.PHONY: all help build rebuild test lint fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema generate-schema-agent-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
|
||||
.PHONY: all help build rebuild test test-plan lint format-check fmt policy edition-test interface-integrity authoritative-interface-integrity coverage-gate coverage-gate-platform update-interface-baseline reset-interface-baseline schema-compatibility skill-command-integrity cli-smoke mock-mcp-smoke test-schema-agent-examples generate-schema generate-schema-agent-metadata generate-schema-catalog package release release-pre release-stable changelog-pre changelog-stable publish-homebrew-formula setup-hooks
|
||||
|
||||
all: setup-hooks fmt lint build test rebuild
|
||||
|
||||
@@ -14,13 +16,15 @@ help:
|
||||
@printf "Available targets:\n"
|
||||
@printf " make build - Build the dws CLI binary\n"
|
||||
@printf " make test - Run the Go test suite\n"
|
||||
@printf " make lint - Run formatting checks and golangci-lint when available\n"
|
||||
@printf " make fmt - Format Go source files\n"
|
||||
@printf " make test-plan - Verify every default Go package belongs to one CI test shard\n"
|
||||
@printf " make lint - Run formatting checks, go vet, and staticcheck\n"
|
||||
@printf " make format-check - Check all repository Go source files with gofmt\n"
|
||||
@printf " make fmt - Format all repository Go source files\n"
|
||||
@printf " make policy - Check the built dws plus open-source and Schema policies\n"
|
||||
@printf " make interface-integrity - Check historical commands and help contracts still work\n"
|
||||
@printf " make authoritative-interface-integrity BASE_REF=<ref> - Check the Git-owned PR merge-base\n"
|
||||
@printf " make coverage-gate BASE_REF=<ref> - Enforce overall non-regression and changed-code coverage\n"
|
||||
@printf " make coverage-gate-platform BASE_REF=<ref> PROFILE=<file> - Enforce native-platform changed-code coverage\n"
|
||||
@printf " make coverage-gate BASE_REF=<ref> - Enforce overall non-regression and 100%% changed-code coverage\n"
|
||||
@printf " make coverage-gate-platform BASE_REF=<ref> PROFILE=<file> - Enforce 100%% native changed-code coverage\n"
|
||||
@printf " make update-interface-baseline - Add new CLI contracts without removing history\n"
|
||||
@printf " make reset-interface-baseline - DANGEROUS: replace all CLI compatibility history\n"
|
||||
@printf " make schema-compatibility BASE_REF=<ref> - Check the complete Schema contract against the PR merge-base\n"
|
||||
@@ -45,13 +49,32 @@ rebuild:
|
||||
@./scripts/dev/build.sh
|
||||
|
||||
test:
|
||||
@./test/scripts/run_all_tests.sh
|
||||
@DWS_PACKAGE_VERSION="$(DWS_PACKAGE_VERSION)" $(GO) test -count=1 -timeout=10m ./...
|
||||
|
||||
test-plan:
|
||||
@./scripts/ci/test-packages.sh verify
|
||||
|
||||
lint:
|
||||
@./scripts/dev/lint.sh
|
||||
|
||||
format-check:
|
||||
@set -eu; \
|
||||
go_files="$$(mktemp "$${TMPDIR:-/tmp}/dws-go-files.XXXXXX")"; \
|
||||
trap 'rm -f "$$go_files"' EXIT HUP INT TERM; \
|
||||
$(GO_SOURCE_LIST) > "$$go_files"; \
|
||||
unformatted="$$(xargs -0 sh -c 'if [ "$$#" -gt 0 ]; then exec gofmt -l -- "$$@"; fi' sh < "$$go_files")"; \
|
||||
if [ -n "$$unformatted" ]; then \
|
||||
printf '%s\n' "$$unformatted"; \
|
||||
printf '%s\n' "Go files are not formatted. Run 'make fmt'." >&2; \
|
||||
exit 1; \
|
||||
fi
|
||||
|
||||
fmt:
|
||||
@find cmd internal test scripts/policy -name '*.go' -print0 2>/dev/null | xargs -0r gofmt -w
|
||||
@set -eu; \
|
||||
go_files="$$(mktemp "$${TMPDIR:-/tmp}/dws-go-files.XXXXXX")"; \
|
||||
trap 'rm -f "$$go_files"' EXIT HUP INT TERM; \
|
||||
$(GO_SOURCE_LIST) > "$$go_files"; \
|
||||
xargs -0 sh -c 'if [ "$$#" -gt 0 ]; then exec gofmt -w -- "$$@"; fi' sh < "$$go_files"
|
||||
|
||||
policy:
|
||||
@mkdir -p "$(POLICY_GOTMPDIR)"
|
||||
|
||||
@@ -283,16 +283,22 @@ Credentials are securely persisted after first login (Keychain). Subsequent runs
|
||||
<details>
|
||||
<summary><strong>Multiple organizations (profiles)</strong></summary>
|
||||
|
||||
`dws` can stay logged in to several DingTalk organizations at once. Each organization is one **profile**; the current profile decides which org a command runs against (credentials are stored per organization).
|
||||
`dws` can stay logged in to several DingTalk accounts at once, including multiple accounts in the same organization. A profile is uniquely identified by `corpId:userId`; the current profile decides which identity a command runs as.
|
||||
|
||||
```bash
|
||||
dws auth login # log in to another org → adds a profile (first login becomes the primary)
|
||||
dws profile list # list logged-in orgs (primary / current marker, status)
|
||||
dws profile switch <name|corpId> # switch the default org (use - to toggle back to the previous one)
|
||||
dws --profile <name|corpId> contact user search --query "..." # run one command against a specific org, without changing the default
|
||||
dws auth login # add or refresh one account
|
||||
dws profile list # list every logged-in account
|
||||
dws profile switch <corpId:userId> # persistently switch; use - to toggle back
|
||||
dws profile switch "<corpName>:<userName>" # friendly input; names must be unique
|
||||
dws --profile <corpId> contact user search --query "..." # use that org's explicitly recorded current account
|
||||
dws --profile <corpId:userId> contact user search --query "..." # use one exact account without changing the default
|
||||
```
|
||||
|
||||
Cross-org reads are orchestrated by the agent rather than a built-in `--all-orgs`: list the profiles, run the query per org with `--profile`, then merge. Writes default to the current org only — confirm the target org before writing across orgs.
|
||||
Selectors support `corpId:userId`, `corpId:userName`, `corpName:userId`, and `corpName:userName`. Friendly names are input aliases only; use the stable `profile` value returned by `profile list` for automation. Duplicate organization or account names fail with explicit `corpId:userId` candidates. If an organization has multiple accounts but no recorded current account, `--profile <corpId>` fails instead of choosing the first or most recently used account.
|
||||
|
||||
`currentProfile`, `previousProfile`, and per-organization defaults are stored as exact identities. `primaryProfile` remains in JSON only for compatibility and is not used for selection. `profile list` reads status and expiry from each real identity Token without refreshing it. `auth logout --profile <corpId>` removes all local accounts in that organization; an exact selector or local profile name removes one account.
|
||||
|
||||
Cross-org reads are orchestrated by the agent rather than a built-in `--all-orgs`: list profiles, group by `corpId`, and use the unique `isOrgCurrent=true` account for each organization. If a multi-account organization has no default, ask the user to choose an account first. Writes default to the current account — confirm both organization and account before cross-org writes.
|
||||
|
||||
On macOS, an unreadable registered token slot blocks a new OAuth login rather than risking a mixed Keychain/file-DEK state. If normal terminal commands can still read the login while a sandbox using `DWS_DISABLE_KEYCHAIN=1` cannot, migrate the legacy and profile auth entries without exposing tokens:
|
||||
|
||||
@@ -302,7 +308,7 @@ env -u DWS_DISABLE_KEYCHAIN dws auth migrate-keychain --to file-dek --yes --form
|
||||
DWS_DISABLE_KEYCHAIN=1 dws auth status --format json
|
||||
```
|
||||
|
||||
The migration validates every selected auth ciphertext before writing, ignores unrelated application secrets, and can be rerun after an interrupted commit. If validation identifies genuinely damaged ciphertext, remove only the affected profile with `dws auth logout --profile <name|corpId>`, then log in again. Use `dws auth reset` only when you intend to discard every local profile.
|
||||
The migration validates every selected auth ciphertext before writing, ignores unrelated application secrets, and can be rerun after an interrupted commit. If validation identifies genuinely damaged ciphertext, remove only the affected account with `dws auth logout --profile <corpId:userId>`, or all accounts in one organization with `--profile <corpId>`, then log in again. Use `dws auth reset` only when you intend to discard every local profile.
|
||||
|
||||
</details>
|
||||
|
||||
@@ -489,6 +495,9 @@ dws event consume user_im_message_receive_at -f ndjson
|
||||
# Listen for one-to-one messages with a specified user
|
||||
dws event consume user_im_message_receive_o2o --user <userId> -f ndjson
|
||||
|
||||
# Listen by openDingtalkId (external contact, bot, or cross-organization identity)
|
||||
dws event consume user_im_message_receive_o2o --open-dingtalk-id <openDingtalkId> -f ndjson
|
||||
|
||||
# Listen for messages in a specified group
|
||||
dws event consume user_im_message_receive_group --group <openConversationId> -f ndjson
|
||||
|
||||
@@ -497,6 +506,8 @@ dws event status
|
||||
dws event stop <subscribe_id>
|
||||
```
|
||||
|
||||
For one-to-one and specified-sender events, use exactly one target identity: `--user` for an internal `userId`, or `--open-dingtalk-id` for an `openDingtalkId`. The CLI does not infer or convert between these identity types.
|
||||
|
||||
| Feature | Details |
|
||||
|---------|---------|
|
||||
| Managed lifecycle | `consume` creates or reuses the personal subscription; `stop` cancels it and cleans local state |
|
||||
|
||||
+18
-7
@@ -280,16 +280,22 @@ dws auth login --client-id <your-app-key> --client-secret <your-app-secret>
|
||||
<details>
|
||||
<summary><strong>多组织(profile)</strong></summary>
|
||||
|
||||
`dws` 可以同时登录多个钉钉组织。一个组织就是一个 **profile**,当前 profile 决定本次命令操作哪个组织(凭证按组织分别存储)。
|
||||
`dws` 可以同时登录多个钉钉账号,同一组织也能保留多个账号。一个 profile 由 `corpId + userId` 唯一确定。
|
||||
|
||||
```bash
|
||||
dws auth login # 再登录一个组织 → 新增一个 profile(首次登录的为主组织)
|
||||
dws profile list # 列出已登录组织(主 / 当前标记、状态)
|
||||
dws profile switch <名称|corpId> # 切换默认组织(用 - 切回上一个)
|
||||
dws --profile <名称|corpId> contact user search --query "..." # 单次对指定组织执行,不改默认组织
|
||||
dws auth login # 新增或刷新一个账号
|
||||
dws profile list # 列出全部账号,profile 字段是稳定的 corpId:userId
|
||||
dws profile switch <corpId:userId> # 持久切换账号;用 - 切回上一个
|
||||
dws profile switch "组织名:用户名" # 名称输入要求唯一
|
||||
dws --profile <corpId> contact user search --query "..." # 使用该组织明确记录的当前账号
|
||||
dws --profile <corpId:userId> contact user search --query "..." # 单次精确指定账号,不改默认账号
|
||||
```
|
||||
|
||||
跨组织读取由 agent 编排,而非内置 `--all-orgs`:先 `dws profile list` 拿到组织,再对每个组织带 `--profile` 各查一遍,然后合并。写操作默认只在当前组织进行——跨组织写之前先确认目标组织。
|
||||
支持 `corpId:userId`、`corpId:userName`、`corpName:userId`、`corpName:userName`。名称只用于输入,自动化应使用 `profile list` 返回的稳定 `profile`。组织名或用户名重名时会列出候选并报错;同组织多账号但没有明确当前账号时,只传组织也会报错,不会选择第一项或最近使用账号。
|
||||
|
||||
`currentProfile`、`previousProfile` 和组织默认账号都保存精确身份。`primaryProfile` 只为 JSON 兼容保留,不再参与选择。`profile list` 直接读取各身份 Token 计算状态和到期时间,不触发刷新。`auth logout --profile <corpId>` 退出该组织全部账号;精确选择器或本地 profile 名只退出一个账号。
|
||||
|
||||
跨组织读取由 agent 编排,而非内置 `--all-orgs`:先 `dws profile list`,每个组织使用唯一的 `isOrgCurrent=true` 账号;若多账号组织没有默认账号,先让用户指定账号。写操作默认只在当前账号执行——跨组织写之前先确认目标组织和账号。
|
||||
|
||||
macOS 下,如果已登记的 token slot 无法解密,为避免把系统 Keychain 和 file-DEK 写成混合状态,新的 OAuth 登录会直接拒绝。如果普通终端仍能读取登录态、只有设置 `DWS_DISABLE_KEYCHAIN=1` 的沙箱读不到,可在不暴露 token 的情况下迁移 legacy 与各 profile 的认证条目:
|
||||
|
||||
@@ -299,7 +305,7 @@ env -u DWS_DISABLE_KEYCHAIN dws auth migrate-keychain --to file-dek --yes --form
|
||||
DWS_DISABLE_KEYCHAIN=1 dws auth status --format json
|
||||
```
|
||||
|
||||
迁移会先验证全部认证密文再写入、忽略无关的应用密钥;提交中断后可安全重跑。如果预检确认是密文本身损坏,报错会给出对应 `corpId`;只清理这个组织可执行 `dws auth logout --profile <名称|corpId>`,再重新登录。只有确认要丢弃全部本地 profile 时才用 `dws auth reset`。
|
||||
迁移会先验证全部认证密文再写入、忽略无关的应用密钥;提交中断后可安全重跑。如果预检确认是密文本身损坏,优先使用 `dws auth logout --profile <corpId:userId>` 只清理受影响账号;只有确认要丢弃全部本地 profile 时才用 `dws auth reset`。
|
||||
|
||||
</details>
|
||||
|
||||
@@ -483,6 +489,9 @@ dws event consume user_im_message_receive_at -f ndjson
|
||||
# 监听与指定用户的单聊消息
|
||||
dws event consume user_im_message_receive_o2o --user <userId> -f ndjson
|
||||
|
||||
# 使用 openDingtalkId 监听外部联系人、机器人或跨组织身份
|
||||
dws event consume user_im_message_receive_o2o --open-dingtalk-id <openDingtalkId> -f ndjson
|
||||
|
||||
# 监听指定群的消息
|
||||
dws event consume user_im_message_receive_group --group <openConversationId> -f ndjson
|
||||
|
||||
@@ -491,6 +500,8 @@ dws event status
|
||||
dws event stop <subscribe_id>
|
||||
```
|
||||
|
||||
单聊和指定发送人事件必须且只能选择一种目标身份:企业内部 `userId` 使用 `--user`,`openDingtalkId` 使用 `--open-dingtalk-id`。CLI 不会自动猜测或转换身份类型。
|
||||
|
||||
| 特性 | 说明 |
|
||||
|------|------|
|
||||
| 自动编排 | `consume` 创建或复用个人订阅,`stop` 取消订阅并清理本地状态 |
|
||||
|
||||
@@ -43,3 +43,38 @@
|
||||
- `skills/`: bundled agent skills (mono/ and multi/ layouts)
|
||||
- `test/`: CLI, integration, contract, unit, and skill E2E tests
|
||||
- `scripts/`: install scripts, policy checks, and CI helpers
|
||||
|
||||
## Quality Pipeline
|
||||
|
||||
Quality enforcement is layered so a pull request receives fast, deterministic
|
||||
admission feedback without pretending that downstream integration has already
|
||||
run.
|
||||
|
||||
```mermaid
|
||||
flowchart TB
|
||||
PR["Pull request"] --> CA["CI"]
|
||||
subgraph CA_CHECKS["Nine required contexts"]
|
||||
L["Lint"]
|
||||
T["Test"]
|
||||
C["Coverage"]
|
||||
P["Policy"]
|
||||
E["Edition"]
|
||||
I["Interface Integrity"]
|
||||
A["AI Behavior"]
|
||||
S["CLI Smoke"]
|
||||
M["Mock MCP"]
|
||||
end
|
||||
CA --> CA_CHECKS
|
||||
CA_CHECKS --> MAIN["Protected main"]
|
||||
MAIN --> MP["Main Integration — 主干集成<br/>Multi-profile E2E"]
|
||||
MAIN --> PLATFORM["Risk-selected / release native platform validation"]
|
||||
MP --> RELEASE["Release delivery"]
|
||||
PLATFORM --> RELEASE
|
||||
```
|
||||
|
||||
Complete Multi-profile E2E and the ordinary full native-platform matrix are
|
||||
downstream of PR admission. PRs still run primary-environment assurance and
|
||||
fast cross-platform compilation; auth, keychain, OS-specific, installer, and
|
||||
release changes additionally select native platform tests before merge. See
|
||||
[`docs/ci-pr-gates.md`](ci-pr-gates.md) for the exact context and ruleset
|
||||
contract.
|
||||
|
||||
+62
-6
@@ -65,21 +65,77 @@ git diff --check
|
||||
## Homebrew Formula PR Automation
|
||||
|
||||
Official tag releases require the repository Actions secret
|
||||
`HOMEBREW_PR_TOKEN`. The `DingTalk-Real-AI` organization currently does not
|
||||
allow fine-grained personal access tokens to target this repository, so use a
|
||||
classic personal access token owned by a maintainer or release-bot account with
|
||||
`HOMEBREW_PR_TOKEN`. Prefer a fine-grained personal access token owned by a
|
||||
maintainer or release-bot account, limited to this repository with
|
||||
`Contents: write` and `Pull requests: write`. If organization policy prevents
|
||||
that account from targeting the repository, use a dedicated classic token with
|
||||
only the `public_repo` scope. Do not reuse a broad developer token.
|
||||
|
||||
Store the non-expiring token as the `HOMEBREW_PR_TOKEN` repository Actions
|
||||
secret. Replace it immediately if it is exposed, its owner loses repository
|
||||
access, or the release-bot ownership changes. The Release workflow uses this
|
||||
Store the dedicated token as the `HOMEBREW_PR_TOKEN` repository Actions secret
|
||||
and rotate it before its configured expiration. Replace it immediately if it is
|
||||
exposed, its owner loses repository access, or the release-bot ownership
|
||||
changes. The Release workflow uses this
|
||||
dedicated token only to push an `automation/homebrew-*` branch and open the
|
||||
stable or beta Formula PR. It does not push Formula changes directly to `main`.
|
||||
The default-branch governance preflight and every tag contract authenticate the
|
||||
token before publication, reject over-scoped classic tokens, confirm its
|
||||
identity, and run a controlled write canary. The canary pushes a unique
|
||||
`automation/homebrew-token-canary-*` branch with a `[skip ci]` commit, creates a
|
||||
draft PR, closes it, and deletes the branch with the same token. This proves both
|
||||
Contents and Pull requests write access before publication without merging
|
||||
anything. The gate also rejects reuse of `RELEASE_GOVERNANCE_TOKEN`.
|
||||
No maintainer environment variable is required when creating a tag. Using the
|
||||
built-in `GITHUB_TOKEN` is insufficient because organization policy prevents
|
||||
Actions from creating pull requests, and its generated PR events may require
|
||||
separate workflow approval.
|
||||
|
||||
## Release Governance and Recovery
|
||||
|
||||
Store `RELEASE_GOVERNANCE_TOKEN` as a dedicated Actions secret with only
|
||||
repository `Administration: read`. The immutable-releases REST endpoint is an
|
||||
administration setting and cannot be read by the workflow's built-in
|
||||
`GITHUB_TOKEN`. Both the default-branch governance preflight and the tag
|
||||
contract use this same credential so a missing or expired identity is detected
|
||||
before an irreversible tag is created.
|
||||
|
||||
Create a protected `release-recovery` environment limited to protected
|
||||
branches, with a required reviewer, self-review disabled, and administrator
|
||||
bypass disabled. The workflow reads the environment through the GitHub API and
|
||||
fails closed unless the required-reviewer, prevent-self-review, and protected-
|
||||
branch rules are present.
|
||||
Recovery is restricted to an existing annotated tag whose exact tag object,
|
||||
commit, and failed tag-push run all match; it then reuses the normal release
|
||||
jobs. Do not put publication secrets in temporary branches or create ad-hoc
|
||||
recovery workflows.
|
||||
|
||||
Cloud-sealed releases mirror to OSS only when the repository variable
|
||||
`ENABLE_OSS_MIRROR` is exactly `true`. Leave the variable unset while no Bucket
|
||||
is provisioned; GitHub, npm, and Homebrew delivery can then complete without
|
||||
running the OSS step. Once enabled, missing credentials, an invalid Bucket, or
|
||||
an upload failure remains fail-closed. The cloud tag immutably records the
|
||||
decision as `OSS-Mirror: enabled|deferred`; publication and withdrawal consume
|
||||
that sealed value instead of the variable's later state. Deferred releases
|
||||
cannot use `repair_oss_version`; enabling OSS applies to later release tags
|
||||
until an audited immutable repair marker is implemented.
|
||||
|
||||
If an immutable GitHub Release and npm package were delivered but an enabled
|
||||
downstream China mirror failed, dispatch the normal `Release` workflow from the
|
||||
protected default branch with exactly one of `repair_gitee_version` or
|
||||
`repair_oss_version`. Channel repair accepts a fully successful exact release,
|
||||
or a failed exact-tag run only when its latest attempt completed the release
|
||||
contract, build, Apple signature, immutable GitHub publication, and npm
|
||||
delivery checks for the exact tagged commit. OSS repair additionally requires
|
||||
the tag's sealed policy to be `enabled`. It then downloads and re-verifies the
|
||||
immutable assets before invoking only the selected mirror. For a failed
|
||||
release, an OSS repair requires the OSS step itself to be the recorded failure.
|
||||
A Gitee repair accepts either a failed Gitee job or a Gitee job that was
|
||||
skipped behind that OSS failure; the latter is an explicit Gitee backfill and
|
||||
does not claim that OSS has been repaired. Gitee repair requires `GITEE_TOKEN`,
|
||||
`GITEE_USER`, and `GITEE_REPO`; OSS repair requires `OSS_ACCESS_KEY_ID`,
|
||||
`OSS_ACCESS_KEY_SECRET`, `OSS_ENDPOINT`, and `OSS_BUCKET` (with optional
|
||||
`OSS_PREFIX`) as Actions secrets. Missing credentials fail the selected repair
|
||||
closed.
|
||||
|
||||
## Handoff Checklist
|
||||
|
||||
Before handoff, include:
|
||||
|
||||
+133
-108
@@ -1,133 +1,158 @@
|
||||
# Pull request quality gates
|
||||
# CI — PR 合入门禁
|
||||
|
||||
The repository defines five focused checks in addition to its existing CI:
|
||||
The pull-request admission layer has exactly nine required external contexts:
|
||||
|
||||
- **Interface Integrity** enforces backwards compatibility. Every historical
|
||||
command path and alias must still resolve, every historical command must
|
||||
still render `-h`, and historical flags must keep their type and shorthand.
|
||||
New commands, aliases, and flags are allowed. The same job compares the full
|
||||
complete `dws schema --all` contract with the PR merge-base, blocking removed
|
||||
products/tools/parameters, incompatible parameter or interface mappings,
|
||||
constraint drift, and safety-semantic drift. It also checks that executable
|
||||
`dws ...` references in `skills/**/*.md` resolve to real commands.
|
||||
Help compatibility covers command/alias/flag spelling, flag type and
|
||||
shorthand; descriptive prose may evolve without breaking the gate.
|
||||
- **Coverage** runs unit tests on every pull request and prints both overall and
|
||||
changed-code statement coverage. During the migration to the 80% repository
|
||||
target, overall coverage may not regress from a profile generated from the
|
||||
merge-base with the same test command, while changed production Go
|
||||
statements must meet 80%. Linux, Windows, and macOS each generate a native
|
||||
coverage profile for changed packages and enforce the threshold against
|
||||
changed files buildable on that platform, so build-tagged source cannot be
|
||||
hidden by an Ubuntu-only profile. Overall non-regression allows 0.1 percentage point of measurement
|
||||
variance to avoid failing unchanged code on test-path noise. Set
|
||||
`COVERAGE_ENFORCE_OVERALL=true` once repository coverage reaches 80% to make
|
||||
the overall target fail closed as well.
|
||||
CI generates the candidate, supporting, and merge-base profiles on three
|
||||
independent runners, then downloads all profiles into the aggregate
|
||||
`Coverage` job and applies the same fail-closed gate. The split changes only
|
||||
scheduling: the tested packages, profile contents, merge-base comparison,
|
||||
and final coverage thresholds remain unchanged.
|
||||
- **CLI Smoke** builds the release binary, reads the root command list from the
|
||||
structured Interface contract, and renders offline help for every public
|
||||
top-level command. It rejects Cobra's unknown-command root-help fallback and
|
||||
fails when the checked-in development fixture is stale.
|
||||
- **Mock MCP Smoke** runs the existing HTTP and stdio MCP lifecycle tests
|
||||
(`Initialize -> ListTools -> CallTool`).
|
||||
- **AI Behavior Check** applies to pull requests labeled `ai-generated`. It
|
||||
limits the change to 30 files and blocks release/CI infrastructure changes,
|
||||
including policy implementations and the checked-in Interface fixture.
|
||||
It uses `pull_request_target` without checking out PR code, so the policy
|
||||
cannot be bypassed by changing the workflow in the same pull request. The
|
||||
evaluator writes an `AI Behavior Check` commit status to the PR head SHA so
|
||||
GitHub rulesets can require it.
|
||||
| Required context | Contract |
|
||||
|---|---|
|
||||
| `Lint` | Stable PR revision classification, formatting, `go vet`, and Actionlint |
|
||||
| `Test` | Race/unit/release-script tests plus fast cross-platform compilation |
|
||||
| `Coverage` | Overall non-regression and 100% changed-code coverage |
|
||||
| `Policy` | Repository policy and the fail-closed CHANGELOG contract |
|
||||
| `Edition` | Edition contract tests |
|
||||
| `Interface Integrity` | CLI, Schema, Skill, and stable-release compatibility |
|
||||
| `AI Behavior` | Base-owned policy for PRs labeled `ai-generated` |
|
||||
| `CLI Smoke` | Offline help for every public top-level command |
|
||||
| `Mock MCP` | HTTP and stdio MCP lifecycle smoke tests |
|
||||
|
||||
## Running the compatibility gates
|
||||
The workflow display name is `CI`. Parallel helper
|
||||
jobs may implement `Test` and `Coverage`, but they are not ruleset contexts.
|
||||
Do not require an aggregate alias or a downstream integration check in place of
|
||||
the nine contracts above.
|
||||
|
||||
Run:
|
||||
`AI Behavior` is evaluated by a `pull_request_target` workflow that never
|
||||
checks out or executes PR code. It writes the exact `AI Behavior` status to the
|
||||
current PR head. Its Files API read is bracketed by base/head revision checks,
|
||||
so a synchronize race fails closed. The same workflow supplies a successful
|
||||
`AI Behavior` check run on protected `main` pushes for release governance.
|
||||
|
||||
## Exact CHANGELOG-only fast path
|
||||
|
||||
A pull request qualifies only when GitHub reports exactly one changed file,
|
||||
that file is an in-place modification of `CHANGELOG.md`, and the base and head
|
||||
both retain it as a regular non-executable `100644` blob. Add, delete, rename,
|
||||
symlink, executable-mode, and second-file changes do not qualify.
|
||||
|
||||
`Lint` classifies the Files API result only after verifying that the API's base
|
||||
and head equal the event revision both before and after pagination. `Policy`
|
||||
checks out GitHub's PR merge ref and verifies its parents:
|
||||
|
||||
```text
|
||||
HEAD^1 = pull_request.base.sha
|
||||
HEAD^2 = pull_request.head.sha
|
||||
```
|
||||
|
||||
It then runs:
|
||||
|
||||
```sh
|
||||
./scripts/policy/check-changelog-pr.sh \
|
||||
--fast-path "$PR_BASE_SHA" HEAD
|
||||
```
|
||||
|
||||
Because the verified PR diff contains only `CHANGELOG.md`, the validator and
|
||||
its policy dependencies in that merge tree are byte-for-byte the current base
|
||||
versions. Validation targets the synthetic merge tree, not the feature-branch
|
||||
tree, so a stale branch cannot supply an older validator or combine with newer
|
||||
base notes into an invalid final CHANGELOG.
|
||||
|
||||
All nine admission contexts are still emitted and must succeed. Expensive
|
||||
implementation helpers are skipped; the named contexts record that their code
|
||||
surface is unaffected. After merge, the protected `main` push executes the
|
||||
full admission suite.
|
||||
|
||||
Any PR that touches `CHANGELOG.md` but also changes another file runs the same
|
||||
content contract in `Policy` with `--content-only`. That mode permits the
|
||||
second file but still rejects invalid dates or versions, missing bullets,
|
||||
placeholder `TODO`/`TBD`, unmanaged-section changes, and unsafe tree modes.
|
||||
Adding a second file therefore cannot bypass CHANGELOG validation.
|
||||
|
||||
## Platform and downstream boundaries
|
||||
|
||||
Ordinary PRs run the primary Linux assurance plus fast Darwin/Windows compile
|
||||
checks. Full native macOS/Windows tests and platform coverage run on a PR only
|
||||
when its diff touches auth, keychain, OS-specific Go files, installers,
|
||||
packaging, Formulae, or release automation. Protected `main` pushes run the
|
||||
complete native matrix.
|
||||
|
||||
Complete `Multi-profile E2E` is not a PR admission context. It belongs to the
|
||||
`Main Integration — 主干集成` workflow and runs only after a push to `main` (or
|
||||
an explicit manual dispatch). A failing downstream run remains a real
|
||||
regression and must be repaired, but it must not be represented by a synthetic
|
||||
successful PR check.
|
||||
|
||||
```mermaid
|
||||
flowchart TB
|
||||
PR["Pull request"] --> ADMISSION["CI"]
|
||||
ADMISSION --> L["Lint"]
|
||||
ADMISSION --> T["Test"]
|
||||
ADMISSION --> C["Coverage"]
|
||||
ADMISSION --> P["Policy"]
|
||||
ADMISSION --> E["Edition"]
|
||||
ADMISSION --> I["Interface Integrity"]
|
||||
ADMISSION --> A["AI Behavior"]
|
||||
ADMISSION --> S["CLI Smoke"]
|
||||
ADMISSION --> M["Mock MCP"]
|
||||
ADMISSION --> MAIN["Protected main"]
|
||||
MAIN --> NATIVE["Full native platform matrix"]
|
||||
MAIN --> E2E["Multi-profile E2E"]
|
||||
MAIN --> RELEASE["Release delivery"]
|
||||
```
|
||||
|
||||
## Running focused gates locally
|
||||
|
||||
Run the contracts relevant to the change:
|
||||
|
||||
```sh
|
||||
make build
|
||||
make policy
|
||||
make interface-integrity
|
||||
make authoritative-interface-integrity BASE_REF=<merge-base>
|
||||
make schema-compatibility BASE_REF=<merge-base>
|
||||
make skill-command-integrity
|
||||
make cli-smoke
|
||||
# Run on the corresponding native runner with its generated profile:
|
||||
make coverage-gate-platform BASE_REF=<merge-base> PROFILE=<coverage-profile>
|
||||
make mock-mcp-smoke
|
||||
go test -v -count=1 ./pkg/editiontest/...
|
||||
```
|
||||
|
||||
`make coverage-gate` is the enforcement step, not a profile generator. It
|
||||
expects the candidate, policy, shortcut, and merge-base profiles
|
||||
(`coverage.txt`, `coverage-policy.txt`, `coverage-shortcut.txt`, and
|
||||
`coverage-base.txt`) produced by the parallel CI profile jobs. A clean local
|
||||
checkout can reproduce the Linux/overall CI gate sequentially with:
|
||||
For an exact CHANGELOG-only branch:
|
||||
|
||||
```sh
|
||||
base_ref=$(git merge-base HEAD origin/main)
|
||||
root=$(pwd)
|
||||
base_worktree=$(mktemp -d "${TMPDIR:-/tmp}/dws-coverage-base.XXXXXX")
|
||||
rmdir "$base_worktree"
|
||||
cleanup() { git worktree remove --force "$base_worktree" >/dev/null 2>&1 || true; }
|
||||
trap cleanup EXIT HUP INT TERM
|
||||
|
||||
go test -count=1 -p 1 -coverprofile=coverage.txt -covermode=atomic \
|
||||
./ ./cmd/... ./internal/... ./skills/...
|
||||
go test -count=1 -coverprofile=coverage-policy.txt -covermode=atomic \
|
||||
./pkg/... ./scripts/policy/...
|
||||
go test -count=1 \
|
||||
-run '^(TestAllShortcuts|TestCrossPlatformCoverage)' \
|
||||
-coverpkg=./internal/app,./internal/helpers,./internal/shortcut/... \
|
||||
-coverprofile=coverage-shortcut.txt \
|
||||
-covermode=atomic \
|
||||
./internal/app ./internal/helpers ./internal/shortcut/...
|
||||
git worktree add --detach "$base_worktree" "$base_ref"
|
||||
(
|
||||
cd "$base_worktree"
|
||||
go test -count=1 -p 1 \
|
||||
-coverprofile="$root/coverage-base.txt" -covermode=atomic \
|
||||
./ ./cmd/... ./internal/... ./skills/...
|
||||
)
|
||||
COVERAGE_ADDITIONAL_PROFILE=coverage-shortcut.txt \
|
||||
make coverage-gate BASE_REF="$base_ref"
|
||||
./scripts/policy/check-changelog-pr.sh --fast-path "$base_ref" HEAD
|
||||
```
|
||||
|
||||
The native-platform target likewise expects `PROFILE` to have already been
|
||||
generated on that operating system. CI owns those generation steps; copying
|
||||
only either enforcement command into a clean checkout is intentionally an
|
||||
incomplete invocation.
|
||||
`make coverage-gate` is an enforcement step, not a profile generator. CI
|
||||
generates the candidate, supporting, merge-base, and (when risk-selected)
|
||||
native profiles before the aggregate `Coverage` context evaluates them. The
|
||||
aggregate and native gates require 100% coverage for changed executable Go
|
||||
statements. Overall coverage remains an unrounded, zero-tolerance merge-base
|
||||
non-regression check. Candidate and baseline profiles are evaluated by the
|
||||
same block-deduplicating checker; supporting policy and shortcut profiles
|
||||
contribute to changed-code coverage only. The checked-in badge is presentation
|
||||
only and is never read as a gate input.
|
||||
|
||||
CI derives the authoritative Interface snapshots from both the PR merge-base
|
||||
and the latest reachable stable release tag. The complete Schema snapshot comes
|
||||
from the PR merge-base, which contains the registry-first Schema introduced on
|
||||
`main`. The candidate branch cannot bless a breaking change by editing a
|
||||
fixture. Schema additions are allowed; historical products, tools, parameters,
|
||||
parameter mappings, positional execution fields, constraints, and safety
|
||||
semantics remain protected. Positional descriptions are documentation and may
|
||||
change without breaking compatibility.
|
||||
|
||||
`make update-interface-baseline` still extends the local checked-in Interface
|
||||
fixture used by `make interface-integrity`. Updates are monotonic: they add new
|
||||
commands and flags without removing history.
|
||||
|
||||
For an intentional compatibility reset at a major-version boundary, run
|
||||
`make reset-interface-baseline`. This replaces all CLI compatibility history
|
||||
with the current command tree and must receive explicit human review.
|
||||
Compatibility checks derive authoritative Interface snapshots from the PR
|
||||
merge-base and the latest reachable stable release. The candidate cannot bless
|
||||
a breaking change by editing a fixture. Schema additions are allowed;
|
||||
historical products, tools, parameters, mappings, positional execution fields,
|
||||
constraints, and safety semantics remain protected.
|
||||
|
||||
## Required GitHub repository settings
|
||||
|
||||
Create a ruleset for `main` that requires pull requests and code-owner review,
|
||||
then mark these aggregate status checks as required:
|
||||
The `main` quality ruleset must enable strict required-status-check policy
|
||||
(`strict_required_status_checks_policy=true`) so a PR is revalidated whenever
|
||||
`main` advances. It must require these exact contexts and no legacy aliases:
|
||||
|
||||
- `CI Gate`
|
||||
- `Multi Profile E2E`
|
||||
- `AI Behavior Check`
|
||||
- `Lint`
|
||||
- `Test`
|
||||
- `Coverage`
|
||||
- `Policy`
|
||||
- `Edition`
|
||||
- `Interface Integrity`
|
||||
- `AI Behavior`
|
||||
- `CLI Smoke`
|
||||
- `Mock MCP`
|
||||
|
||||
`CI Gate` fails closed unless every first-layer CI job succeeds, including
|
||||
lint, tests, native Linux/Windows/macOS coverage, policy,
|
||||
Interface/Schema/Skill integrity, and smoke tests. Requiring the aggregate
|
||||
check keeps repository rules stable when an internal job is renamed or split.
|
||||
|
||||
The `ai-generated` label must be applied by the PR-creation automation or by a
|
||||
maintainer; GitHub cannot infer reliably whether a human-authored PR contains
|
||||
AI-generated code.
|
||||
Do not require helper jobs, `Multi-profile E2E`, or an aggregate admission
|
||||
alias. Update ruleset contexts only after the new names have appeared on the
|
||||
protected branch, so a rename cannot silently remove enforcement or leave an
|
||||
unproducible required context.
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
# Event consume — AI subprocess contract
|
||||
|
||||
Aligns `dws event consume` with the "AI subprocess contract" that
|
||||
`lark-cli event consume` exposes, so any orchestrator (Claude Code's
|
||||
Monitor, a bash bridge, systemd, an agent plugin) can drive it with zero
|
||||
ambiguity: know when it is ready, stop it cleanly, and machine-read why it
|
||||
exited.
|
||||
Defines the stable `dws event consume` subprocess contract so an
|
||||
orchestrator can determine when the consumer is ready, stop it cleanly,
|
||||
and machine-read why it exited.
|
||||
|
||||
Scope of this branch: the four **contract** items below. Reconnect
|
||||
resilience (keeping the stream alive across a transient upstream drop) is
|
||||
@@ -16,7 +14,6 @@ tracked separately and intentionally out of scope here.
|
||||
- `--duration D` — wall-clock budget (exit 0). Kept as `--duration`, NOT
|
||||
aliased to `--timeout`: the global `--timeout` is the HTTP request
|
||||
timeout (int seconds) and would collide (different type and meaning).
|
||||
Docs note the lark-cli name difference.
|
||||
- Bus idle-shutdown fires only with **zero** consumers, so a connected
|
||||
consumer is never idle-killed.
|
||||
- SIGINT/SIGTERM already cancel the run context and return cleanly.
|
||||
@@ -28,7 +25,7 @@ tracked separately and intentionally out of scope here.
|
||||
On connect, emit a fixed stderr line **before** any stdout event:
|
||||
|
||||
```
|
||||
[event] ready event_key=<key> bus_pid=<pid>
|
||||
[event] ready event_key=<key> bus_pid=<pid> subscribe_id=<id>
|
||||
```
|
||||
|
||||
Parents block on stderr until this line, then read stdout. Suppressed
|
||||
@@ -75,7 +72,7 @@ or runtime failure (permissions, network, params) = non-zero, with no
|
||||
|
||||
### 4. Cleanup on exit (no `kill -9`)
|
||||
|
||||
Ownership-based, matching lark-cli:
|
||||
Ownership-based cleanup:
|
||||
- If this run **created** the subscription (no `--subscribe-id`), a clean
|
||||
exit (SIGTERM / SIGINT / stdin-EOF / limit / timeout) **unsubscribes**
|
||||
it server-side and sends Bye.
|
||||
|
||||
+102
-20
@@ -1,10 +1,54 @@
|
||||
# 发布手册(预发 / 正式)
|
||||
|
||||
发布只走一条链路:本地脚本负责封板、验证并推送 annotated tag;GitHub Actions 负责构建和发布最终产物。不要直接运行 `goreleaser release`,也不要手工补打或移动 tag。
|
||||
发布只走一条受控链路:GitHub Actions 的 `Release` workflow 负责版本分配、封板、构建、签名和下游发布;Homebrew 以 workflow 自动创建的 Formula PR 经独立审核合入为交付边界。本地 `dws-release` 仍是兼容入口,但不再要求某一台固定电脑承担打包;不要直接运行 `goreleaser release`,也不要手工补打、移动或复用 tag。
|
||||
|
||||
发布前必须完成平台治理:目标 GitHub 仓库已启用 immutable releases,`main` 要求 `CI Gate`,操作机已安装并登录 `gh`。本地脚本会在封 tag 前通过 API 检查 immutable releases、当前 SHA 的 `CI Gate` 和在途 Release;`v*` tag ruleset 仍需仓库管理员预先配置并由操作人确认。
|
||||
发布前必须完成平台治理:目标 GitHub 仓库已启用 immutable releases,`main` 精确要求 `CI` workflow 的九个 context:`Lint`、`Test`、`Coverage`、`Policy`、`Edition`、`Interface Integrity`、`AI Behavior`、`CLI Smoke`、`Mock MCP`。云端和本地入口都会在封 tag 前检查 immutable releases、当前 SHA 的全部九个 context 和在途 Release;`v*` tag ruleset 仍需仓库管理员预先配置。
|
||||
|
||||
## 日常只用一个入口
|
||||
## 推荐入口:GitHub 云端发布
|
||||
|
||||
任何具有仓库写权限、因而可以手动运行 Actions workflow 的成员,都可以基于当时最新的 `main` 发起发布:
|
||||
|
||||
1. 在 GitHub Actions 打开 `Release`,选择 `Run workflow`,分支必须是默认分支 `main`。
|
||||
2. `release_operation=plan`,选择 `release_channel=beta|stable`;仅在开始新 beta 线时选择 `release_bump=patch|minor|major`。
|
||||
3. workflow summary 会给出唯一的下一版本。把对应的精确 `CHANGELOG.md` 章节通过 PR 合入 `main`。
|
||||
4. 再次运行,改为 `release_operation=publish`,并输入 `PUBLISH beta` 或 `PUBLISH stable`。
|
||||
|
||||
`plan` 是纯只读操作,不创建 tag、预留版本号或生成包。CHANGELOG 合入期间若另一个发布先占用了该版本,`publish` 会重新分配并因 CHANGELOG 章节不匹配而拒绝,需要重新 plan。`publish` 会先再次确认 dispatch SHA 仍是当前 `main`、Code Admission 和平台治理均通过,再由唯一的 write job 使用 GitHub API 原子创建 annotated tag;同一次 run 随即进入既有的跨平台构建、GitHub/npm、可选 OSS/Gitee 发布和 Homebrew PR DAG。内置 `GITHUB_TOKEN` 创建的 tag 不依赖第二条 workflow 被再次触发。
|
||||
|
||||
OSS 镜像默认不参与发布 DAG,适用于尚未创建 Bucket 的仓库。云端封板会把当时的仓库变量 `ENABLE_OSS_MIRROR=true` 记录为不可变 tag 元数据 `OSS-Mirror: enabled`,否则记录为 `deferred`;后续发布和撤回只读取该 sealed policy,不读取变量的当前值。`enabled` 继续对缺失凭据、无效 Bucket、上传、pointer 和撤回失败保持 fail-closed;`deferred` 明确跳过不存在的渠道。为避免补发后撤回遗漏,deferred 版本暂不接受 `repair_oss_version`,启用 OSS 只影响后续新 tag,直到补齐可审计的不可变 repair 证明。
|
||||
|
||||
## 自动版本规则
|
||||
|
||||
- beta:如果存在尚未封正式版的最高版本线,自动取 `beta.N+1`;否则从最新已分配正式版按所选 patch/minor/major 开新线并取 `beta.1`。
|
||||
- stable:先锁定最高开放版本线上的最新已分配 beta,再要求它已成功交付且未撤回;不会跳过失败/撤回的最新 beta 去选择更早版本。正式版 core 与该 beta 完全相同。
|
||||
- `vX.Y.Z`、`vX.Y.Z-beta.N` 一经分配就永久占用。撤回时创建 `withdrawn/v...` 墓碑,原编号永不复用。
|
||||
- 例如撤回 `v1.0.53-beta.5` 后,下一 beta 是 `v1.0.53-beta.6`;撤回正式版 `v1.0.53` 后,下一 patch 修复线是 `v1.0.54-beta.1`,验证后再发布 `v1.0.54`。
|
||||
- 如果最新 beta 已撤回,禁止直接用更早 beta 晋级正式版;必须先构建下一个 beta。
|
||||
|
||||
## 全平台撤回与回滚
|
||||
|
||||
已公开版本出现问题时,在 GitHub Actions 运行 `Withdraw release`,分支必须选择当前默认分支 `main`,并填写:
|
||||
|
||||
- `version`:精确版本,例如 `v1.0.53` 或 `v1.0.53-beta.5`。
|
||||
- `reason`:8–300 字符的单行公开原因。
|
||||
- `confirmation`:精确输入 `WITHDRAW <version>`,例如 `WITHDRAW v1.0.53`。
|
||||
|
||||
该 workflow 使用与发布相同的串行 publication lock,并进入受保护的 `release-withdrawal` environment。它只接受已经由 Release workflow 完整交付的 public immutable release,自动选择同一渠道中最新的、更早且未撤回的完整版本作为回退目标,然后按以下顺序执行:
|
||||
|
||||
1. 先创建永久 annotated tag `withdrawn/<version>`,记录原 tag object、commit、原因、申请人和 workflow run。这个墓碑是版本号永久占用记录,永不移动、永不删除。
|
||||
2. 先验证 Homebrew Formula;若它仍指向问题版本,先创建回退 PR,再继续其他渠道撤回。这样 PR 创建失败时只留下可安全续跑的墓碑,不会先造成渠道分裂。若 Formula 尚未指向问题版本或已经处于安全版本,则直接校验。
|
||||
3. GitHub Release 先标记为 withdrawn;npm 精确版本执行 `deprecate`,并把 `latest` / `beta` dist-tag 回退;只有目标 tag 封存了 `OSS-Mirror: enabled` 时,OSS 才会先补齐回退版本资产,再移动 `latest.txt` / `beta.txt` 并删除问题版本目录;启用 Gitee 时同样先补齐回退 Release,再删除问题 Release 和 tag。
|
||||
4. npm 以及目标 tag 启用或发布时配置的镜像渠道均已验证安全后,删除 GitHub 上的问题 Release 和原 `v...` tag,并验证 `/releases/latest` 对正式版回到安全版本。若本次创建了 Homebrew PR,run 最后故意保持失败,直到另一名维护者审核合入;合入后,从新的 `main` 使用完全相同的 version、reason 和 confirmation 重跑并完成。永久 `withdrawn/v...` 墓碑始终保留。
|
||||
|
||||
GitHub、npm、OSS、Gitee 和 Homebrew 的“回滚”指新的安装、升级和渠道解析不再拿到问题版本。已经装到用户电脑上的二进制无法被服务端强制降级;用户必须重新安装回退版本、安装后续修复版,或使用 CLI 自带的本地 rollback 能力。npm 不执行 `unpublish`:问题版本保留明确的弃用警告,但 `latest` / `beta` 不再指向它;即使 registry 允许删除,已发布过的版本号也不会重新使用。
|
||||
|
||||
撤回前必须存在同一渠道中更早、完整交付且未撤回的安全版本;若目标是该渠道第一个版本、没有安全候选,workflow 会在创建墓碑或修改任何渠道前 fail closed,需要先决定明确的替代策略。CLI 本地 rollback 也只有在本机仍保留上一次升级备份时可用。
|
||||
|
||||
撤回以“精确版本”为单位,不会因为正式版曾由某个 beta 晋级就隐式级联修改另一个渠道。若同一缺陷同时存在于正式版及其 beta,应先撤回正式版,再撤回对应 beta,并分别使用各自的精确确认串;每次都只会把该渠道回退到自己的安全候选。
|
||||
|
||||
撤回正式版 `v1.0.53` 后,`v1.0.53` 仍被墓碑视为已分配。下一次 patch 发布从 `v1.0.54-beta.1` 开始,验证后晋级 `v1.0.54`。撤回 `v1.0.53-beta.5` 后,同一开放版本线继续为 `v1.0.53-beta.6`;不会退回或复用 `beta.5`。
|
||||
|
||||
## 兼容入口:本地发布
|
||||
|
||||
安装发布 Skill 后直接运行:
|
||||
|
||||
@@ -25,11 +69,11 @@ dws-release config --remote origin
|
||||
```text
|
||||
main 上的候选代码 + beta CHANGELOG
|
||||
→ vX.Y.Z-beta.N(预发验证)
|
||||
→ 只允许补正式 CHANGELOG,源码不得再变化
|
||||
→ vX.Y.Z(正式发布)
|
||||
→ 补正式 CHANGELOG;允许继续通过 PR 合入新 commit
|
||||
→ vX.Y.Z(正式发布,封板提交必须包含该 beta 提交)
|
||||
```
|
||||
|
||||
正式版必须显式指定本次验证过的 beta。脚本会比较两者:除 `CHANGELOG.md` 外只要有任何文件变化,就拒绝正式发布。这样预发测过的代码、命令树和正式发布的代码是同一份。
|
||||
云端入口自动选择本次最新、已交付且未撤回的 beta;本地入口必须显式指定。流水线要求该 beta 已成功交付、未撤回,且 beta 提交必须位于正式发布封板提交的历史中——不能跳过 beta 直接发正式版,但允许在 beta 之后把经过 review 合入 `main` 的 commit 一起发布。
|
||||
|
||||
## 预发发布
|
||||
|
||||
@@ -45,13 +89,13 @@ dws-release v1.2.3-beta.1
|
||||
dws-release v1.2.3-beta.1
|
||||
```
|
||||
|
||||
预检包含测试、策略检查、旧正式版命令树兼容检查、全平台打包、npm 安装验证,以及 macOS 环境下的 Homebrew 安装验证。通过后发布:
|
||||
预检包含测试、策略检查、旧正式版命令树兼容检查、全平台打包、npm 安装验证,以及 macOS 环境下的 Homebrew 安装验证。它还会从默认分支触发一次无发布权限的 `Release governance preflight`,用正式流水线相同的身份检查该精确 commit 的九个 Code Admission context 和 immutable releases。通过后会在当前 Git worktree 的私有 Git 状态目录写入一个有效期六小时的证明,绑定版本、精确 commit、发布仓库、beta/stable 基线和远端 `main`:
|
||||
|
||||
```bash
|
||||
dws-release v1.2.3-beta.1 --publish
|
||||
```
|
||||
|
||||
命令会在所有预检完成后要求再次输入完整版本号。统一入口不提供跳过确认的参数。
|
||||
若源码、版本、远端身份和 stable 基线均未变化,`--publish` 会复用该证明,只执行远端契约、发布身份和最终治理复核,不再重复测试与打包。也可以直接运行 `--publish`;没有可复用证明时只会完整执行一次预检。命令在封 tag 前仍要求再次输入完整版本号,统一入口不提供跳过确认的参数。
|
||||
|
||||
## 正式发布
|
||||
|
||||
@@ -68,7 +112,7 @@ dws-release v1.2.3 --from-beta v1.2.3-beta.1
|
||||
dws-release v1.2.3 --from-beta v1.2.3-beta.1 --publish
|
||||
```
|
||||
|
||||
`FROM_BETA` 不会自动推断,并会写入 stable annotated tag 的 `From-Beta` 元数据,CI 会再次读取和验证。
|
||||
本地入口的 `FROM_BETA` 不会自动推断;云端入口会按上述规则唯一选择。两种入口都会把它写入 stable annotated tag 的 `From-Beta` 元数据,CI 会再次读取和验证。
|
||||
|
||||
## CHANGELOG 契约
|
||||
|
||||
@@ -86,28 +130,66 @@ dws-release v1.2.3 --from-beta v1.2.3-beta.1 --publish
|
||||
|
||||
## CI/CD 保证
|
||||
|
||||
- 只接受 `vX.Y.Z-beta.N` 和 `vX.Y.Z`,且新版本必须高于上一正式版。这里的“上一正式版”必须同时具备公开非草稿 GitHub Release 和同 tag/commit 的成功 Release workflow;只有 tag、没有交付成功的孤儿版本会阻断后续发布,要求先重跑补齐。历史版本若曾通过专用 recovery workflow 完成交付,只能使用仓库内 `delivered-stable-recoveries.json` 中精确到 tag、commit、run、workflow SHA 与 attempt 的 reviewed 证据;验证仍要求 release、Darwin 签名和最终发布三个 job 全部成功,不能接受任意 workflow_dispatch。
|
||||
- tag 必须是 annotated tag;本地脚本在推送前重新确认 HEAD 与远端 `main` 完全一致,CI 允许其后 `main` 前进,但要求封板提交仍位于 `main` 历史中。
|
||||
- 只接受 `vX.Y.Z-beta.N` 和 `vX.Y.Z`,且新版本必须高于上一正式版。这里的“上一正式版”必须同时具备公开非草稿 GitHub Release 和同 tag/commit 的成功 Release workflow;只有 tag、没有交付成功的孤儿版本会阻断后续发布,要求走受保护恢复补齐。云端 tag 会固定 `Release-Run`、requester、commit 和版本分配指纹,交付验证按该精确 run/attempt 及完整 job graph 取证,不接受任意 `workflow_dispatch`。历史版本若曾通过专用 recovery workflow 完成交付,只能使用仓库内 `delivered-stable-recoveries.json` 中精确到 tag、commit、run、workflow SHA 与 attempt 的 reviewed 证据。
|
||||
- tag 必须是 annotated tag;本地脚本要求封板提交已通过 PR 合入并包含在远端 `main` 历史中,发布只推送 tag。CI 允许其后 `main` 继续前进,但始终要求封板提交位于 `main` 历史中。
|
||||
- 日常 CI 和发布前都会对比“最新已交付正式版”的完整命令树;若长时间预检期间该 baseline 发生变化,会针对新的 baseline 重新比较。
|
||||
- GoReleaser 只构建;Darwin 重签、checksums 重算和 npm 安装验证通过后,才统一上传 GitHub Release 的最终产物。
|
||||
- 六个平台归档会逐个解包并核验二进制内嵌版本;公开资产集合、checksums 集合和 npm tarball integrity 都必须精确一致。npm tarball 固定由 npm `10.9.2` 打包,避免重跑时因 runner 自带 npm 漂移产生不同字节。
|
||||
- stable 发布到 npm `latest`,更新 OSS `latest.txt` 和共享安装脚本;prerelease 发布到 npm `beta`,只更新 OSS `beta.txt`,不会覆盖稳定入口。
|
||||
- Release workflow 使用一个最多容纳 100 个 pending run 的串行 publication queue;本地入口仍要求上一条 Release 完成后才能封下一个 tag。
|
||||
- 本地 tag push 失败时会删除本次新建的本地 tag。tag 一旦成功推送,后续发布归 CI 所有,禁止改 tag 指向或复用版本号。
|
||||
- stable 发布到 npm `latest`;prerelease 发布到 npm `beta`。启用 `ENABLE_OSS_MIRROR=true` 后,stable 同步 OSS `latest.txt` 和共享安装脚本,prerelease 只同步 OSS `beta.txt`,不会覆盖稳定入口。
|
||||
- Release workflow 使用一个最多容纳 100 个 pending run 的串行 publication queue;版本规划、云端封板、发布、恢复、修复和撤回共享同一发布锁。
|
||||
- 本地 tag push 失败时会删除本次新建的本地 tag。远端 tag 一旦创建,后续发布归 CI 所有;发布中途失败时走受保护恢复,禁止改 tag 指向或复用版本号。只有已经公开版本经过受保护的全渠道撤回并留下永久 `withdrawn/...` 墓碑后,撤回 workflow 才会在最后一步删除原 tag。
|
||||
|
||||
npm 补发只允许从默认分支触发 Release workflow 的 `repair_npm_version`。它只支持启用 immutable releases 后、由本流水线成功产出的公开 immutable release:目标必须是 `main` 历史中的 annotated tag,并且同 commit 的 `Build immutable GitHub Release` job 已成功。即使后续 npm 分发失败,这个独立的产物封存边界仍可作为补发依据。补发会用目标 commit 的 npm 模板重组包,逐平台核验资产和二进制版本,再发布到隔离的 `backfill` dist-tag,不会回滚 `latest` / `beta`。历史 mutable release 不进入自动补发路径,避免把可被替换的资产带入 npm。
|
||||
|
||||
OSS/Gitee 分发失败时直接重跑该 tag 的 `Publish npm and mirrors` failed job;各步会复用 immutable GitHub 资产并保持 channel 单调。独立 Gitee release workflow 和本地直发脚本已停用,避免绕开 publication queue 或用重新构建的不同字节覆盖镜像。
|
||||
已启用的 OSS 或 Gitee 分发失败且 GitHub immutable Release、npm 已交付时,从受保护的默认分支触发
|
||||
Release workflow,并且只填写 `repair_oss_version` 或 `repair_gitee_version` 之一。channel
|
||||
repair 会精确绑定失败 tag run 的最新 attempt,且 OSS repair 要求 tag 的 sealed policy 为 `enabled`;contract、构建、Developer ID 签名、
|
||||
immutable GitHub 发布和 npm delivery 必须全部成功,且只能有一个 OSS/Gitee 下游失败,
|
||||
随后才会下载并重新校验原始资产、修复所选镜像。OSS repair 必须匹配失败的 OSS step;
|
||||
Gitee repair 还允许其 job 因该 OSS 失败而 skipped,此时只代表 Gitee backfill 成功,
|
||||
不会把仍未修复的 OSS 标成成功。该证据不能用于 beta → stable 或
|
||||
stable baseline,后两者仍要求整条 Release 成功或受保护 recovery 成功。不要重跑旧
|
||||
attempt 的单个 failed job,以免在 attempts 之间拼接交付证据。独立 Gitee release
|
||||
workflow 和本地直发脚本已停用,避免绕开 publication queue 或用重新构建的不同字节覆盖镜像。
|
||||
|
||||
OSS 的 `latest.txt` / `beta.txt` 当前是镜像频道元数据;仓库内安装器仍从 GitHub/Gitee 解析版本,不能把 OSS pointer 当成已接入的安装通道。
|
||||
## 既有 tag 的紧急恢复
|
||||
|
||||
Homebrew 当前只属于本机预检/手工公式通道:预检会在当前 macOS 架构真实安装,但 Release workflow 不发布 tap,CI 生成的单主机公式也不应当作 Darwin 双架构正式交付。正式自动交付范围是 GitHub Release、npm、OSS,以及显式开启时的 Gitee fallback;Homebrew 双架构 tap 发布需另立需求。
|
||||
云端封板或本地 tag push 已成功、但 Release workflow 失败且 GitHub Release 尚未公开时,不要新建临时 workflow、移动 tag 或跳过门禁。在最新且干净的 `main` worktree 运行:
|
||||
|
||||
```bash
|
||||
dws-release recover v1.2.3-beta.1
|
||||
```
|
||||
|
||||
命令会自动解析 annotated tag object、peeled commit,以及 tag 绑定的失败云端 run 或最近一次匹配的失败 tag-push run;也可以用 `--failed-run <run-id>` 精确指定。确认完整版本号后,它从默认分支触发受保护的恢复模式并等待完成。恢复模式必须满足:
|
||||
|
||||
- 输入精确绑定原 annotated tag object、commit 和失败的 sealed `Release` run;云端 run 还必须与 tag 内的 run ID、attempt、requester 完全一致,commit 必须仍在 `main` 历史中。
|
||||
- 目标只允许不存在 GitHub Release 或仍为 Draft;已经公开的版本不能全量重建:单个下游故障走对应的 channel repair,版本本身有问题则走受保护的全平台 withdrawal。
|
||||
- `release-recovery` environment 必须限制为受保护分支、配置至少一名 required reviewer,并禁止自审;workflow 会通过 API 复核这些设置,未配置时 fail closed。
|
||||
- 恢复复用正常的 contract、构建、Developer ID 签名、资产校验、immutable 发布、Homebrew、npm,以及已启用的 OSS jobs,不存在 recovery 专用 publisher 或门禁跳过。
|
||||
- 如果 GitHub Release 已在 recovery 中封存、后续 Homebrew/npm 校验发生瞬时失败,只重跑该 run 的 failed jobs;流水线仅在隐藏 run marker、tag object、commit 和 finalized artifact 字节全部精确一致时复用公开 Release。
|
||||
|
||||
成功的默认分支恢复 run 会成为后续 beta → stable 和 stable baseline 验证的可审计交付证据;历史临时分支恢复仍只接受 reviewed manifest 中的固定证据。
|
||||
|
||||
云端 seal 后不要使用 GitHub 的 “Re-run failed jobs” 作为交付修复:annotated tag 永久绑定最初的 run attempt,普通 rerun 不会成为可接受的交付证据。GitHub Release 尚未公开时走上述 protected recovery;已经公开且仅 npm/OSS/Gitee 某一渠道失败时走对应 repair;版本内容本身有问题时走 withdrawal。
|
||||
|
||||
OSS 的 `latest.txt` / `beta.txt` 是镜像频道元数据;当前仓库安装器仍主要从 GitHub/Gitee 解析版本。启用 OSS 后,发布和撤回把它作为受控分发渠道处理,保证一旦外部消费者接入该 pointer,也不会继续解析到已撤回版本;未启用时两条流程都明确跳过不存在的 OSS 渠道。
|
||||
|
||||
Release workflow 会生成 Darwin/Linux 双架构 Formula,并分别为 stable/beta 打开 Homebrew PR;tap 的默认分支仍以独立审核合入为交付边界。撤回 workflow 使用相同模板和回退版本 checksums 打开反向 PR;问题 GitHub Release 会先被移除以阻止新安装,永久墓碑和 workflow 日志承担审计/续跑依据。
|
||||
|
||||
## 平台治理前置
|
||||
|
||||
仓库管理员还需要在 GitHub 平台配置两项不可由脚本替代的规则:
|
||||
仓库管理员还需要在 GitHub 平台配置以下不可由脚本替代的规则:
|
||||
|
||||
- `main` 必须要求精确的 `CI Gate`;tag workflow 也会通过 Checks API 再确认该封板 SHA 已通过。
|
||||
- `main` 必须精确要求 `Lint`、`Test`、`Coverage`、`Policy`、`Edition`、`Interface Integrity`、`AI Behavior`、`CLI Smoke`、`Mock MCP` 九个 Code Admission context;tag workflow 也会通过 Checks API 再确认该封板 SHA 上九项全部成功。
|
||||
- 必须启用 immutable releases;它只保护启用后发布的 release,因此应在第一次使用新流水线前配置。为 `v*` 增加 tag ruleset,限制创建权限,并在 release 发布前保护 tag 的短暂窗口。
|
||||
- tag ruleset 还必须覆盖 `withdrawn/v*`:只允许受保护的撤回 workflow 创建墓碑,禁止更新或删除墓碑;同时应允许 Release workflow 创建新的 `v*`,允许撤回 workflow 在全部渠道回退后删除精确的问题 `v*`。若组织级规则阻止这两个 workflow 的预期动作,发布或撤回会 fail closed,不能靠手工移动 tag 绕过。
|
||||
- 配置 `RELEASE_GOVERNANCE_TOKEN` Actions secret,只授予目标仓库 `Administration: read`;内置 `GITHUB_TOKEN` 不具备 immutable-releases API 所需的仓库治理权限。每次本地预检和 tag workflow 都使用这一个身份进行 fail-closed 验证。
|
||||
- 配置 `APPLE_CERTIFICATE_P12_BASE64`、`APPLE_CERTIFICATE_PASSWORD` 和具备发布权限的 `NPM_TOKEN`;撤回还要求该 npm 身份能够执行 `deprecate` 和修改 dist-tag。
|
||||
- 启用 OSS 镜像时,先创建有效 Bucket,再设置仓库变量 `ENABLE_OSS_MIRROR=true`,并配置 `OSS_ACCESS_KEY_ID`、`OSS_ACCESS_KEY_SECRET`、`OSS_ENDPOINT`、`OSS_BUCKET`,按需配置 `OSS_PREFIX`。启用后发布保持 fail-closed;撤回身份必须能够补齐安全版本资产、写 `latest.txt` / `beta.txt` 并删除问题版本前缀。尚未 provision Bucket 时保持该变量未设置或不等于 `true`,新 tag 会封存 `OSS-Mirror: deferred` 并跳过 OSS;该版本不能通过现有 repair 流程事后改成启用。
|
||||
- 若启用 Gitee fallback,设置 `ENABLE_GITEE_UPLOAD_FALLBACK=true`,并配置 `GITEE_TOKEN`、`GITEE_USER`、`GITEE_REPO`;该身份必须能够创建和删除目标仓库的 Release 与 tag。
|
||||
- 单独配置 `HOMEBREW_PR_TOKEN`,优先使用仅授权本仓库且具备 `Contents: write`、`Pull requests: write` 的 fine-grained PAT;若组织策略不允许该账号使用 fine-grained PAT,则回退到仅带 `public_repo` scope 的专用 classic PAT。治理预检和 tag contract 会验证 token 身份、classic scope,并用 `[skip ci]` 临时分支和 draft PR 完成真实写权限 canary,随后立即关闭 PR、删除分支;任何清理失败都会 fail closed。门禁也会拒绝与治理 token 复用。
|
||||
- 创建 `release-recovery` environment,只允许受保护分支,设置 required reviewer、禁止自审并关闭管理员绕过。workflow 会读取 environment 的 required-reviewer、prevent-self-review 和 protected-branch 规则;规则缺失时紧急恢复会失败,正常 beta/stable tag 发布不受影响。
|
||||
- 创建 `release-withdrawal` environment,只允许受保护分支,设置至少一名 required reviewer、禁止申请人自审并关闭管理员绕过。撤回 workflow 会通过 API 复核这些规则;任何一项缺失都会在触碰 npm、OSS、Gitee、Homebrew 或 GitHub Release 前失败。
|
||||
- 仓库或组织的 Actions 策略必须允许 `Release` 与 `Withdraw release` workflow 的 `GITHUB_TOKEN` 获得各 job 声明的 `contents: write`。若上述发布凭证采用 environment secret,确认 `release-withdrawal` 审批完成后能够读取撤回所需的 npm、OSS、Gitee 和 Homebrew 凭证。
|
||||
|
||||
immutable releases 或 `CI Gate` 缺失时,发布脚本会自动拒绝封 tag。tag ruleset 可能来自组织层,脚本不自动推断其最终作用范围;管理员确认不能省略,脚本约定也不能替代平台强制。
|
||||
immutable releases,或任一 Code Admission context 缺失、未成功时,发布脚本会自动拒绝封 tag。tag ruleset 可能来自组织层,脚本不自动推断其最终作用范围;管理员确认不能省略,脚本约定也不能替代平台强制。
|
||||
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,223 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
type tokenManagerSnapshotProvider struct {
|
||||
load func() (*authpkg.TokenData, error)
|
||||
}
|
||||
|
||||
func (p tokenManagerSnapshotProvider) GetAccessToken(context.Context) (string, error) {
|
||||
data, err := p.load()
|
||||
if err != nil || data == nil {
|
||||
return "", err
|
||||
}
|
||||
return data.AccessToken, nil
|
||||
}
|
||||
|
||||
func (p tokenManagerSnapshotProvider) GetTokenSnapshot(context.Context) (*authpkg.TokenData, error) {
|
||||
return p.load()
|
||||
}
|
||||
|
||||
type tokenManagerLegacyGetter struct {
|
||||
token string
|
||||
err error
|
||||
}
|
||||
|
||||
func (g tokenManagerLegacyGetter) GetToken() (string, string, error) {
|
||||
return g.token, "file", g.err
|
||||
}
|
||||
|
||||
func installTokenManagerFakes(t *testing.T, load func() (*authpkg.TokenData, error)) {
|
||||
t.Helper()
|
||||
oldProvider, oldLegacy := newAccessTokenProvider, newLegacyTokenManager
|
||||
oldEdition := edition.Get()
|
||||
edition.Override(&edition.Hooks{})
|
||||
newAccessTokenProvider = func(string) accessTokenGetter {
|
||||
return tokenManagerSnapshotProvider{load: load}
|
||||
}
|
||||
newLegacyTokenManager = func(string) legacyTokenGetter {
|
||||
return tokenManagerLegacyGetter{err: authpkg.ErrTokenDataNotFound}
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
newAccessTokenProvider, newLegacyTokenManager = oldProvider, oldLegacy
|
||||
edition.Override(oldEdition)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerCachesUntilMarkerRevisionChanges(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
token := "token-a"
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: token, ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
|
||||
manager := NewTokenManager()
|
||||
first, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || first.AccessToken != "token-a" {
|
||||
t.Fatalf("first token = %#v, %v", first, err)
|
||||
}
|
||||
second, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || second.AccessToken != "token-a" || calls.Load() != 1 {
|
||||
t.Fatalf("cached token = %#v, %v, calls=%d", second, err, calls.Load())
|
||||
}
|
||||
|
||||
token = "token-b"
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rotated, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || rotated.AccessToken != "token-b" || calls.Load() != 2 {
|
||||
t.Fatalf("rotated token = %#v, %v, calls=%d", rotated, err, calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerDoesNotCacheWithoutExpiryOrRevision(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token"}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerTreatsMalformedMarkerAsUncacheable(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(configDir, "token.json"), []byte("{"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if snapshot, err := manager.Get(context.Background(), configDir, ""); err != nil || snapshot.AccessToken != "token" {
|
||||
t.Fatalf("snapshot = %#v, error = %v", snapshot, err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerDoesNotCacheOpaqueEditionStorageWithProviderFallback(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
edition.Override(&edition.Hooks{
|
||||
LoadToken: func(string) ([]byte, error) { return nil, nil },
|
||||
TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
|
||||
return fallback()
|
||||
},
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerCoalescesConcurrentLoads(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
release := make(chan struct{})
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
<-release
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
const workers = 8
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
errs := make(chan error, workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_, err := manager.Get(context.Background(), configDir, "")
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
for calls.Load() == 0 {
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
close(release)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 1 {
|
||||
t.Fatalf("provider calls = %d, want 1", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerPreservesProviderFailure(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
want := errors.New("keychain permission denied")
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) { return nil, want })
|
||||
_, err := NewTokenManager().Get(context.Background(), configDir, "")
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("provider failure was misclassified as missing credentials: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenResolutionErrorOnlyClassifiesTrueMissingCredential(t *testing.T) {
|
||||
missing := tokenResolutionError(authpkg.ErrTokenDataNotFound)
|
||||
var typed interface{ Unwrap() error }
|
||||
if !errors.As(missing, &typed) || !errors.Is(missing, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("missing error = %v", missing)
|
||||
}
|
||||
want := errors.New("decrypt failed")
|
||||
if got := tokenResolutionError(want); !errors.Is(got, want) || errors.Is(got, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("storage error = %v", got)
|
||||
}
|
||||
if got := tokenResolutionError(context.Canceled); !errors.Is(got, context.Canceled) {
|
||||
t.Fatalf("cancellation = %v", got)
|
||||
}
|
||||
}
|
||||
@@ -21,19 +21,61 @@ import (
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
const accessTokenRefreshWindow = 5 * time.Minute
|
||||
|
||||
type legacyTokenGetter interface {
|
||||
GetToken() (string, string, error)
|
||||
}
|
||||
|
||||
type accessTokenSnapshotGetter interface {
|
||||
GetTokenSnapshot(context.Context) (*authpkg.TokenData, error)
|
||||
}
|
||||
|
||||
// AccessTokenSnapshot is the minimal bearer view needed by the process cache.
|
||||
// Refresh-token material never leaves the auth package.
|
||||
type AccessTokenSnapshot struct {
|
||||
AccessToken string
|
||||
ExpiresAt time.Time
|
||||
Source string
|
||||
}
|
||||
|
||||
type tokenManagerKey struct {
|
||||
configDir string
|
||||
profile string
|
||||
}
|
||||
|
||||
type tokenManagerEntry struct {
|
||||
mu sync.Mutex
|
||||
snapshot AccessTokenSnapshot
|
||||
revision string
|
||||
}
|
||||
|
||||
// TokenManager is the only process cache for user access tokens. Cache entries
|
||||
// are isolated by config directory and profile, expiry-aware, and invalidated
|
||||
// by the credential publication marker written by auth storage.
|
||||
type TokenManager struct {
|
||||
mu sync.Mutex
|
||||
entries map[tokenManagerKey]*tokenManagerEntry
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func NewTokenManager() *TokenManager {
|
||||
return &TokenManager{entries: make(map[tokenManagerKey]*tokenManagerEntry), now: time.Now}
|
||||
}
|
||||
|
||||
var runtimeTokenManager = NewTokenManager()
|
||||
|
||||
var (
|
||||
newAccessTokenProvider = func(configDir string) accessTokenGetter {
|
||||
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, disc)
|
||||
discard := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, discard)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
return provider
|
||||
}
|
||||
@@ -44,55 +86,218 @@ var (
|
||||
}
|
||||
)
|
||||
|
||||
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
|
||||
// the same host compatibility hooks as MCP. It mirrors the former body of
|
||||
// getCachedRuntimeToken (excluding process-level cache and timing).
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
provider := newAccessTokenProvider(configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token), nil
|
||||
}
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
return "", tokenErr
|
||||
}
|
||||
manager := newLegacyTokenManager(configDir)
|
||||
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
|
||||
return strings.TrimSpace(leg), nil
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
|
||||
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
|
||||
// the active edition config directory, the same process-cached path as MCP is used.
|
||||
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
if t := strings.TrimSpace(explicitToken); t != "" {
|
||||
return t, nil
|
||||
// Get resolves an access token for the active runtime profile.
|
||||
func (m *TokenManager) Get(ctx context.Context, configDir, explicitToken string) (AccessTokenSnapshot, error) {
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return AccessTokenSnapshot{AccessToken: token, Source: "explicit"}, nil
|
||||
}
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
|
||||
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
|
||||
return tok, nil
|
||||
key := tokenManagerKey{
|
||||
configDir: canonicalTokenConfigDir(configDir),
|
||||
profile: strings.TrimSpace(authpkg.RuntimeProfile()),
|
||||
}
|
||||
entry := m.entry(key)
|
||||
entry.mu.Lock()
|
||||
defer entry.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if m != nil && m.now != nil {
|
||||
now = m.now()
|
||||
}
|
||||
revision, present, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if tokenSnapshotUsable(entry.snapshot, now) && present && revision != "" && revision == entry.revision {
|
||||
return entry.snapshot, nil
|
||||
}
|
||||
|
||||
// Treat the marker and credential as one optimistic snapshot. A concurrent
|
||||
// login/refresh between the reads causes a retry instead of caching stale A
|
||||
// under the publication marker for B.
|
||||
for attempt := 0; attempt < 4; attempt++ {
|
||||
beforeRevision, beforePresent, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
snapshot, err := resolveTokenSnapshotWithEdition(ctx, configDir, key.profile)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
afterRevision, afterPresent, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if beforePresent != afterPresent || beforeRevision != afterRevision {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(snapshot.AccessToken) == "" {
|
||||
return AccessTokenSnapshot{}, noCredentialsError()
|
||||
}
|
||||
if tokenSnapshotUsable(snapshot, now) && afterPresent && afterRevision != "" {
|
||||
entry.snapshot = snapshot
|
||||
entry.revision = afterRevision
|
||||
} else {
|
||||
entry.snapshot = AccessTokenSnapshot{}
|
||||
entry.revision = ""
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
tok, err := resolveAccessTokenFromDir(ctx, configDir)
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("token publication changed repeatedly while resolving credentials")
|
||||
}
|
||||
|
||||
func (m *TokenManager) entry(key tokenManagerKey) *tokenManagerEntry {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.entries == nil {
|
||||
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
|
||||
}
|
||||
entry := m.entries[key]
|
||||
if entry == nil {
|
||||
entry = &tokenManagerEntry{}
|
||||
m.entries[key] = entry
|
||||
}
|
||||
return entry
|
||||
}
|
||||
|
||||
func (m *TokenManager) Invalidate() {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
func resolveTokenSnapshotWithEdition(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
|
||||
hooks := edition.Get()
|
||||
opaqueStorage := hooks.LoadToken != nil || hooks.SaveToken != nil || hooks.DeleteToken != nil
|
||||
provider := hooks.TokenProvider
|
||||
if provider == nil {
|
||||
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
// Opaque edition storage hooks have no publication-revision contract.
|
||||
// Resolve them on every logical request instead of caching a token that
|
||||
// may be replaced outside the default auth store.
|
||||
if opaqueStorage {
|
||||
snapshot.ExpiresAt = time.Time{}
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
var fallbackSnapshot AccessTokenSnapshot
|
||||
var fallbackCalled bool
|
||||
token, err := provider(ctx, func() (string, error) {
|
||||
fallbackCalled = true
|
||||
var fallbackErr error
|
||||
fallbackSnapshot, fallbackErr = resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
|
||||
if fallbackErr != nil {
|
||||
return "", fallbackErr
|
||||
}
|
||||
return fallbackSnapshot.AccessToken, nil
|
||||
})
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("edition token provider: %w", err)
|
||||
}
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return AccessTokenSnapshot{}, noCredentialsError()
|
||||
}
|
||||
if fallbackCalled && token == fallbackSnapshot.AccessToken {
|
||||
if opaqueStorage {
|
||||
fallbackSnapshot.ExpiresAt = time.Time{}
|
||||
}
|
||||
return fallbackSnapshot, nil
|
||||
}
|
||||
// Edition providers expose no lifetime metadata, so resolve them on every
|
||||
// logical request instead of recreating a process-lifetime string cache.
|
||||
return AccessTokenSnapshot{AccessToken: token, Source: "edition"}, nil
|
||||
}
|
||||
|
||||
func resolveAccessTokenSnapshotFromDir(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
|
||||
provider := newAccessTokenProvider(configDir)
|
||||
if snapshotProvider, ok := provider.(accessTokenSnapshotGetter); ok {
|
||||
data, err := snapshotProvider.GetTokenSnapshot(ctx)
|
||||
if err == nil && data != nil && strings.TrimSpace(data.AccessToken) != "" {
|
||||
return AccessTokenSnapshot{
|
||||
AccessToken: strings.TrimSpace(data.AccessToken),
|
||||
ExpiresAt: data.ExpiresAt,
|
||||
Source: "oauth",
|
||||
}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
return resolveLegacyToken(configDir, err)
|
||||
}
|
||||
|
||||
token, err := provider.GetAccessToken(ctx)
|
||||
if err == nil && strings.TrimSpace(token) != "" {
|
||||
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: "oauth_compat"}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
return resolveLegacyToken(configDir, err)
|
||||
}
|
||||
|
||||
func resolveLegacyToken(configDir string, oauthErr error) (AccessTokenSnapshot, error) {
|
||||
token, source, err := newLegacyTokenManager(configDir).GetToken()
|
||||
if err == nil && strings.TrimSpace(token) != "" {
|
||||
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: source}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if oauthErr != nil {
|
||||
return AccessTokenSnapshot{}, oauthErr
|
||||
}
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, authpkg.RuntimeProfile())
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if tok != "" {
|
||||
return tok, nil
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves every non-runner bearer token through
|
||||
// the same TokenManager used by MCP tool calls.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
snapshot, err := runtimeTokenManager.Get(ctx, configDir, explicitToken)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
func tokenSnapshotUsable(snapshot AccessTokenSnapshot, now time.Time) bool {
|
||||
return strings.TrimSpace(snapshot.AccessToken) != "" &&
|
||||
!snapshot.ExpiresAt.IsZero() &&
|
||||
now.Before(snapshot.ExpiresAt.Add(-accessTokenRefreshWindow))
|
||||
}
|
||||
|
||||
func canonicalTokenConfigDir(configDir string) string {
|
||||
if absolute, err := filepath.Abs(configDir); err == nil {
|
||||
return filepath.Clean(absolute)
|
||||
}
|
||||
return filepath.Clean(configDir)
|
||||
}
|
||||
|
||||
func noCredentialsError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return fmt.Errorf("认证信息已失效,请重新认证")
|
||||
return fmt.Errorf("认证信息已失效,请重新认证: %w", authpkg.ErrTokenDataNotFound)
|
||||
}
|
||||
return fmt.Errorf("no credentials found, run: dws auth login")
|
||||
return fmt.Errorf("no credentials found, run: dws auth login: %w", authpkg.ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,15 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
func TestResolveAuxiliaryAccessToken_explicitToken(t *testing.T) {
|
||||
@@ -24,3 +32,48 @@ func TestResolveAuxiliaryAccessToken_emptyConfigDir(t *testing.T) {
|
||||
t.Fatal("expected error for empty config directory")
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveAccessTokenFromDirPreservesRefreshFailure(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
configDir := filepath.Join(root, "config")
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
t.Setenv(keychain.StorageDirEnv, filepath.Join(root, "keychain"))
|
||||
|
||||
if err := authpkg.SaveTokenData(configDir, &authpkg.TokenData{
|
||||
AccessToken: "expired-access",
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(-time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
CorpID: "corp_refresh",
|
||||
UserID: "user_refresh",
|
||||
ClientID: "client_refresh",
|
||||
Source: "mcp",
|
||||
}); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
|
||||
originalTransport := http.DefaultTransport
|
||||
t.Cleanup(func() {
|
||||
http.DefaultTransport = originalTransport
|
||||
})
|
||||
http.DefaultTransport = refreshFailureRoundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return nil, errors.New("refresh endpoint rejected token")
|
||||
})
|
||||
|
||||
token, err := resolveAccessTokenFromDir(context.Background(), configDir)
|
||||
if token != "" {
|
||||
t.Fatalf("token = %q, want empty", token)
|
||||
}
|
||||
if err == nil {
|
||||
t.Fatal("resolveAccessTokenFromDir() error = nil")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "refresh endpoint rejected token") {
|
||||
t.Fatalf("error = %q, want original refresh failure", err)
|
||||
}
|
||||
}
|
||||
|
||||
type refreshFailureRoundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f refreshFailureRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
@@ -34,6 +34,10 @@ func (g fakeAccessTokenGetter) GetAccessToken(context.Context) (string, error) {
|
||||
return g.token, g.err
|
||||
}
|
||||
|
||||
func (g fakeAccessTokenGetter) ForceRefreshRejectedToken(context.Context, string) (string, error) {
|
||||
return g.token, g.err
|
||||
}
|
||||
|
||||
type fakeLegacyTokenGetter struct {
|
||||
token string
|
||||
err error
|
||||
@@ -212,11 +216,19 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
|
||||
if _, err := resolveAccessTokenFromDir(context.Background(), "unused"); !errors.Is(err, authpkg.ErrTokenDecryption) {
|
||||
t.Fatalf("decryption error = %v", err)
|
||||
}
|
||||
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{err: errors.New("missing")} }
|
||||
newAccessTokenProvider = func(string) accessTokenGetter {
|
||||
return fakeAccessTokenGetter{err: authpkg.ErrTokenDataNotFound}
|
||||
}
|
||||
newLegacyTokenManager = func(string) legacyTokenGetter { return fakeLegacyTokenGetter{token: " legacy "} }
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); err != nil || got != "legacy" {
|
||||
t.Fatalf("legacy token = %q, %v", got, err)
|
||||
}
|
||||
authpkg.SetRuntimeProfile("corp:user")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("explicit profile fallback = token %q error %v, want profile error", got, err)
|
||||
}
|
||||
authpkg.SetRuntimeProfile("")
|
||||
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{err: errors.New("load")} }
|
||||
newLegacyTokenManager = func(string) legacyTokenGetter { return fakeLegacyTokenGetter{err: errors.New("missing")} }
|
||||
edition.Override(&edition.Hooks{})
|
||||
@@ -240,10 +252,10 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T) {
|
||||
oldMark, oldFactory := markAccessTokenStale, newRefreshProvider
|
||||
oldLoad, oldFactory := loadRefreshTokenData, newRefreshProvider
|
||||
oldStop := stopStdio
|
||||
t.Cleanup(func() {
|
||||
markAccessTokenStale, newRefreshProvider = oldMark, oldFactory
|
||||
loadRefreshTokenData, newRefreshProvider = oldLoad, oldFactory
|
||||
stopStdio = oldStop
|
||||
stdioMu.Lock()
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
@@ -251,11 +263,13 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
|
||||
})
|
||||
fail := errors.New("failure")
|
||||
_ = oldFactory(t.TempDir())
|
||||
markAccessTokenStale = func(string) error { return fail }
|
||||
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) { return nil, fail }
|
||||
if _, err := ForceRefreshAccessToken(context.Background(), "config"); !errors.Is(err, fail) {
|
||||
t.Fatalf("mark stale error = %v", err)
|
||||
t.Fatalf("load rejected token error = %v", err)
|
||||
}
|
||||
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{AccessToken: "rejected"}, nil
|
||||
}
|
||||
markAccessTokenStale = func(string) error { return nil }
|
||||
for _, tc := range []struct {
|
||||
getter fakeAccessTokenGetter
|
||||
want string
|
||||
@@ -264,7 +278,7 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
|
||||
{getter: fakeAccessTokenGetter{token: " "}, want: "empty"},
|
||||
{getter: fakeAccessTokenGetter{token: " refreshed "}},
|
||||
} {
|
||||
newRefreshProvider = func(string) accessTokenGetter { return tc.getter }
|
||||
newRefreshProvider = func(string) rejectedAccessTokenRefresher { return tc.getter }
|
||||
got, err := ForceRefreshAccessToken(context.Background(), "config")
|
||||
if tc.want != "" && (err == nil || !strings.Contains(err.Error(), tc.want)) {
|
||||
t.Fatalf("refresh error = %v, want %q", err, tc.want)
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/audit"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/output"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
@@ -117,9 +118,29 @@ func newAuditVerifyCommand() *cobra.Command {
|
||||
}
|
||||
}
|
||||
|
||||
valid, brokenAt, err := auditVerify(target)
|
||||
if err != nil {
|
||||
return fmt.Errorf("校验失败: %w", err)
|
||||
valid, brokenAt, verifyErr := auditVerify(target)
|
||||
if output.ResolveFormat(cmd, output.FormatTable) == output.FormatJSON {
|
||||
if verifyErr != nil && brokenAt == 0 {
|
||||
return fmt.Errorf("校验失败: %w", verifyErr)
|
||||
}
|
||||
payload := map[string]any{
|
||||
"valid": valid,
|
||||
"file": target,
|
||||
"brokenAt": brokenAt,
|
||||
}
|
||||
if verifyErr != nil {
|
||||
payload["reason"] = verifyErr.Error()
|
||||
}
|
||||
if err := output.WriteCommandPayload(cmd, payload, output.FormatTable); err != nil {
|
||||
return err
|
||||
}
|
||||
if !valid {
|
||||
auditExit(1)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if verifyErr != nil {
|
||||
return fmt.Errorf("校验失败: %w", verifyErr)
|
||||
}
|
||||
if valid {
|
||||
fmt.Printf("✓ %s 哈希链完整(全部通过)\n", filepath.Base(target))
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
// 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.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func executeAuditVerifyJSON(t *testing.T, verify func(string) (bool, int, error)) (map[string]any, int, error) {
|
||||
t.Helper()
|
||||
previousVerify, previousExit := auditVerify, auditExit
|
||||
t.Cleanup(func() {
|
||||
auditVerify, auditExit = previousVerify, previousExit
|
||||
})
|
||||
auditVerify = verify
|
||||
exitCode := 0
|
||||
auditExit = func(code int) { exitCode = code }
|
||||
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.SilenceErrors = true
|
||||
root.SilenceUsage = true
|
||||
root.PersistentFlags().String("format", "json", "output format")
|
||||
root.AddCommand(newAuditVerifyCommand())
|
||||
var stdout bytes.Buffer
|
||||
root.SetOut(&stdout)
|
||||
root.SetArgs([]string{"verify", "--file", "/tmp/audit.jsonl"})
|
||||
err := root.Execute()
|
||||
|
||||
var payload map[string]any
|
||||
if decodeErr := json.Unmarshal(stdout.Bytes(), &payload); decodeErr != nil {
|
||||
t.Fatalf("audit verify stdout must be one JSON document: %v\n%s", decodeErr, stdout.String())
|
||||
}
|
||||
return payload, exitCode, err
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAuditVerifyJSONOutputIsSingleDocument(t *testing.T) {
|
||||
payload, exitCode, err := executeAuditVerifyJSON(t, func(string) (bool, int, error) {
|
||||
return true, 0, nil
|
||||
})
|
||||
if err != nil || exitCode != 0 {
|
||||
t.Fatalf("audit verify returned err=%v exit=%d", err, exitCode)
|
||||
}
|
||||
if payload["valid"] != true || payload["file"] != "/tmp/audit.jsonl" || payload["brokenAt"] != float64(0) {
|
||||
t.Fatalf("unexpected audit payload: %#v", payload)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageAuditVerifyBrokenJSONIncludesReasonBeforeExit(t *testing.T) {
|
||||
payload, exitCode, err := executeAuditVerifyJSON(t, func(string) (bool, int, error) {
|
||||
return false, 3, errors.New("prev_hash mismatch")
|
||||
})
|
||||
if err != nil || exitCode != 1 {
|
||||
t.Fatalf("broken audit verify returned err=%v exit=%d", err, exitCode)
|
||||
}
|
||||
if payload["valid"] != false || payload["brokenAt"] != float64(3) || payload["reason"] != "prev_hash mismatch" {
|
||||
t.Fatalf("unexpected broken audit payload: %#v", payload)
|
||||
}
|
||||
}
|
||||
+168
-34
@@ -30,6 +30,7 @@ import (
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/helpers"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/pat"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
@@ -85,7 +86,7 @@ func buildAuthCommand(patCaller edition.ToolCaller) *cobra.Command {
|
||||
newAuthMigrateKeychainCommand(),
|
||||
newAuthExportCommand(),
|
||||
newAuthImportCommand(),
|
||||
newAuthExchangeCommand(),
|
||||
newAuthExchangeCommand(patCaller),
|
||||
newAuthResetCommand(),
|
||||
)
|
||||
return cmd
|
||||
@@ -146,6 +147,9 @@ func newAuthLoginCommand(patCaller edition.ToolCaller) *cobra.Command {
|
||||
provider := authpkg.NewDeviceFlowProvider(configDir, nil)
|
||||
provider.Output = cmd.ErrOrStderr()
|
||||
provider.NoBrowser, _ = cmd.Flags().GetBool("no-browser")
|
||||
provider.IdentityEnricher = func(ctx context.Context, data *authpkg.TokenData) error {
|
||||
return enrichAuthLoginProfileFromContact(ctx, configDir, patCaller, data)
|
||||
}
|
||||
tokenData, err = authDeviceLogin(provider, loginCtx)
|
||||
if err != nil {
|
||||
return apperrors.NewAuth(fmt.Sprintf("device authorization failed: %v", err))
|
||||
@@ -158,6 +162,9 @@ func newAuthLoginCommand(patCaller edition.ToolCaller) *cobra.Command {
|
||||
provider.Output = cmd.ErrOrStderr()
|
||||
provider.NoBrowser, _ = cmd.Flags().GetBool("no-browser")
|
||||
provider.TargetCorpID = cfg.TargetCorpID
|
||||
provider.IdentityEnricher = func(ctx context.Context, data *authpkg.TokenData) error {
|
||||
return enrichAuthLoginProfileFromContact(ctx, configDir, patCaller, data)
|
||||
}
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
tokenData, err = authOAuthLogin(provider, loginCtx, authLoginForcesAuthorization(cfg))
|
||||
if err != nil {
|
||||
@@ -167,17 +174,13 @@ func newAuthLoginCommand(patCaller edition.ToolCaller) *cobra.Command {
|
||||
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
if tokenData != nil && strings.TrimSpace(tokenData.CorpID) != "" {
|
||||
_ = enrichAuthLoginProfileFromContact(cmd.Context(), configDir, patCaller, tokenData)
|
||||
ResetRuntimeTokenCache()
|
||||
clearCompatCache()
|
||||
}
|
||||
|
||||
w := cmd.OutOrStdout()
|
||||
runPostLoginAuthorization := func() error {
|
||||
if !recommendAuthMode {
|
||||
return nil
|
||||
}
|
||||
restoreProfile := replaceRuntimeProfile(authpkg.TokenProfileSelector(tokenData))
|
||||
defer restoreProfile()
|
||||
recommendScopeMode := pat.LoginRecommendScopeRecommended
|
||||
var initialPlan *pat.LoginRecommendPlan
|
||||
if postLoginTUIMode {
|
||||
@@ -303,7 +306,10 @@ var (
|
||||
authRunLoginRecommend = pat.RunLoginRecommendAuthorizationWithOptions
|
||||
authRunDirectPATWait = runDirectPATAuthCheckWaitOnly
|
||||
authResolveProfile = authpkg.ResolveProfile
|
||||
authResolveProfileDeletion = authpkg.ResolveProfileDeletionScope
|
||||
authRevokeToken = authpkg.RevokeTokenRemote
|
||||
authRevokeTokenForData = authpkg.RevokeTokenRemoteForData
|
||||
authLoadTokenForProfile = authpkg.LoadTokenDataForProfile
|
||||
authDeleteProfileToken = authpkg.DeleteTokenDataForProfile
|
||||
authEnsureProfilesMigration = authpkg.EnsureProfilesMigration
|
||||
authLoadProfiles = authpkg.LoadProfiles
|
||||
@@ -423,9 +429,10 @@ func newAuthLogoutCommand() *cobra.Command {
|
||||
Short: "清除认证信息(默认退出所有组织)",
|
||||
Long: `清除本机钉钉登录态。
|
||||
|
||||
默认退出所有已登录组织 profile;指定 --profile 时只退出该组织,不影响其他组织。`,
|
||||
默认退出全部账号。--profile 传组织时退出该组织全部账号;传精确账号或本地 profile 名时只退出一个账号。`,
|
||||
Example: ` dws auth logout
|
||||
dws auth logout --profile <corpId>
|
||||
dws auth logout --profile <corpId>:<userId>
|
||||
dws auth logout --profile "钉钉"`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -455,7 +462,7 @@ func newAuthLogoutCommand() *cobra.Command {
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("profile", "", "指定要退出的 profile 名或 corpId")
|
||||
cmd.Flags().String("profile", "", "指定组织或账号:corpId、corpName、corpId:userId、corpId:userName、corpName:userId、corpName:userName 或本地 profile 名")
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -468,7 +475,8 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
指定 --profile 时只读取并刷新被选中的 token slot,不会修改 currentProfile。`,
|
||||
Example: ` dws auth status
|
||||
dws auth status --profile <corpId>
|
||||
dws auth status --profile "钉钉"
|
||||
dws auth status --profile <corpId>:<userId>
|
||||
dws auth status --profile "钉钉:孙博文"
|
||||
dws auth status --profile <corpId> --format json`,
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -477,6 +485,17 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
if err != nil {
|
||||
return apperrors.NewInternal("failed to read --profile")
|
||||
}
|
||||
profileSelector = strings.TrimSpace(profileSelector)
|
||||
if profileSelector != "" {
|
||||
selected, resolveErr := authpkg.ResolveProfile(configDir, profileSelector)
|
||||
if resolveErr != nil {
|
||||
return apperrors.NewValidation(resolveErr.Error())
|
||||
}
|
||||
if selected == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("profile %q not found", profileSelector))
|
||||
}
|
||||
profileSelector = authpkg.ProfileSelector(*selected)
|
||||
}
|
||||
restoreProfile := pushRuntimeProfile(profileSelector)
|
||||
defer restoreProfile()
|
||||
|
||||
@@ -484,6 +503,7 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
refreshed := false
|
||||
var tokenData *authpkg.TokenData
|
||||
var statusErr error
|
||||
var refreshFailure error
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
if data, err := authOAuthStatus(provider); err == nil {
|
||||
@@ -498,18 +518,23 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
refreshed = true
|
||||
}
|
||||
} else if edition.Get().AutoPurgeToken {
|
||||
refreshFailure = refreshErr
|
||||
_ = authDeleteTokenData(configDir)
|
||||
} else if tokenData != nil {
|
||||
_ = authMarkProfileStatus(configDir, tokenData.CorpID, authpkg.ProfileStatusExpired)
|
||||
refreshFailure = refreshErr
|
||||
_ = authMarkProfileStatus(configDir, authpkg.TokenProfileSelector(tokenData), authpkg.ProfileStatusExpired)
|
||||
}
|
||||
}
|
||||
if authStatusAuthenticated(tokenData) {
|
||||
if refreshFailure == nil && authStatusAuthenticated(tokenData) {
|
||||
authenticated = true
|
||||
}
|
||||
} else {
|
||||
statusErr = err
|
||||
}
|
||||
diagnostic := authStatusDiagnosticFromError(statusErr)
|
||||
if refreshFailure != nil {
|
||||
diagnostic = authStatusRefreshDiagnostic(refreshFailure)
|
||||
}
|
||||
|
||||
// Check if JSON output is requested
|
||||
format, _ := cmd.Root().PersistentFlags().GetString("format")
|
||||
@@ -554,7 +579,7 @@ func newAuthStatusCommand() *cobra.Command {
|
||||
return nil
|
||||
},
|
||||
}
|
||||
cmd.Flags().String("profile", "", "指定要查看的 profile 名或 corpId")
|
||||
cmd.Flags().String("profile", "", "指定组织或账号:corpId、corpName、corpId:userId、corpId:userName、corpName:userId、corpName:userName 或本地 profile 名")
|
||||
return cmd
|
||||
}
|
||||
|
||||
@@ -618,13 +643,33 @@ func newAuthMigrateKeychainCommand() *cobra.Command {
|
||||
}
|
||||
|
||||
func logoutOneProfile(_ *cobra.Command, ctx context.Context, configDir, selector string) error {
|
||||
if _, err := authResolveProfile(configDir, selector); err != nil {
|
||||
selected, exact, err := authResolveProfileDeletion(configDir, selector)
|
||||
if err != nil {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
restoreProfile := pushRuntimeProfile(selector)
|
||||
defer restoreProfile()
|
||||
_ = authRevokeToken(ctx)
|
||||
if err := authDeleteProfileToken(configDir, selector); err != nil {
|
||||
if selected == nil {
|
||||
return apperrors.NewValidation(fmt.Sprintf("profile %q not found", selector))
|
||||
}
|
||||
stableSelector := selected.CorpID
|
||||
if exact {
|
||||
stableSelector = authpkg.ProfileSelector(*selected)
|
||||
if data, loadErr := authLoadTokenForProfile(configDir, stableSelector); loadErr == nil {
|
||||
_ = authRevokeTokenForData(ctx, data)
|
||||
}
|
||||
} else if cfg, loadErr := authLoadProfiles(configDir); loadErr == nil {
|
||||
for _, profile := range cfg.Profiles {
|
||||
if profile.CorpID != selected.CorpID {
|
||||
continue
|
||||
}
|
||||
if data, tokenErr := authLoadTokenForProfile(configDir, authpkg.ProfileSelector(profile)); tokenErr == nil {
|
||||
_ = authRevokeTokenForData(ctx, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := authDeleteProfileToken(configDir, stableSelector); err != nil {
|
||||
if strings.Contains(err.Error(), "not found") {
|
||||
return apperrors.NewValidation(err.Error())
|
||||
}
|
||||
return apperrors.NewInternal(fmt.Sprintf("failed to clear token data: %v", err))
|
||||
}
|
||||
return nil
|
||||
@@ -642,9 +687,9 @@ func logoutAllProfiles(_ *cobra.Command, ctx context.Context, configDir string)
|
||||
_ = authRevokeToken(ctx)
|
||||
} else {
|
||||
for _, profile := range cfg.Profiles {
|
||||
restoreProfile := pushRuntimeProfile(profile.CorpID)
|
||||
_ = authRevokeToken(ctx)
|
||||
restoreProfile()
|
||||
if data, tokenErr := authLoadTokenForProfile(configDir, authpkg.ProfileSelector(profile)); tokenErr == nil {
|
||||
_ = authRevokeTokenForData(ctx, data)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := authDeleteAllTokenData(configDir); err != nil {
|
||||
@@ -665,6 +710,14 @@ func pushRuntimeProfile(selector string) func() {
|
||||
}
|
||||
}
|
||||
|
||||
func replaceRuntimeProfile(selector string) func() {
|
||||
previous := authpkg.RuntimeProfile()
|
||||
authpkg.SetRuntimeProfile(strings.TrimSpace(selector))
|
||||
return func() {
|
||||
authpkg.SetRuntimeProfile(previous)
|
||||
}
|
||||
}
|
||||
|
||||
func newAuthExportCommand() *cobra.Command {
|
||||
return newAuthExportCommandWithSupport(authpkg.PortableExportSupportError)
|
||||
}
|
||||
@@ -810,7 +863,7 @@ func newAuthImportCommandWithSupport(supportError func() error) *cobra.Command {
|
||||
return cmd
|
||||
}
|
||||
|
||||
func newAuthExchangeCommand() *cobra.Command {
|
||||
func newAuthExchangeCommand(caller edition.ToolCaller) *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "exchange",
|
||||
Short: "Exchange an authorization code for credentials",
|
||||
@@ -832,6 +885,9 @@ func newAuthExchangeCommand() *cobra.Command {
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
provider := authpkg.NewOAuthProvider(configDir, nil)
|
||||
provider.IdentityEnricher = func(ctx context.Context, data *authpkg.TokenData) error {
|
||||
return enrichAuthLoginProfileFromContact(ctx, configDir, caller, data)
|
||||
}
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
exchangeCtx, cancel := context.WithTimeout(cmd.Context(), time.Minute)
|
||||
defer cancel()
|
||||
@@ -1154,6 +1210,19 @@ func resolveAuthLoginConfig(cmd *cobra.Command) (authLoginConfig, error) {
|
||||
if err != nil {
|
||||
return authLoginConfig{}, err
|
||||
}
|
||||
flow := "oauth"
|
||||
if strings.TrimSpace(token) != "" {
|
||||
flow = "token"
|
||||
} else if device {
|
||||
flow = "device"
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.request",
|
||||
"flow", flow,
|
||||
"profile_selector", strings.TrimSpace(profileSelector),
|
||||
"target_corp_id", targetCorpID,
|
||||
"recommend", recommend,
|
||||
)
|
||||
return authLoginConfig{
|
||||
Token: strings.TrimSpace(token),
|
||||
Force: force,
|
||||
@@ -1189,7 +1258,11 @@ type contactProfileIdentity struct {
|
||||
UserName string
|
||||
}
|
||||
|
||||
func enrichAuthLoginProfileFromContact(ctx context.Context, configDir string, caller edition.ToolCaller, data *authpkg.TokenData) error {
|
||||
type tokenOverrideToolCaller interface {
|
||||
CallToolWithToken(ctx context.Context, token, productID, toolName string, args map[string]any) (*edition.ToolResult, error)
|
||||
}
|
||||
|
||||
func enrichAuthLoginProfileFromContact(ctx context.Context, _ string, caller edition.ToolCaller, data *authpkg.TokenData) error {
|
||||
if caller == nil || data == nil {
|
||||
return nil
|
||||
}
|
||||
@@ -1197,24 +1270,62 @@ func enrichAuthLoginProfileFromContact(ctx context.Context, configDir string, ca
|
||||
if corpID == "" {
|
||||
return nil
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.lookup.start",
|
||||
"corp_id", corpID,
|
||||
"user_id", strings.TrimSpace(data.UserID),
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"corp_name", strings.TrimSpace(data.CorpName),
|
||||
)
|
||||
if strings.TrimSpace(data.CorpName) != "" && strings.TrimSpace(data.UserID) != "" && strings.TrimSpace(data.UserName) != "" {
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.lookup.result",
|
||||
"source", "token_exchange",
|
||||
"corp_id", corpID,
|
||||
"user_id", strings.TrimSpace(data.UserID),
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"corp_name", strings.TrimSpace(data.CorpName),
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
restoreProfile := pushRuntimeProfile(corpID)
|
||||
defer restoreProfile()
|
||||
ResetRuntimeTokenCache()
|
||||
|
||||
result, err := caller.CallTool(ctx, "contact", "get_current_user_profile", map[string]any{
|
||||
"profile": corpID,
|
||||
})
|
||||
var (
|
||||
result *edition.ToolResult
|
||||
err error
|
||||
)
|
||||
if tokenCaller, ok := caller.(tokenOverrideToolCaller); ok && strings.TrimSpace(data.AccessToken) != "" {
|
||||
result, err = tokenCaller.CallToolWithToken(ctx, data.AccessToken, "contact", "get_current_user_profile", nil)
|
||||
} else {
|
||||
if strings.TrimSpace(data.UserID) == "" {
|
||||
return fmt.Errorf("login identity lookup requires an in-memory token override")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.lookup.error",
|
||||
"corp_id", corpID,
|
||||
"existing_user_id", strings.TrimSpace(data.UserID),
|
||||
"error", err,
|
||||
)
|
||||
if strings.TrimSpace(data.UserID) != "" {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
identity, ok := contactProfileIdentityFromToolResult(result)
|
||||
if !ok {
|
||||
logging.AuthDebug("auth.login.identity.lookup.empty", "corp_id", corpID)
|
||||
return nil
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.lookup.result",
|
||||
"source", "contact.get_current_user_profile",
|
||||
"corp_id", strings.TrimSpace(identity.CorpID),
|
||||
"user_id", strings.TrimSpace(identity.UserID),
|
||||
"user_name", strings.TrimSpace(identity.UserName),
|
||||
"corp_name", strings.TrimSpace(identity.CorpName),
|
||||
)
|
||||
if identity.CorpID != "" && identity.CorpID != corpID {
|
||||
return fmt.Errorf("contact profile corpId %q does not match login corpId %q", identity.CorpID, corpID)
|
||||
}
|
||||
@@ -1230,12 +1341,23 @@ func enrichAuthLoginProfileFromContact(ctx context.Context, configDir string, ca
|
||||
updated.UserName = identity.UserName
|
||||
}
|
||||
if updated.CorpName == data.CorpName && updated.UserID == data.UserID && updated.UserName == data.UserName {
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.resolved",
|
||||
"corp_id", corpID,
|
||||
"user_id", strings.TrimSpace(data.UserID),
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"changed", false,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
if err := authSaveTokenData(configDir, &updated); err != nil {
|
||||
return err
|
||||
}
|
||||
*data = updated
|
||||
logging.AuthDebug(
|
||||
"auth.login.identity.resolved",
|
||||
"corp_id", strings.TrimSpace(data.CorpID),
|
||||
"user_id", strings.TrimSpace(data.UserID),
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"changed", true,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1262,6 +1384,7 @@ func contactProfileIdentityFromJSON(data []byte) (contactProfileIdentity, bool)
|
||||
OrgName string `json:"orgName"`
|
||||
UserID string `json:"userId"`
|
||||
UserIDLower string `json:"userid"`
|
||||
OrgUserID string `json:"orgUserId"`
|
||||
OrgUserName string `json:"orgUserName"`
|
||||
Name string `json:"name"`
|
||||
} `json:"orgEmployeeModel"`
|
||||
@@ -1277,7 +1400,7 @@ func contactProfileIdentityFromJSON(data []byte) (contactProfileIdentity, bool)
|
||||
identity := contactProfileIdentity{
|
||||
CorpID: strings.TrimSpace(org.CorpID),
|
||||
CorpName: strings.TrimSpace(org.OrgName),
|
||||
UserID: firstNonEmptyString(org.UserID, org.UserIDLower),
|
||||
UserID: firstNonEmptyString(org.UserID, org.UserIDLower, org.OrgUserID),
|
||||
UserName: firstNonEmptyString(org.OrgUserName, org.Name),
|
||||
}
|
||||
return identity, identity.CorpID != "" || identity.CorpName != "" || identity.UserID != "" || identity.UserName != ""
|
||||
@@ -1364,6 +1487,17 @@ func authStatusDiagnosticFromError(err error) *authStatusDiagnostic {
|
||||
}
|
||||
}
|
||||
|
||||
func authStatusRefreshDiagnostic(err error) *authStatusDiagnostic {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
return &authStatusDiagnostic{
|
||||
Reason: "token_refresh_failed",
|
||||
Message: fmt.Sprintf("Token 刷新失败: %v", err),
|
||||
Hint: "请重新运行 dws auth login 完成授权。",
|
||||
}
|
||||
}
|
||||
|
||||
func writeAuthStatusJSON(w io.Writer, authenticated, refreshed bool, data *authpkg.TokenData, diagnostic *authStatusDiagnostic) error {
|
||||
resp := authStatusResponse{
|
||||
Success: true,
|
||||
|
||||
@@ -28,6 +28,9 @@ type authCoverageCaller struct {
|
||||
func (c *authCoverageCaller) CallTool(context.Context, string, string, map[string]any) (*edition.ToolResult, error) {
|
||||
return c.result, c.err
|
||||
}
|
||||
func (c *authCoverageCaller) CallToolWithToken(ctx context.Context, _ string, productID, toolName string, args map[string]any) (*edition.ToolResult, error) {
|
||||
return c.CallTool(ctx, productID, toolName, args)
|
||||
}
|
||||
func (*authCoverageCaller) Format() string { return "json" }
|
||||
func (*authCoverageCaller) DryRun() bool { return false }
|
||||
func (*authCoverageCaller) Fields() string { return "" }
|
||||
@@ -67,12 +70,16 @@ func authCoverageRunLogin(t *testing.T, caller edition.ToolCaller, format string
|
||||
|
||||
func TestCrossPlatformCoverageAuthCoverageFormsParentAndTargets(t *testing.T) {
|
||||
oldEdition := edition.Get()
|
||||
oldClientID := authpkg.ClientID()
|
||||
oldClientSecret := authpkg.ClientSecret()
|
||||
oldRunForm := authRunForm
|
||||
oldPrompt := authLoginManualCredentialsPrompt
|
||||
oldSaveConfig := authSaveAppConfig
|
||||
oldResolve := authResolveProfile
|
||||
t.Cleanup(func() {
|
||||
edition.Override(oldEdition)
|
||||
authpkg.SetClientID(oldClientID)
|
||||
authpkg.SetClientSecret(oldClientSecret)
|
||||
authRunForm = oldRunForm
|
||||
authLoginManualCredentialsPrompt = oldPrompt
|
||||
authSaveAppConfig = oldSaveConfig
|
||||
@@ -352,14 +359,22 @@ func TestCrossPlatformCoverageAuthCoverageContactEnrichment(t *testing.T) {
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", &authCoverageCaller{err: errors.New("call")}, &authpkg.TokenData{CorpID: "ding"}); err == nil {
|
||||
t.Fatal("caller error should propagate")
|
||||
}
|
||||
if err := enrichAuthLoginProfileFromContact(
|
||||
ctx,
|
||||
"cfg",
|
||||
&authCoverageCaller{err: errors.New("call")},
|
||||
&authpkg.TokenData{CorpID: "ding", UserID: "known", AccessToken: "token"},
|
||||
); err != nil {
|
||||
t.Fatalf("optional contact metadata failure with known userId = %v", err)
|
||||
}
|
||||
for _, text := range []string{"", "not-json", `{"result":[]}`, `{"result":[{"orgEmployeeModel":{}}]}`} {
|
||||
caller := &authCoverageCaller{result: &edition.ToolResult{Content: []edition.ContentBlock{{Text: text}}}}
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", caller, &authpkg.TokenData{CorpID: "ding"}); err != nil {
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", caller, &authpkg.TokenData{CorpID: "ding", AccessToken: "token"}); err != nil {
|
||||
t.Fatalf("invalid contact %q: %v", text, err)
|
||||
}
|
||||
}
|
||||
mismatch := &authCoverageCaller{result: &edition.ToolResult{Content: []edition.ContentBlock{{Text: `{"result":[{"orgEmployeeModel":{"corpId":"other"}}]}`}}}}
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", mismatch, &authpkg.TokenData{CorpID: "ding"}); err == nil {
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", mismatch, &authpkg.TokenData{CorpID: "ding", AccessToken: "token"}); err == nil {
|
||||
t.Fatal("corp mismatch should fail")
|
||||
}
|
||||
same := &authCoverageCaller{result: &edition.ToolResult{Content: []edition.ContentBlock{{Text: `{"result":[{"orgEmployeeModel":{"corpId":"ding","orgName":"Corp","userid":"u","name":"User"}}]}`}}}}
|
||||
@@ -371,11 +386,7 @@ func TestCrossPlatformCoverageAuthCoverageContactEnrichment(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
authSaveTokenData = func(string, *authpkg.TokenData) error { return errors.New("save") }
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", same, &authpkg.TokenData{CorpID: "ding"}); err == nil {
|
||||
t.Fatal("enrichment save should fail")
|
||||
}
|
||||
authSaveTokenData = func(string, *authpkg.TokenData) error { return nil }
|
||||
data := &authpkg.TokenData{CorpID: "ding"}
|
||||
data := &authpkg.TokenData{CorpID: "ding", AccessToken: "token"}
|
||||
if err := enrichAuthLoginProfileFromContact(ctx, "cfg", same, data); err != nil || data.CorpName != "Corp" || data.UserID != "u" {
|
||||
t.Fatalf("enriched = %#v, %v", data, err)
|
||||
}
|
||||
@@ -394,6 +405,7 @@ func TestCrossPlatformCoverageAuthCoverageDefaultSeamClosures(t *testing.T) {
|
||||
device := authpkg.NewDeviceFlowProvider(configDir, nil)
|
||||
_, _ = authDeviceLogin(device, ctx)
|
||||
oauth := authpkg.NewOAuthProvider(configDir, nil)
|
||||
oauth.NoBrowser = true
|
||||
_, _ = authOAuthStatus(oauth)
|
||||
_, _ = authOAuthAccessToken(oauth, ctx)
|
||||
_, _ = authOAuthLogin(oauth, ctx, true)
|
||||
@@ -409,7 +421,10 @@ func TestCrossPlatformCoverageAuthCoverageStatusAndLogout(t *testing.T) {
|
||||
oldDelete := authDeleteTokenData
|
||||
oldMark := authMarkProfileStatus
|
||||
oldResolve := authResolveProfile
|
||||
oldResolveDeletion := authResolveProfileDeletion
|
||||
oldRevoke := authRevokeToken
|
||||
oldRevokeForData := authRevokeTokenForData
|
||||
oldLoadTokenForProfile := authLoadTokenForProfile
|
||||
oldDeleteProfile := authDeleteProfileToken
|
||||
oldMigrate := authEnsureProfilesMigration
|
||||
oldLoadProfiles := authLoadProfiles
|
||||
@@ -421,7 +436,10 @@ func TestCrossPlatformCoverageAuthCoverageStatusAndLogout(t *testing.T) {
|
||||
authDeleteTokenData = oldDelete
|
||||
authMarkProfileStatus = oldMark
|
||||
authResolveProfile = oldResolve
|
||||
authResolveProfileDeletion = oldResolveDeletion
|
||||
authRevokeToken = oldRevoke
|
||||
authRevokeTokenForData = oldRevokeForData
|
||||
authLoadTokenForProfile = oldLoadTokenForProfile
|
||||
authDeleteProfileToken = oldDeleteProfile
|
||||
authEnsureProfilesMigration = oldMigrate
|
||||
authLoadProfiles = oldLoadProfiles
|
||||
@@ -510,20 +528,51 @@ func TestCrossPlatformCoverageAuthCoverageStatusAndLogout(t *testing.T) {
|
||||
t.Fatalf("mark-expired = %v, marked=%v", err, marked)
|
||||
}
|
||||
|
||||
authResolveProfile = func(string, string) (*authpkg.Profile, error) { return nil, errors.New("missing") }
|
||||
authResolveProfileDeletion = func(string, string) (*authpkg.Profile, bool, error) { return nil, false, errors.New("missing") }
|
||||
if err := logoutOneProfile(nil, context.Background(), "cfg", "x"); err == nil {
|
||||
t.Fatal("missing profile should fail")
|
||||
}
|
||||
authResolveProfile = func(string, string) (*authpkg.Profile, error) { return &authpkg.Profile{CorpID: "ding"}, nil }
|
||||
authRevokeToken = func(context.Context) error { return errors.New("ignored") }
|
||||
authDeleteProfileToken = func(string, string) error { return errors.New("delete") }
|
||||
authResolveProfileDeletion = func(string, string) (*authpkg.Profile, bool, error) {
|
||||
return &authpkg.Profile{CorpID: "ding", UserID: "user"}, true, nil
|
||||
}
|
||||
authLoadTokenForProfile = func(string, string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{CorpID: "ding", UserID: "user"}, nil
|
||||
}
|
||||
authRevokeTokenForData = func(context.Context, *authpkg.TokenData) error { return errors.New("ignored") }
|
||||
var deletedSelector string
|
||||
authDeleteProfileToken = func(_ string, selector string) error {
|
||||
deletedSelector = selector
|
||||
return errors.New("delete")
|
||||
}
|
||||
if err := logoutOneProfile(nil, context.Background(), "cfg", "x"); err == nil {
|
||||
t.Fatal("profile delete should fail")
|
||||
}
|
||||
authDeleteProfileToken = func(string, string) error { return nil }
|
||||
if deletedSelector != "ding:user" {
|
||||
t.Fatalf("exact deletion selector = %q, want stable identity selector", deletedSelector)
|
||||
}
|
||||
authDeleteProfileToken = func(_ string, selector string) error {
|
||||
deletedSelector = selector
|
||||
return nil
|
||||
}
|
||||
if err := logoutOneProfile(nil, context.Background(), "cfg", "x"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if deletedSelector != "ding:user" {
|
||||
t.Fatalf("exact deletion selector = %q, want stable identity selector", deletedSelector)
|
||||
}
|
||||
|
||||
authResolveProfileDeletion = func(string, string) (*authpkg.Profile, bool, error) {
|
||||
return &authpkg.Profile{CorpID: "ding", UserID: "user"}, false, nil
|
||||
}
|
||||
authLoadProfiles = func(string) (*authpkg.ProfilesConfig, error) {
|
||||
return &authpkg.ProfilesConfig{Profiles: []authpkg.Profile{{CorpID: "ding", UserID: "user"}}}, nil
|
||||
}
|
||||
if err := logoutOneProfile(nil, context.Background(), "cfg", "organization-name"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if deletedSelector != "ding" {
|
||||
t.Fatalf("organization deletion selector = %q, want stable corpId", deletedSelector)
|
||||
}
|
||||
|
||||
authEnsureProfilesMigration = func(string) error { return errors.New("migrate") }
|
||||
if err := logoutAllProfiles(nil, context.Background(), "cfg"); err == nil {
|
||||
@@ -536,6 +585,10 @@ func TestCrossPlatformCoverageAuthCoverageStatusAndLogout(t *testing.T) {
|
||||
}
|
||||
revokes := 0
|
||||
authRevokeToken = func(context.Context) error { revokes++; return nil }
|
||||
authRevokeTokenForData = func(context.Context, *authpkg.TokenData) error { revokes++; return nil }
|
||||
authLoadTokenForProfile = func(string, string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{AccessToken: "token"}, nil
|
||||
}
|
||||
authLoadProfiles = func(string) (*authpkg.ProfilesConfig, error) { return nil, nil }
|
||||
authDeleteAllTokenData = func(string) error { return nil }
|
||||
if err := logoutAllProfiles(nil, context.Background(), "cfg"); err != nil || revokes != 1 {
|
||||
@@ -734,7 +787,7 @@ func TestCrossPlatformCoverageAuthCoveragePortableExchangeAndReset(t *testing.T)
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
exchange := newAuthExchangeCommand()
|
||||
exchange := newAuthExchangeCommand(nil)
|
||||
badCode := &cobra.Command{}
|
||||
badCode.Flags().Bool("code", false, "")
|
||||
if err := exchange.RunE(badCode, nil); err == nil {
|
||||
|
||||
@@ -20,6 +20,7 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -498,7 +499,7 @@ func TestAuthStatusDiagnosticReportsCiphertextKeyMismatch(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
func TestAuthStatusRefreshFailureReportsUnauthenticatedDiagnostic(t *testing.T) {
|
||||
// Isolate keychain storage to a per-test directory so the saved
|
||||
// token can't leak into other test packages running in parallel.
|
||||
t.Setenv(keychain.StorageDirEnv, t.TempDir())
|
||||
@@ -517,6 +518,9 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
ExpiresAt: time.Now().Add(-time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
CorpID: "dingcorp",
|
||||
UserID: "user-dingcorp",
|
||||
ClientID: "client-dingcorp",
|
||||
Source: "mcp",
|
||||
})
|
||||
if err != nil {
|
||||
t.Skipf("SaveTokenData() unavailable in this environment: %v", err)
|
||||
@@ -534,7 +538,7 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"auth", "status"})
|
||||
cmd.SetArgs([]string{"--format", "json", "auth", "status"})
|
||||
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("Execute() error = %v\noutput:\n%s", err, out.String())
|
||||
@@ -545,8 +549,18 @@ func TestAuthStatusRefreshFailureLeavesStoredTokenIntact(t *testing.T) {
|
||||
t.Fatal("secure token data should remain in keychain after refresh failure")
|
||||
}
|
||||
|
||||
if !bytes.Contains(out.Bytes(), []byte("\"authenticated\"")) {
|
||||
t.Fatalf("output should still report authenticated status:\n%s", out.String())
|
||||
var resp authStatusResponse
|
||||
if err := json.Unmarshal(out.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if resp.Authenticated {
|
||||
t.Fatalf("authenticated = true after refresh failure: %+v", resp)
|
||||
}
|
||||
if resp.Reason != "token_refresh_failed" {
|
||||
t.Fatalf("reason = %q, want token_refresh_failed: %+v", resp.Reason, resp)
|
||||
}
|
||||
if !strings.Contains(resp.Message, "refresh failed") {
|
||||
t.Fatalf("message = %q, want original refresh failure", resp.Message)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -594,8 +608,35 @@ func TestAuthStatusProfileOverrideDoesNotSwitchCurrentProfile(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_secondary" {
|
||||
t.Fatalf("currentProfile = %q, want unchanged corp_secondary", cfg.CurrentProfile)
|
||||
if cfg.CurrentProfile != "corp_secondary:user-corp_secondary" {
|
||||
t.Fatalf("currentProfile = %q, want unchanged exact secondary identity", cfg.CurrentProfile)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthStatusRejectsAmbiguousProfileSelector(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_first")
|
||||
first.CorpName = "Shared Org"
|
||||
second := authLogoutTestToken("corp_second")
|
||||
second.CorpName = "Shared Org"
|
||||
setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"--format", "json", "auth", "status", "--profile", "Shared Org"})
|
||||
err := cmd.Execute()
|
||||
if err == nil {
|
||||
t.Fatalf("auth status accepted ambiguous profile selector\noutput:\n%s", out.String())
|
||||
}
|
||||
var appErr *apperrors.Error
|
||||
if !errors.As(err, &appErr) || appErr.Category != apperrors.CategoryValidation {
|
||||
t.Fatalf("error = %T %v, want validation error", err, err)
|
||||
}
|
||||
for _, candidate := range []string{"corp_first", "corp_second"} {
|
||||
if !strings.Contains(err.Error(), candidate) {
|
||||
t.Fatalf("error = %q, want candidate %q", err.Error(), candidate)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -769,8 +810,8 @@ func TestAuthLogoutProfileDeletesOnlySelectedProfile(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.PrimaryProfile != "corp_secondary" || cfg.CurrentProfile != "corp_secondary" {
|
||||
t.Fatalf("profiles pointers = primary %q current %q, want corp_secondary/corp_secondary", cfg.PrimaryProfile, cfg.CurrentProfile)
|
||||
if cfg.PrimaryProfile != "" || cfg.CurrentProfile != "corp_secondary:user-corp_secondary" {
|
||||
t.Fatalf("profiles pointers = primary %q current %q", cfg.PrimaryProfile, cfg.CurrentProfile)
|
||||
}
|
||||
if len(cfg.Profiles) != 1 || cfg.Profiles[0].CorpID != "corp_secondary" {
|
||||
t.Fatalf("profiles = %#v, want only corp_secondary retained", cfg.Profiles)
|
||||
@@ -790,6 +831,102 @@ func TestAuthLogoutProfileDeletesOnlySelectedProfile(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLogoutExactProfilePreservesSameCorpAccount(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_same")
|
||||
first.UserID = "user_1"
|
||||
second := authLogoutTestToken("corp_same")
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
configDir := setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
originalTransport := http.DefaultTransport
|
||||
t.Cleanup(func() {
|
||||
http.DefaultTransport = originalTransport
|
||||
})
|
||||
http.DefaultTransport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
return nil, errors.New("remote revoke disabled in unit test")
|
||||
})
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"auth", "logout", "--profile", "corp_same:user_2"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("auth logout exact profile error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if len(cfg.Profiles) != 1 || cfg.Profiles[0].UserID != "user_1" {
|
||||
t.Fatalf("profiles = %#v, want only user_1 retained", cfg.Profiles)
|
||||
}
|
||||
if authpkg.TokenDataExistsKeychainForIdentity("corp_same", "user_2") {
|
||||
t.Fatal("selected identity token should be deleted")
|
||||
}
|
||||
loaded, err := authpkg.LoadTokenDataForProfile(configDir, "corp_same")
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenDataForProfile(org) error = %v", err)
|
||||
}
|
||||
if loaded.UserID != "user_1" || loaded.AccessToken != first.AccessToken {
|
||||
t.Fatalf("org current token = %#v, want retained user_1", loaded)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLogoutLocalProfileNameRevokesOnlySelectedAccount(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_same")
|
||||
first.UserID = "user_1"
|
||||
first.UserName = "账号一"
|
||||
second := authLogoutTestToken("corp_same")
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
second.UserName = "账号二"
|
||||
configDir := setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
var selector string
|
||||
for _, profile := range cfg.Profiles {
|
||||
if profile.UserID == "user_2" {
|
||||
selector = profile.Name
|
||||
}
|
||||
}
|
||||
if selector == "" || selector == second.CorpName {
|
||||
t.Fatalf("second local profile name = %q, want unique non-org alias", selector)
|
||||
}
|
||||
|
||||
requests := 0
|
||||
originalTransport := http.DefaultTransport
|
||||
t.Cleanup(func() {
|
||||
http.DefaultTransport = originalTransport
|
||||
})
|
||||
http.DefaultTransport = roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
requests++
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader("")),
|
||||
}, nil
|
||||
})
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"auth", "logout", "--profile", selector})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("auth logout local profile error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if requests != 1 {
|
||||
t.Fatalf("remote revoke requests = %d, want 1", requests)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLoginPostLoginTUIModeRespectsRecommendAndFormat(t *testing.T) {
|
||||
newRoot := func(t *testing.T) *cobra.Command {
|
||||
t.Helper()
|
||||
@@ -875,8 +1012,15 @@ func TestLoginRecommendProductLabelMatchesTUITarget(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestResolveAuthLoginConfigReadsInheritedYes(t *testing.T) {
|
||||
t.Setenv("DWS_DEBUG_AUTH", "1")
|
||||
var logs bytes.Buffer
|
||||
previousLogger := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, &slog.HandlerOptions{Level: slog.LevelDebug})))
|
||||
t.Cleanup(func() { slog.SetDefault(previousLogger) })
|
||||
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.PersistentFlags().Bool("yes", false, "")
|
||||
root.PersistentFlags().String("profile", "", "")
|
||||
login := &cobra.Command{Use: "login"}
|
||||
login.Flags().String("token", "", "")
|
||||
login.Flags().Bool("device", false, "")
|
||||
@@ -901,6 +1045,11 @@ func TestResolveAuthLoginConfigReadsInheritedYes(t *testing.T) {
|
||||
if !cfg.Yes {
|
||||
t.Fatal("Yes = false, want true")
|
||||
}
|
||||
if got := logs.String(); !strings.Contains(got, `"msg":"auth.login.request"`) ||
|
||||
!strings.Contains(got, `"profile_selector":""`) ||
|
||||
!strings.Contains(got, `"target_corp_id":""`) {
|
||||
t.Fatalf("login request diagnostic log missing selector resolution:\n%s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLoginForcesAuthorizationByDefault(t *testing.T) {
|
||||
@@ -951,6 +1100,12 @@ func TestAuthLoginRecommendSkipsPostLoginTUI(t *testing.T) {
|
||||
`{"success":true,"data":{"items":[{"scope":"calendar.event:read","productCode":"calendar","productName":"日历"}],"selectedScopes":["calendar.event:read"]}}`,
|
||||
`{"success":true,"data":{"grantedScopes":["calendar.event:read"]}}`,
|
||||
}}
|
||||
authpkg.SetRuntimeProfile("corp_old:user_old")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
var authorizationProfiles []string
|
||||
fake.beforeCall = func(string) {
|
||||
authorizationProfiles = append(authorizationProfiles, authpkg.RuntimeProfile())
|
||||
}
|
||||
cmd := newAuthLoginCommand(fake)
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
@@ -969,6 +1124,61 @@ func TestAuthLoginRecommendSkipsPostLoginTUI(t *testing.T) {
|
||||
if got := fake.args[0]["recommend"]; got != true {
|
||||
t.Fatalf("--recommend plan recommend = %#v, want true", got)
|
||||
}
|
||||
for _, profile := range authorizationProfiles {
|
||||
if profile != "" {
|
||||
t.Fatalf("manual token post-login profile = %q, want empty runtime selector", profile)
|
||||
}
|
||||
}
|
||||
if got := authpkg.RuntimeProfile(); got != "corp_old:user_old" {
|
||||
t.Fatalf("runtime profile after authorization = %q, want restored selector", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLoginRecommendUsesNewExactIdentity(t *testing.T) {
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
t.Setenv(keychain.StorageDirEnv, t.TempDir())
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
oldOAuthLogin := authOAuthLogin
|
||||
oldInteractive := authLoginInteractiveTerminal
|
||||
t.Cleanup(func() {
|
||||
authOAuthLogin = oldOAuthLogin
|
||||
authLoginInteractiveTerminal = oldInteractive
|
||||
authpkg.SetRuntimeProfile("")
|
||||
})
|
||||
authLoginInteractiveTerminal = func() bool { return false }
|
||||
authOAuthLogin = func(*authpkg.OAuthProvider, context.Context, bool) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{
|
||||
AccessToken: "new-token",
|
||||
CorpID: "corp_same",
|
||||
UserID: "user_new",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
}, nil
|
||||
}
|
||||
fake := &authLoginRecommendSequenceCaller{responses: []string{
|
||||
`{"success":true,"data":{"items":[],"selectedScopes":[]}}`,
|
||||
}}
|
||||
var authorizationProfiles []string
|
||||
fake.beforeCall = func(string) {
|
||||
authorizationProfiles = append(authorizationProfiles, authpkg.RuntimeProfile())
|
||||
}
|
||||
authpkg.SetRuntimeProfile("corp_same:user_old")
|
||||
|
||||
cmd := newAuthLoginCommand(fake)
|
||||
cmd.SetOut(io.Discard)
|
||||
cmd.SetErr(io.Discard)
|
||||
cmd.SetArgs([]string{"--recommend"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("auth login --recommend error = %v", err)
|
||||
}
|
||||
for _, profile := range authorizationProfiles {
|
||||
if profile != "corp_same:user_new" {
|
||||
t.Fatalf("post-login authorization profile = %q, want new exact identity", profile)
|
||||
}
|
||||
}
|
||||
if got := authpkg.RuntimeProfile(); got != "corp_same:user_old" {
|
||||
t.Fatalf("runtime profile after authorization = %q, want restored old identity", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthLoginDefaultTUIModeSkipsSelectorWhenAllGranted(t *testing.T) {
|
||||
@@ -1193,7 +1403,7 @@ func TestAuthLoginDefaultTUIRunsAfterLoginTokenSaved(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnrichAuthLoginProfileFromContactPersistsCorpName(t *testing.T) {
|
||||
func TestEnrichAuthLoginProfileFromContactBeforePersist(t *testing.T) {
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
t.Setenv(keychain.StorageDirEnv, t.TempDir())
|
||||
configDir := t.TempDir()
|
||||
@@ -1208,12 +1418,8 @@ func TestEnrichAuthLoginProfileFromContactPersistsCorpName(t *testing.T) {
|
||||
ClientID: "client-id",
|
||||
Source: "mcp",
|
||||
}
|
||||
if err := authpkg.SaveTokenData(configDir, token); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
|
||||
fake := &authLoginRecommendSequenceCaller{responses: []string{
|
||||
`{"success":true,"result":[{"orgEmployeeModel":{"corpId":"ding32fff839a3e0105d","orgName":"钉钉(中国)信息技术有限公司","userId":"011352590165863362195","orgUserName":"玄玦(主用钉)"}}]}`,
|
||||
`{"success":true,"result":[{"isAdmin":false,"orgEmployeeModel":{"jobNumber":"202397","orgId":null,"orgName":"钉钉(中国)信息技术有限公司","orgUserId":"011352590165863362195","orgUserName":"玄玦(主用钉)"}}]}`,
|
||||
}}
|
||||
if err := enrichAuthLoginProfileFromContact(context.Background(), configDir, fake, token); err != nil {
|
||||
t.Fatalf("enrichAuthLoginProfileFromContact() error = %v", err)
|
||||
@@ -1224,6 +1430,16 @@ func TestEnrichAuthLoginProfileFromContactPersistsCorpName(t *testing.T) {
|
||||
if token.UserID != "011352590165863362195" || token.UserName != "玄玦(主用钉)" {
|
||||
t.Fatalf("token user identity = (%q, %q), want contact result", token.UserID, token.UserName)
|
||||
}
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() before persist error = %v", err)
|
||||
}
|
||||
if len(cfg.Profiles) != 0 {
|
||||
t.Fatalf("identity enrichment persisted token early: %#v", cfg.Profiles)
|
||||
}
|
||||
if err := authpkg.SaveTokenData(configDir, token); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
|
||||
loaded, err := authpkg.LoadTokenDataForProfile(configDir, "ding32fff839a3e0105d")
|
||||
if err != nil {
|
||||
@@ -1235,8 +1451,50 @@ func TestEnrichAuthLoginProfileFromContactPersistsCorpName(t *testing.T) {
|
||||
if len(fake.tools) != 1 || fake.tools[0] != "get_current_user_profile" {
|
||||
t.Fatalf("tool calls = %v, want get_current_user_profile", fake.tools)
|
||||
}
|
||||
if got := fake.args[0]["profile"]; got != "ding32fff839a3e0105d" {
|
||||
t.Fatalf("contact profile arg = %#v, want ding32fff839a3e0105d", got)
|
||||
if len(fake.args[0]) != 0 {
|
||||
t.Fatalf("contact profile args = %#v, want no arguments", fake.args[0])
|
||||
}
|
||||
if len(fake.tokens) != 1 || fake.tokens[0] != "access-token" {
|
||||
t.Fatalf("token overrides = %v, want access-token", fake.tokens)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnrichAuthLoginProfileLogsIdentityResolutionWithoutCredentials(t *testing.T) {
|
||||
t.Setenv("DWS_DEBUG_AUTH", "1")
|
||||
var logs bytes.Buffer
|
||||
previousLogger := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, &slog.HandlerOptions{Level: slog.LevelDebug})))
|
||||
t.Cleanup(func() { slog.SetDefault(previousLogger) })
|
||||
|
||||
token := &authpkg.TokenData{
|
||||
AccessToken: "secret-access-token",
|
||||
RefreshToken: "secret-refresh-token",
|
||||
CorpID: "ding_same_corp",
|
||||
}
|
||||
fake := &authLoginRecommendSequenceCaller{responses: []string{
|
||||
`{"success":true,"result":[{"orgEmployeeModel":{"corpId":"ding_same_corp","orgName":"同一组织","userId":"user_two","orgUserName":"账号二"}}]}`,
|
||||
}}
|
||||
|
||||
if err := enrichAuthLoginProfileFromContact(context.Background(), t.TempDir(), fake, token); err != nil {
|
||||
t.Fatalf("enrichAuthLoginProfileFromContact() error = %v", err)
|
||||
}
|
||||
|
||||
got := logs.String()
|
||||
for _, want := range []string{
|
||||
`"msg":"auth.login.identity.lookup.start"`,
|
||||
`"msg":"auth.login.identity.lookup.result"`,
|
||||
`"corp_id":"ding_same_corp"`,
|
||||
`"user_id":"user_two"`,
|
||||
`"user_name":"账号二"`,
|
||||
} {
|
||||
if !strings.Contains(got, want) {
|
||||
t.Fatalf("diagnostic logs missing %q:\n%s", want, got)
|
||||
}
|
||||
}
|
||||
for _, secret := range []string{"secret-access-token", "secret-refresh-token"} {
|
||||
if strings.Contains(got, secret) {
|
||||
t.Fatalf("diagnostic logs exposed credential %q:\n%s", secret, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1250,6 +1508,7 @@ type authLoginRecommendSequenceCaller struct {
|
||||
responses []string
|
||||
tools []string
|
||||
args []map[string]any
|
||||
tokens []string
|
||||
beforeCall func(toolName string)
|
||||
}
|
||||
|
||||
@@ -1271,6 +1530,11 @@ func (f *authLoginRecommendSequenceCaller) CallTool(_ context.Context, _ string,
|
||||
return &edition.ToolResult{Content: []edition.ContentBlock{{Type: "text", Text: response}}}, nil
|
||||
}
|
||||
|
||||
func (f *authLoginRecommendSequenceCaller) CallToolWithToken(ctx context.Context, token, productID, toolName string, args map[string]any) (*edition.ToolResult, error) {
|
||||
f.tokens = append(f.tokens, token)
|
||||
return f.CallTool(ctx, productID, toolName, args)
|
||||
}
|
||||
|
||||
func (f *authLoginRecommendSequenceCaller) Format() string { return "table" }
|
||||
|
||||
func (f *authLoginRecommendSequenceCaller) DryRun() bool { return false }
|
||||
|
||||
@@ -15,6 +15,15 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/authretry"
|
||||
)
|
||||
|
||||
// authRetryingKey marks a context that has already attempted one
|
||||
@@ -23,8 +32,35 @@ import (
|
||||
// to the user instead.
|
||||
type authRetryingKeyType struct{}
|
||||
|
||||
type authRefreshFailureError struct {
|
||||
rejection error
|
||||
refresh error
|
||||
}
|
||||
|
||||
func (e *authRefreshFailureError) Error() string {
|
||||
return "automatic access token refresh failed"
|
||||
}
|
||||
|
||||
func (e *authRefreshFailureError) Unwrap() []error {
|
||||
if e == nil {
|
||||
return nil
|
||||
}
|
||||
return []error{e.rejection, e.refresh}
|
||||
}
|
||||
|
||||
var authRetryingKey = authRetryingKeyType{}
|
||||
|
||||
var (
|
||||
runnerForceRefreshRejectedAccessToken = forceRefreshRejectedAccessToken
|
||||
runnerExecuteAuthRetry func(*runtimeRunner, context.Context, string, executor.Invocation) (executor.Result, error)
|
||||
)
|
||||
|
||||
func init() {
|
||||
runnerExecuteAuthRetry = func(r *runtimeRunner, ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
|
||||
return r.executeInvocation(ctx, endpoint, invocation)
|
||||
}
|
||||
}
|
||||
|
||||
// IsAuthRetrying reports whether the current context is already inside an
|
||||
// AuthRefreshRequired retry. Mirrors IsPatRetrying.
|
||||
func IsAuthRetrying(ctx context.Context) bool {
|
||||
@@ -34,3 +70,103 @@ func IsAuthRetrying(ctx context.Context) bool {
|
||||
v, _ := ctx.Value(authRetryingKey).(bool)
|
||||
return v
|
||||
}
|
||||
|
||||
func withAuthRetrying(ctx context.Context) context.Context {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
return context.WithValue(ctx, authRetryingKey, true)
|
||||
}
|
||||
|
||||
func authRefreshLogger() *slog.Logger {
|
||||
if logger := FileLoggerInstance(); logger != nil {
|
||||
return logger
|
||||
}
|
||||
return slog.Default()
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) managesRuntimeOAuth(hasPluginAuth bool) bool {
|
||||
if r == nil || hasPluginAuth {
|
||||
return false
|
||||
}
|
||||
return r.globalFlags == nil || strings.TrimSpace(r.globalFlags.Token) == ""
|
||||
}
|
||||
|
||||
// retryAuthRefreshRequired consumes only the explicit edition marker. It does
|
||||
// not infer retryability from free text, generic auth categories, HTTP 403, or
|
||||
// ordinary business errors.
|
||||
func (r *runtimeRunner) retryAuthRefreshRequired(
|
||||
ctx context.Context,
|
||||
endpoint string,
|
||||
invocation executor.Invocation,
|
||||
rejectedAccessToken string,
|
||||
markerErr error,
|
||||
hasPluginAuth bool,
|
||||
) (executor.Result, error, bool) {
|
||||
marker, marked := authretry.As(markerErr)
|
||||
if !marked {
|
||||
return executor.Result{}, nil, false
|
||||
}
|
||||
cause := marker.Cause
|
||||
if cause == nil {
|
||||
cause = markerErr
|
||||
}
|
||||
|
||||
// Explicit --token and plugin credentials are not backed by the default
|
||||
// OAuth refresh store. Preserve the overlay cause without mutating an
|
||||
// unrelated persisted login.
|
||||
if !r.managesRuntimeOAuth(hasPluginAuth) {
|
||||
return executor.Result{}, cause, true
|
||||
}
|
||||
if IsAuthRetrying(ctx) {
|
||||
authRefreshLogger().Warn("auth.runtime.refresh.retry_exhausted",
|
||||
"product", invocation.CanonicalProduct,
|
||||
"tool", invocation.Tool,
|
||||
)
|
||||
return executor.Result{}, cause, true
|
||||
}
|
||||
|
||||
if _, err := runnerForceRefreshRejectedAccessToken(ctx, defaultConfigDir(), rejectedAccessToken); err != nil {
|
||||
// Keep every log credential-safe. The returned error chain retains the
|
||||
// complete cause for in-process diagnosis; even DWS_DEBUG_AUTH must not
|
||||
// serialize an OAuth response body or other attacker-controlled text.
|
||||
authRefreshLogger().Warn("auth.runtime.refresh.failed",
|
||||
"product", invocation.CanonicalProduct,
|
||||
"tool", invocation.Tool,
|
||||
"stage", "force_refresh_rejected_token",
|
||||
"error_type", fmt.Sprintf("%T", err),
|
||||
)
|
||||
logging.AuthDebug("auth.runtime.refresh.failed.detail",
|
||||
"product", invocation.CanonicalProduct,
|
||||
"tool", invocation.Tool,
|
||||
"stage", "force_refresh_rejected_token",
|
||||
"error_type", fmt.Sprintf("%T", err),
|
||||
)
|
||||
combined := &authRefreshFailureError{rejection: cause, refresh: err}
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"automatic access token refresh failed",
|
||||
apperrors.WithOperation("auth/token/refresh"),
|
||||
apperrors.WithReason("auth_refresh_failed"),
|
||||
apperrors.WithHint("本地凭证已保留;可稍后重试,若持续失败请查看认证诊断日志。"),
|
||||
apperrors.WithCause(combined),
|
||||
), true
|
||||
}
|
||||
|
||||
logging.AuthDebug("auth.runtime.refresh.succeeded",
|
||||
"product", invocation.CanonicalProduct,
|
||||
"tool", invocation.Tool,
|
||||
)
|
||||
result, err := runnerExecuteAuthRetry(r, withAuthRetrying(ctx), endpoint, invocation)
|
||||
return result, err, true
|
||||
}
|
||||
|
||||
// isRefreshableTransportAuthError deliberately excludes HTTP/RPC 403 and
|
||||
// generic CategoryAuth values. OnAuthError may request a refresh only for an
|
||||
// exact transport-level unauthorized signal.
|
||||
func isRefreshableTransportAuthError(err error) bool {
|
||||
var typed *apperrors.Error
|
||||
if !errors.As(err, &typed) || typed.Category != apperrors.CategoryAuth {
|
||||
return false
|
||||
}
|
||||
return typed.Reason == "http_401" || typed.RPCCode == 401
|
||||
}
|
||||
|
||||
@@ -0,0 +1,354 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/audit"
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/authretry"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func installAuthRefreshRunnerSeams(t *testing.T) {
|
||||
t.Helper()
|
||||
previousHooks := edition.Get()
|
||||
previousCall := runnerCallTool
|
||||
previousPreflight := runnerPreflightDocDownload
|
||||
previousRefresh := runnerForceRefreshRejectedAccessToken
|
||||
previousRetry := runnerExecuteAuthRetry
|
||||
previousCapture := runnerCaptureRuntimeFailure
|
||||
previousProfile := authpkg.RuntimeProfile()
|
||||
|
||||
pluginAuthMu.Lock()
|
||||
previousPlugins := pluginAuthRegistry
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
pluginAuthMu.Unlock()
|
||||
|
||||
runnerPreflightDocDownload = func(*runtimeRunner, context.Context, *transport.Client, string, executor.Invocation) error {
|
||||
return nil
|
||||
}
|
||||
runnerCaptureRuntimeFailure = func(executor.Invocation, error, error) {}
|
||||
authpkg.SetRuntimeProfile("")
|
||||
runtimeTokenManager.Invalidate()
|
||||
t.Setenv("DWS_CONFIG_DIR", "")
|
||||
t.Setenv("DWS_DEBUG_AUTH", "0")
|
||||
|
||||
t.Cleanup(func() {
|
||||
edition.Override(previousHooks)
|
||||
runnerCallTool = previousCall
|
||||
runnerPreflightDocDownload = previousPreflight
|
||||
runnerForceRefreshRejectedAccessToken = previousRefresh
|
||||
runnerExecuteAuthRetry = previousRetry
|
||||
runnerCaptureRuntimeFailure = previousCapture
|
||||
authpkg.SetRuntimeProfile(previousProfile)
|
||||
runtimeTokenManager.Invalidate()
|
||||
pluginAuthMu.Lock()
|
||||
pluginAuthRegistry = previousPlugins
|
||||
pluginAuthMu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
func authRefreshTestRunner(flags *GlobalFlags) *runtimeRunner {
|
||||
return &runtimeRunner{
|
||||
transport: transport.NewClient(nil),
|
||||
globalFlags: flags,
|
||||
auditSink: audit.NopSink{},
|
||||
}
|
||||
}
|
||||
|
||||
func authRefreshTestInvocation() executor.Invocation {
|
||||
return executor.Invocation{
|
||||
CanonicalProduct: "auth-retry-test-product",
|
||||
Tool: "test_tool",
|
||||
Params: map[string]any{"value": "safe"},
|
||||
}
|
||||
}
|
||||
|
||||
func authRefreshTokenHooks(configDir string, token *string, classify func(map[string]any) error) *edition.Hooks {
|
||||
return &edition.Hooks{
|
||||
ConfigDir: func() string { return configDir },
|
||||
TokenProvider: func(context.Context, func() (string, error)) (string, error) {
|
||||
return *token, nil
|
||||
},
|
||||
ClassifyToolResult: classify,
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerRetriesEditionAuthMarkerOnce(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
rejection := apperrors.NewAuth("server rejected access token", apperrors.WithReason("access_token_rejected"))
|
||||
edition.Override(authRefreshTokenHooks(configDir, &token, func(content map[string]any) error {
|
||||
if expired, _ := content["expired"].(bool); expired {
|
||||
return &authretry.AuthRefreshRequired{Cause: rejection}
|
||||
}
|
||||
return nil
|
||||
}))
|
||||
|
||||
var callTokens []string
|
||||
runnerCallTool = func(client *transport.Client, _ context.Context, _, _ string, _ map[string]any) (transport.ToolCallResult, error) {
|
||||
callTokens = append(callTokens, client.AuthToken)
|
||||
if len(callTokens) == 1 {
|
||||
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
|
||||
}
|
||||
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(_ context.Context, gotDir, rejected string) (string, error) {
|
||||
refreshCalls++
|
||||
if gotDir != configDir || rejected != "old-access" {
|
||||
t.Fatalf("refresh input = dir %q token %q", gotDir, rejected)
|
||||
}
|
||||
token = "new-access"
|
||||
return token, nil
|
||||
}
|
||||
|
||||
result, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if refreshCalls != 1 || len(callTokens) != 2 || callTokens[0] != "old-access" || callTokens[1] != "new-access" {
|
||||
t.Fatalf("refreshes=%d call tokens=%v", refreshCalls, callTokens)
|
||||
}
|
||||
content, _ := result.Response["content"].(map[string]any)
|
||||
if content["value"] != "ok" || content["success"] != true {
|
||||
t.Fatalf("result content = %#v", content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerRefreshFailurePreservesBothCausesAndSafeLog(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
t.Setenv("DWS_DEBUG_AUTH", "1")
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
rejection := apperrors.NewAuth("server rejected access token", apperrors.WithReason("access_token_rejected"))
|
||||
edition.Override(authRefreshTokenHooks(configDir, &token, func(map[string]any) error {
|
||||
return &authretry.AuthRefreshRequired{Cause: rejection}
|
||||
}))
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
|
||||
}
|
||||
refreshErr := errors.New(`oauth refresh response parse failed: body={"access_token":"access-token-secret","refresh_token":"refresh-token-secret","uid":"uid-secret-value"}`)
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", refreshErr
|
||||
}
|
||||
|
||||
var logs bytes.Buffer
|
||||
previousLogger := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, &slog.HandlerOptions{Level: slog.LevelDebug})))
|
||||
t.Cleanup(func() { slog.SetDefault(previousLogger) })
|
||||
|
||||
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
|
||||
if !errors.Is(err, rejection) || !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("error = %v, want rejection and refresh causes", err)
|
||||
}
|
||||
var typed *apperrors.Error
|
||||
if !errors.As(err, &typed) || typed.Category != apperrors.CategoryAuth || typed.Reason != "auth_refresh_failed" || typed.Operation != "auth/token/refresh" {
|
||||
t.Fatalf("refresh envelope = %#v", typed)
|
||||
}
|
||||
var rendered bytes.Buffer
|
||||
if printErr := apperrors.PrintJSON(&rendered, err); printErr != nil {
|
||||
t.Fatal(printErr)
|
||||
}
|
||||
for _, want := range []string{`"category": "auth"`, `"reason": "auth_refresh_failed"`, `"operation": "auth/token/refresh"`} {
|
||||
if !strings.Contains(rendered.String(), want) {
|
||||
t.Fatalf("structured stderr missing %s: %s", want, rendered.String())
|
||||
}
|
||||
}
|
||||
for _, secret := range []string{"access-token-secret", "refresh-token-secret", "uid-secret-value"} {
|
||||
if strings.Contains(err.Error(), secret) || strings.Contains(logs.String(), secret) || strings.Contains(rendered.String(), secret) {
|
||||
t.Fatalf("auth output leaked %q: error=%q logs=%s stderr=%s", secret, err, logs.String(), rendered.String())
|
||||
}
|
||||
}
|
||||
for _, want := range []string{"auth.runtime.refresh.failed", "auth.runtime.refresh.failed.detail", "force_refresh_rejected_token", "error_type"} {
|
||||
if !strings.Contains(logs.String(), want) {
|
||||
t.Fatalf("safe refresh log missing %q: %s", want, logs.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerSecondEditionMarkerReturnsSecondCause(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
firstCause := errors.New("first rejection")
|
||||
secondCause := errors.New("second rejection")
|
||||
edition.Override(authRefreshTokenHooks(configDir, &token, func(content map[string]any) error {
|
||||
attempt, _ := content["attempt"].(int)
|
||||
if attempt == 1 {
|
||||
return &authretry.AuthRefreshRequired{Cause: firstCause}
|
||||
}
|
||||
return &authretry.AuthRefreshRequired{Cause: secondCause}
|
||||
}))
|
||||
calls := 0
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
calls++
|
||||
return transport.ToolCallResult{Content: map[string]any{"attempt": calls}}, nil
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
refreshCalls++
|
||||
token = "new-access"
|
||||
return token, nil
|
||||
}
|
||||
|
||||
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
|
||||
if !errors.Is(err, secondCause) || errors.Is(err, firstCause) {
|
||||
t.Fatalf("error = %v, want only second rejection cause", err)
|
||||
}
|
||||
if calls != 2 || refreshCalls != 1 {
|
||||
t.Fatalf("calls=%d refreshes=%d", calls, refreshCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerOnAuthErrorOnlyRetriesExactUnauthorized(t *testing.T) {
|
||||
t.Run("http 401 marker retries once", func(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
rejection := errors.New("transport rejected token")
|
||||
hookCalls := 0
|
||||
hooks := authRefreshTokenHooks(configDir, &token, nil)
|
||||
hooks.OnAuthError = func(string, error) error {
|
||||
hookCalls++
|
||||
return &authretry.AuthRefreshRequired{Cause: rejection}
|
||||
}
|
||||
edition.Override(hooks)
|
||||
calls := 0
|
||||
var callTokens []string
|
||||
runnerCallTool = func(client *transport.Client, _ context.Context, _, _ string, _ map[string]any) (transport.ToolCallResult, error) {
|
||||
calls++
|
||||
callTokens = append(callTokens, client.AuthToken)
|
||||
if calls == 1 {
|
||||
return transport.ToolCallResult{}, apperrors.NewAuth("unauthorized", apperrors.WithReason("http_401"))
|
||||
}
|
||||
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
refreshCalls++
|
||||
token = "new-access"
|
||||
return token, nil
|
||||
}
|
||||
|
||||
if _, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if hookCalls != 1 || refreshCalls != 1 || calls != 2 || strings.Join(callTokens, ",") != "old-access,new-access" {
|
||||
t.Fatalf("hook=%d refresh=%d calls=%d tokens=%v", hookCalls, refreshCalls, calls, callTokens)
|
||||
}
|
||||
})
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
err error
|
||||
}{
|
||||
{name: "http 403", err: apperrors.NewAuth("forbidden", apperrors.WithReason("http_403"))},
|
||||
{name: "ordinary auth", err: apperrors.NewAuth("load failed", apperrors.WithReason("auth_load_failed"))},
|
||||
} {
|
||||
t.Run(tc.name+" does not enter hook", func(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
hookCalls := 0
|
||||
hooks := authRefreshTokenHooks(configDir, &token, nil)
|
||||
hooks.OnAuthError = func(string, error) error {
|
||||
hookCalls++
|
||||
return &authretry.AuthRefreshRequired{Cause: errors.New("must not run")}
|
||||
}
|
||||
edition.Override(hooks)
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
return transport.ToolCallResult{}, tc.err
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
refreshCalls++
|
||||
return "", nil
|
||||
}
|
||||
|
||||
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
|
||||
if !errors.Is(err, tc.err) || hookCalls != 0 || refreshCalls != 0 {
|
||||
t.Fatalf("error=%v hook=%d refresh=%d", err, hookCalls, refreshCalls)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerDoesNotRefreshExplicitTokenMarker(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
rejection := errors.New("explicit token rejected")
|
||||
edition.Override(&edition.Hooks{ClassifyToolResult: func(map[string]any) error {
|
||||
return &authretry.AuthRefreshRequired{Cause: rejection}
|
||||
}})
|
||||
calls := 0
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
calls++
|
||||
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
refreshCalls++
|
||||
return "", nil
|
||||
}
|
||||
|
||||
_, err := authRefreshTestRunner(&GlobalFlags{Token: "explicit-token"}).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
|
||||
if !errors.Is(err, rejection) || calls != 1 || refreshCalls != 0 {
|
||||
t.Fatalf("error=%v calls=%d refresh=%d", err, calls, refreshCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageRunnerRetriesPreflightEditionMarkerOnce(t *testing.T) {
|
||||
installAuthRefreshRunnerSeams(t)
|
||||
configDir := t.TempDir()
|
||||
token := "old-access"
|
||||
rejection := errors.New("preflight token rejected")
|
||||
edition.Override(authRefreshTokenHooks(configDir, &token, nil))
|
||||
preflightCalls := 0
|
||||
runnerPreflightDocDownload = func(*runtimeRunner, context.Context, *transport.Client, string, executor.Invocation) error {
|
||||
preflightCalls++
|
||||
if preflightCalls == 1 {
|
||||
return &authretry.AuthRefreshRequired{Cause: rejection}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
toolCalls := 0
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
toolCalls++
|
||||
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
|
||||
}
|
||||
refreshCalls := 0
|
||||
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
|
||||
refreshCalls++
|
||||
token = "new-access"
|
||||
return token, nil
|
||||
}
|
||||
|
||||
if _, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if preflightCalls != 2 || toolCalls != 1 || refreshCalls != 1 {
|
||||
t.Fatalf("preflights=%d tools=%d refreshes=%d", preflightCalls, toolCalls, refreshCalls)
|
||||
}
|
||||
}
|
||||
@@ -49,6 +49,15 @@ func RegisterPluginAuth(productID string, auth *PluginAuth) {
|
||||
pluginAuthRegistry[productID] = auth
|
||||
}
|
||||
|
||||
// ClearPluginAuth removes credentials for a plugin product. Registration uses
|
||||
// this before applying an accepted descriptor so a descriptor without custom
|
||||
// auth cannot inherit stale credentials from an earlier root construction.
|
||||
func ClearPluginAuth(productID string) {
|
||||
pluginAuthMu.Lock()
|
||||
defer pluginAuthMu.Unlock()
|
||||
delete(pluginAuthRegistry, productID)
|
||||
}
|
||||
|
||||
// LookupPluginAuth returns the authentication credentials registered
|
||||
// for the given product ID, or nil if none exists.
|
||||
func LookupPluginAuth(productID string) (*PluginAuth, bool) {
|
||||
|
||||
@@ -74,6 +74,19 @@ func TestCrossPlatformCoverageToolCallerAdapterCoverage(t *testing.T) {
|
||||
t.Fatal("adapter flags were not forwarded")
|
||||
}
|
||||
adapter.flags.DryRun = false
|
||||
if _, err := (*toolCallerAdapter)(nil).CallToolWithToken(context.Background(), "token", "doc", "get", nil); err == nil {
|
||||
t.Fatal("nil adapter token override succeeded")
|
||||
}
|
||||
authpkg.SetRuntimeProfile("saved-profile")
|
||||
adapter.flags.Token = "saved-token"
|
||||
runner.result = executor.Result{Response: map[string]any{"content": []any{map[string]any{"type": "text", "text": "ok"}}}}
|
||||
if _, err := adapter.CallToolWithToken(context.Background(), "temporary-token", "doc", "get", nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if adapter.flags.Token != "saved-token" || authpkg.RuntimeProfile() != "saved-profile" {
|
||||
t.Fatal("token override state was not restored")
|
||||
}
|
||||
authpkg.SetRuntimeProfile("")
|
||||
|
||||
runner.err = errors.New("runner failure")
|
||||
if _, err := adapter.CallTool(context.Background(), "doc", "get", map[string]any{"x": 1}); !errors.Is(err, runner.err) {
|
||||
@@ -541,7 +554,7 @@ func TestCrossPlatformCoverageRecoveryRuntimeHTTP(t *testing.T) {
|
||||
defer server.Close()
|
||||
SetDynamicServers([]mcptypes.ServerDescriptor{{Endpoint: server.URL, CLI: mcptypes.CLIOverlay{ID: "devdoc", Tools: []mcptypes.CLITool{{Name: "search_open_platform_docs_rag"}}}}})
|
||||
t.Cleanup(func() { SetDynamicServers(nil) })
|
||||
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client())}
|
||||
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client()), flags: &GlobalFlags{Token: "token"}}
|
||||
got, err := runtime.Search(context.Background(), "query", recovery.RecoveryContext{ToolName: "search"})
|
||||
if err != nil || got.DocSearch.Status != "success" || len(got.KBHits) == 0 {
|
||||
t.Fatalf("recovery search = %#v %v", got, err)
|
||||
@@ -1462,14 +1475,11 @@ func TestCrossPlatformCoveragePersonalEventPureCoverage(t *testing.T) {
|
||||
if err := renderPersonalSchema(io.Discard, def, "yaml"); err == nil {
|
||||
t.Fatal("unsupported schema format succeeded")
|
||||
}
|
||||
for _, key := range []string{"", "unknown", personal.EventMention} {
|
||||
for _, key := range []string{"", "unknown", personal.EventMention, personal.EventFromUser} {
|
||||
if err := ensurePublicPersonalEvent(key); err != nil {
|
||||
t.Fatalf("public event %q: %v", key, err)
|
||||
}
|
||||
}
|
||||
if err := ensurePublicPersonalEvent(personal.EventFromUser); err == nil {
|
||||
t.Fatal("private event accepted")
|
||||
}
|
||||
|
||||
var cfg consume.Config
|
||||
applyPersonalConsumeFilters(nil, personalConsumeOptions{}, "", "")
|
||||
@@ -1604,7 +1614,7 @@ func TestCrossPlatformCoveragePersonalSubscriptionAndSourceCoverage(t *testing.T
|
||||
{"lookup", personalConsumeOptions{SubscribeID: "existing"}, true},
|
||||
{"create", personalConsumeOptions{EventKey: personal.EventMention, Name: "name", FilterJSON: `{"field":"content","op":"eq","value":"x"}`, QueryCSV: "a,b", TTL: time.Minute}, true},
|
||||
{"missing", personalConsumeOptions{}, false},
|
||||
{"private", personalConsumeOptions{EventKey: personal.EventFromUser, UserID: "u"}, false},
|
||||
{"sender", personalConsumeOptions{EventKey: personal.EventFromUser, UserID: "u"}, true},
|
||||
{"bad-rule", personalConsumeOptions{EventKey: personal.EventMention, Rule: "other"}, false},
|
||||
{"bad-filter", personalConsumeOptions{EventKey: personal.EventMention, FilterJSON: "{"}, false},
|
||||
} {
|
||||
@@ -1618,7 +1628,7 @@ func TestCrossPlatformCoveragePersonalSubscriptionAndSourceCoverage(t *testing.T
|
||||
}
|
||||
})
|
||||
}
|
||||
if createCount != 1 {
|
||||
if createCount != 2 {
|
||||
t.Fatalf("create count = %d", createCount)
|
||||
}
|
||||
|
||||
@@ -1640,6 +1650,8 @@ func TestCrossPlatformCoveragePersonalSubscriptionAndSourceCoverage(t *testing.T
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalEventCommandRuntimeCoverage(t *testing.T) {
|
||||
authpkg.SetRuntimeProfile("")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
|
||||
AccessToken: "access", RefreshToken: "refresh", ExpiresAt: time.Now().Add(time.Hour),
|
||||
CorpID: "corp", UserID: "user", ClientID: "client",
|
||||
@@ -2000,6 +2012,31 @@ func TestCrossPlatformCoverageProfileCommandAndModelCoverage(t *testing.T) {
|
||||
if err := authpkg.SaveProfiles(configDir, cfg); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for _, profile := range cfg.Profiles {
|
||||
data := &authpkg.TokenData{
|
||||
AccessToken: "token-" + profile.CorpID,
|
||||
CorpID: profile.CorpID,
|
||||
UserID: profile.UserID,
|
||||
}
|
||||
var err error
|
||||
if profile.UserID == "" {
|
||||
err = authpkg.SaveTokenDataKeychainForCorpID(profile.CorpID, data)
|
||||
} else {
|
||||
err = authpkg.SaveTokenDataKeychainForIdentity(profile.CorpID, profile.UserID, data)
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
for _, profile := range cfg.Profiles {
|
||||
if profile.UserID == "" {
|
||||
_ = authpkg.DeleteTokenDataKeychainForCorpID(profile.CorpID)
|
||||
} else {
|
||||
_ = authpkg.DeleteTokenDataKeychainForIdentity(profile.CorpID, profile.UserID)
|
||||
}
|
||||
}
|
||||
})
|
||||
oldInteractive := profileSwitchInteractiveTerminal
|
||||
oldSelector := profileSwitchSelector
|
||||
t.Cleanup(func() {
|
||||
@@ -2094,15 +2131,11 @@ func TestCrossPlatformCoverageProfileCommandAndModelCoverage(t *testing.T) {
|
||||
}
|
||||
|
||||
_ = profileSwitchSortedProfiles(cfg.Profiles)
|
||||
for _, raw := range []string{"", "bad", now.Format(time.RFC3339)} {
|
||||
_, _ = parseProfileSwitchTime(raw)
|
||||
}
|
||||
for _, p := range cfg.Profiles {
|
||||
_, _ = profileSwitchSortTime(p)
|
||||
_ = profileSwitchOptionLabel(p, cfg)
|
||||
_, _ = profileSwitchProfileCells(p, cfg)
|
||||
}
|
||||
_ = profileSwitchProfileIndex(cfg.Profiles, "missing")
|
||||
_ = profileSwitchProfileIndex(cfg.Profiles, "missing", cfg)
|
||||
_ = profileSwitchBorder("a", "b", "c")
|
||||
_ = profileSwitchTableLine("org", "status")
|
||||
_ = profileSwitchStyledTableLine("org", "status", profileSwitchNormalRowStyle())
|
||||
@@ -2116,9 +2149,9 @@ func TestCrossPlatformCoverageProfileCommandAndModelCoverage(t *testing.T) {
|
||||
_ = profileSwitchTitleStyle()
|
||||
_ = profileSwitchMutedStyle()
|
||||
|
||||
writeProfileListTable(io.Discard, nil)
|
||||
writeProfileListTable(io.Discard, cfg)
|
||||
if err := writeProfileListJSON(io.Discard, cfg); err != nil || writeProfileUseJSON(io.Discard, nil, nil) != nil || writeProfileUseJSON(io.Discard, &cfg.Profiles[0], cfg) != nil {
|
||||
writeProfileListTable(io.Discard, "", nil)
|
||||
writeProfileListTable(io.Discard, "", cfg)
|
||||
if err := writeProfileListJSON(io.Discard, "", cfg); err != nil || writeProfileUseJSON(io.Discard, nil, nil) != nil || writeProfileUseJSON(io.Discard, &cfg.Profiles[0], cfg) != nil {
|
||||
t.Fatal("profile JSON write failed")
|
||||
}
|
||||
_ = profileUseMessage(nil)
|
||||
@@ -2126,8 +2159,8 @@ func TestCrossPlatformCoverageProfileCommandAndModelCoverage(t *testing.T) {
|
||||
_ = profileUseMessage(&p)
|
||||
_ = profileOrgName(p)
|
||||
}
|
||||
_ = profileViews(nil)
|
||||
_ = profileViews(cfg)
|
||||
_ = profileViews("", nil)
|
||||
_ = profileViews("", cfg)
|
||||
for _, limit := range []int{0, 2, 5, 40} {
|
||||
_ = clipProfileCell("abcdefgh", limit)
|
||||
_ = clipProfileDisplayCell("中文abcdefgh", limit)
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -164,7 +165,7 @@ func TestCrossPlatformCoverageRawAPIAndTokenCoverage(t *testing.T) {
|
||||
}
|
||||
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{} }
|
||||
missing := t.TempDir()
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), missing); err != nil || got != "" {
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), missing); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("missing access token = %q, %v", got, err)
|
||||
}
|
||||
if _, err := ResolveAuxiliaryAccessToken(context.Background(), missing, ""); err == nil {
|
||||
|
||||
@@ -56,6 +56,7 @@ var (
|
||||
eventNewEventSource = newEventSource
|
||||
eventNewDingtalkSource = source.New
|
||||
eventResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
eventForceRefreshRejected = forceRefreshRejectedAccessToken
|
||||
eventBusRun = bus.Run
|
||||
eventReadyFDFromEnv = busctl.ReadyFDFromEnv
|
||||
eventResolvePersonal = resolvePersonalEventIdentity
|
||||
@@ -174,6 +175,7 @@ SIGTERM、关 stdin,或先用 dws event stop <subscribe_id> --dry-run 预览
|
||||
"ttl",
|
||||
"ephemeral",
|
||||
"user",
|
||||
"open-dingtalk-id",
|
||||
"group",
|
||||
"personal-event-base-url",
|
||||
); err != nil {
|
||||
@@ -312,7 +314,9 @@ SIGTERM、关 stdin,或先用 dws event stop <subscribe_id> --dry-run 预览
|
||||
"或从外部先用 dws event stop <subscribe_id> --dry-run 预览、确认后加 --yes(会一并退订);"+
|
||||
"请勿 kill -9(会跳过退订、泄漏服务端订阅)")
|
||||
f.StringVar(&personalOpts.UserID, "user", "",
|
||||
"个人单聊对端 userId")
|
||||
"单聊对端或指定发送人的 userId(与 --open-dingtalk-id 二选一)")
|
||||
f.StringVar(&personalOpts.OpenDingTalkID, "open-dingtalk-id", "",
|
||||
"单聊对端或指定发送人的 openDingtalkId(与 --user 二选一)")
|
||||
f.StringVar(&personalOpts.GroupID, "group", "",
|
||||
"group 规则:openConversationId")
|
||||
f.StringVar(&personalOpts.ControlBaseURL, "personal-event-base-url", "",
|
||||
@@ -410,7 +414,7 @@ func eventStreamBusID(streamOpts eventStreamTicketOptions) string {
|
||||
return "portal-ticket-normal:" + sourceID
|
||||
}
|
||||
|
||||
func newEventSource(ctx context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
|
||||
func newEventSource(_ context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
|
||||
if !streamOpts.enabled() {
|
||||
return eventNewDingtalkSource(source.Config{
|
||||
ClientID: clientID,
|
||||
@@ -418,14 +422,6 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
|
||||
})
|
||||
}
|
||||
|
||||
token, err := eventResolveAccessToken(ctx, configDir, "")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("event stream ticket: resolve user token: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(token) == "" {
|
||||
return nil, errors.New("event stream ticket: empty user token")
|
||||
}
|
||||
|
||||
portalClientID := clientID
|
||||
portalClientSecret := clientSecret
|
||||
if streamOpts.usesPortalNormalMode() {
|
||||
@@ -437,8 +433,13 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
|
||||
ClientID: portalClientID,
|
||||
ClientSecret: portalClientSecret,
|
||||
PortalTicket: &source.PortalTicketConfig{
|
||||
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
|
||||
AccessToken: token,
|
||||
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
|
||||
AccessTokenProvider: func(ctx context.Context) (string, error) {
|
||||
return eventResolveAccessToken(ctx, configDir, "")
|
||||
},
|
||||
ForceRefreshToken: func(ctx context.Context, rejectedToken string) (string, error) {
|
||||
return eventForceRefreshRejected(ctx, configDir, rejectedToken)
|
||||
},
|
||||
SourceID: eventStreamSourceID(streamOpts.SourceID),
|
||||
Mode: streamOpts.Mode,
|
||||
ClientID: portalClientID,
|
||||
|
||||
@@ -132,14 +132,18 @@ func TestCrossPlatformCoverageEventSourcesAndForegroundCoverage(t *testing.T) {
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", eventStreamTicketOptions{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
|
||||
stream := eventStreamTicketOptions{Mode: "custom"}
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); !errors.Is(err, fail) {
|
||||
t.Fatalf("stream token error = %v", err)
|
||||
var captured source.Config
|
||||
eventNewDingtalkSource = func(cfg source.Config, _ ...source.SourceOption) (*source.DingtalkSource, error) {
|
||||
captured = cfg
|
||||
return &source.DingtalkSource{}, nil
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return " ", nil }
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err == nil {
|
||||
t.Fatal("empty stream token succeeded")
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err != nil {
|
||||
t.Fatalf("stream source construction = %v", err)
|
||||
}
|
||||
if _, err := captured.PortalTicket.AccessTokenProvider(context.Background()); !errors.Is(err, fail) {
|
||||
t.Fatalf("stream token provider error = %v", err)
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
|
||||
for _, mode := range []string{"custom", "normal"} {
|
||||
|
||||
@@ -0,0 +1,47 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/source"
|
||||
)
|
||||
|
||||
// TestCrossPlatformCoverageNewEventSourceWiresForceRefreshRejectedToken asserts the portal ticket
|
||||
// source receives a ForceRefreshToken callback that forwards the actual
|
||||
// rejected token into the app-level compare-and-refresh chain.
|
||||
func TestCrossPlatformCoverageNewEventSourceWiresForceRefreshRejectedToken(t *testing.T) {
|
||||
oldNew, oldRefresh := eventNewDingtalkSource, eventForceRefreshRejected
|
||||
t.Cleanup(func() { eventNewDingtalkSource, eventForceRefreshRejected = oldNew, oldRefresh })
|
||||
|
||||
var captured source.Config
|
||||
eventNewDingtalkSource = func(cfg source.Config, _ ...source.SourceOption) (*source.DingtalkSource, error) {
|
||||
captured = cfg
|
||||
return &source.DingtalkSource{}, nil
|
||||
}
|
||||
var gotDir, gotRejected string
|
||||
eventForceRefreshRejected = func(_ context.Context, configDir, rejectedToken string) (string, error) {
|
||||
gotDir, gotRejected = configDir, rejectedToken
|
||||
return "fresh", nil
|
||||
}
|
||||
if _, err := newEventSource(context.Background(), "config-dir", "client", "secret", eventStreamTicketOptions{Mode: "custom"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if captured.PortalTicket == nil || captured.PortalTicket.ForceRefreshToken == nil {
|
||||
t.Fatal("ForceRefreshToken not wired into portal ticket config")
|
||||
}
|
||||
tok, err := captured.PortalTicket.ForceRefreshToken(context.Background(), "rejected-token")
|
||||
if err != nil || tok != "fresh" {
|
||||
t.Fatalf("force refresh = %q, %v", tok, err)
|
||||
}
|
||||
if gotDir != "config-dir" || gotRejected != "rejected-token" {
|
||||
t.Fatalf("wiring passed dir %q rejected %q", gotDir, gotRejected)
|
||||
}
|
||||
|
||||
fail := errors.New("refresh failed")
|
||||
eventForceRefreshRejected = func(context.Context, string, string) (string, error) { return "", fail }
|
||||
if _, err := captured.PortalTicket.ForceRefreshToken(context.Background(), "x"); !errors.Is(err, fail) {
|
||||
t.Fatalf("refresh error = %v", err)
|
||||
}
|
||||
}
|
||||
@@ -37,6 +37,7 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/consume"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/source"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/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"
|
||||
@@ -69,6 +70,7 @@ type personalConsumeOptions struct {
|
||||
TTL time.Duration
|
||||
Ephemeral bool
|
||||
UserID string
|
||||
OpenDingTalkID string
|
||||
GroupID string
|
||||
ControlBaseURL string
|
||||
StreamTicketMode string
|
||||
@@ -129,6 +131,7 @@ var (
|
||||
personalFindProcess = os.FindProcess
|
||||
personalSignalProcess = (*os.Process).Signal
|
||||
personalResolveAuxiliaryAccessToken = ResolveAuxiliaryAccessToken
|
||||
personalForceRefreshRejectedToken = forceRefreshRejectedAccessToken
|
||||
personalLoadTokenData = authpkg.LoadTokenData
|
||||
personalClientID = authpkg.ClientID
|
||||
personalResolveAppCredentialsStrict = authpkg.ResolveAppCredentialsStrict
|
||||
@@ -226,8 +229,14 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
if fellback && !opts.Common.Quiet {
|
||||
fmt.Fprintf(c.ErrOrStderr(), "WARN: --format %q has no meaning for event stream; using ndjson\n", rawFormat)
|
||||
}
|
||||
projector := personalEventProjector(opts.DebugRawEvents)
|
||||
|
||||
if opts.Common.DryRun {
|
||||
if strings.TrimSpace(opts.SubscribeID) == "" {
|
||||
if err := validatePersonalSubscriptionOptions(opts); err != nil {
|
||||
return fmt.Errorf("event consume --as user: %w", err)
|
||||
}
|
||||
}
|
||||
cfg := consume.Config{
|
||||
WorkDir: workDir,
|
||||
IPCEndpoint: ipcEndpoint,
|
||||
@@ -240,6 +249,7 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
Format: normalised,
|
||||
OutputDir: opts.Common.OutputDir,
|
||||
Routes: routes,
|
||||
Projector: projector,
|
||||
Stderr: c.ErrOrStderr(),
|
||||
Quiet: opts.Common.Quiet,
|
||||
Foreground: opts.Common.Foreground,
|
||||
@@ -250,7 +260,7 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
return personalConsumeRun(ctx, cfg)
|
||||
}
|
||||
|
||||
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
sub, eventKey, ruleType, err := personalEnsureSubscription(ctx, client, identity, opts)
|
||||
if err != nil {
|
||||
return fmt.Errorf("event consume --as user: %w", err)
|
||||
@@ -272,8 +282,8 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
_ = personalDeleteSubscription(client, context.Background(), sub.SubscribeID)
|
||||
_ = personalRemoveRunStates(workDir, []string{sub.SubscribeID})
|
||||
}
|
||||
// Ownership-based cleanup (AI-subprocess contract, aligned with
|
||||
// lark-cli): a subscription this run CREATED is unsubscribed on exit
|
||||
// Ownership-based cleanup: a subscription this run CREATED is
|
||||
// unsubscribed on exit
|
||||
// (any exit — SIGTERM / stdin-EOF / limit / timeout / error), so nothing
|
||||
// leaks server-side. A subscription REUSED via --subscribe-id is left
|
||||
// intact — the caller owns its lifecycle. --ephemeral forces cleanup
|
||||
@@ -284,22 +294,24 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
}
|
||||
|
||||
cfg := consume.Config{
|
||||
WorkDir: workDir,
|
||||
IPCEndpoint: ipcEndpoint,
|
||||
ClientID: identity.ClientID,
|
||||
SpawnExtraArgs: personalBusSpawnArgs(identity, opts.StreamTicketMode, opts.StreamTicketURL),
|
||||
Compact: opts.Common.Compact,
|
||||
MaxEvents: opts.Common.MaxEvents,
|
||||
Duration: opts.Common.Duration,
|
||||
EventKey: eventKey,
|
||||
Format: normalised,
|
||||
OutputDir: opts.Common.OutputDir,
|
||||
Routes: routes,
|
||||
Stdout: c.OutOrStdout(),
|
||||
Stderr: c.ErrOrStderr(),
|
||||
Quiet: opts.Common.Quiet,
|
||||
Foreground: opts.Common.Foreground,
|
||||
Force: opts.Common.Force,
|
||||
WorkDir: workDir,
|
||||
IPCEndpoint: ipcEndpoint,
|
||||
ClientID: identity.ClientID,
|
||||
SpawnExtraArgs: personalBusSpawnArgs(identity, opts.StreamTicketMode, opts.StreamTicketURL),
|
||||
Compact: opts.Common.Compact,
|
||||
MaxEvents: opts.Common.MaxEvents,
|
||||
Duration: opts.Common.Duration,
|
||||
EventKey: eventKey,
|
||||
Format: normalised,
|
||||
OutputDir: opts.Common.OutputDir,
|
||||
Routes: routes,
|
||||
Projector: projector,
|
||||
ReadySubscribeID: sub.SubscribeID,
|
||||
Stdout: c.OutOrStdout(),
|
||||
Stderr: c.ErrOrStderr(),
|
||||
Quiet: opts.Common.Quiet,
|
||||
Foreground: opts.Common.Foreground,
|
||||
Force: opts.Common.Force,
|
||||
}
|
||||
// Arm the stdin-EOF shutdown watcher only for a pipe-style, unbounded
|
||||
// run (see shouldWatchStdinEOF).
|
||||
@@ -354,6 +366,13 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
return err
|
||||
}
|
||||
|
||||
func personalEventProjector(debugRawEvents bool) consume.Projector {
|
||||
if debugRawEvents {
|
||||
return func(ev transport.Event) (any, error) { return ev, nil }
|
||||
}
|
||||
return personal.ProjectOutput
|
||||
}
|
||||
|
||||
func applyPersonalConsumeFilters(cfg *consume.Config, opts personalConsumeOptions, subscribeID, eventKey string) {
|
||||
if cfg == nil {
|
||||
return
|
||||
@@ -369,6 +388,19 @@ func applyPersonalConsumeFilters(cfg *consume.Config, opts personalConsumeOption
|
||||
cfg.SubscribeID = strings.TrimSpace(subscribeID)
|
||||
}
|
||||
|
||||
func validatePersonalSubscriptionOptions(opts personalConsumeOptions) error {
|
||||
if _, _, err := personal.BuildRuleParam(opts.EventKey, personal.RuleOptions{
|
||||
RuleType: opts.Rule,
|
||||
UserID: opts.UserID,
|
||||
OpenDingTalkID: opts.OpenDingTalkID,
|
||||
GroupID: opts.GroupID,
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
_, _, err := personal.BuildFilter(opts.FilterJSON, opts.QueryCSV)
|
||||
return err
|
||||
}
|
||||
|
||||
func ensurePersonalSubscription(ctx context.Context, client *personal.Client, identity personal.Identity, opts personalConsumeOptions) (*personal.Subscription, string, string, error) {
|
||||
if strings.TrimSpace(opts.SubscribeID) != "" {
|
||||
sub, err := personalGetSubscription(client, ctx, opts.SubscribeID)
|
||||
@@ -398,9 +430,10 @@ func ensurePersonalSubscription(ctx context.Context, client *personal.Client, id
|
||||
return nil, "", "", err
|
||||
}
|
||||
ruleType, ruleParam, err := personal.BuildRuleParam(opts.EventKey, personal.RuleOptions{
|
||||
RuleType: opts.Rule,
|
||||
UserID: opts.UserID,
|
||||
GroupID: opts.GroupID,
|
||||
RuleType: opts.Rule,
|
||||
UserID: opts.UserID,
|
||||
OpenDingTalkID: opts.OpenDingTalkID,
|
||||
GroupID: opts.GroupID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, "", "", err
|
||||
@@ -466,7 +499,7 @@ func runPersonalEventStatus(c *cobra.Command, opts personalStatusOptions) error
|
||||
if status == "" || status == "all" {
|
||||
status = ""
|
||||
}
|
||||
subs, err := personalListSubscriptions(personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
|
||||
subs, err := personalListSubscriptions(newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
|
||||
Status: status,
|
||||
EventKey: opts.EventKey,
|
||||
SubscribeID: opts.SubscribeID,
|
||||
@@ -581,7 +614,7 @@ func runPersonalEventStop(c *cobra.Command, opts personalStopOptions) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("event stop --as user: %w", err)
|
||||
}
|
||||
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
for _, id := range subscribeIDs {
|
||||
if err := personalDeleteSubscription(client, ctx, id); err != nil {
|
||||
return fmt.Errorf("event stop --as user: cancel subscription %s: %w", id, err)
|
||||
@@ -691,7 +724,10 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
|
||||
if err != nil {
|
||||
return personal.Identity{}, err
|
||||
}
|
||||
tokenData, _ := personalLoadTokenData(configDir)
|
||||
tokenData, err := personalLoadTokenData(configDir)
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return personal.Identity{}, fmt.Errorf("load OAuth identity metadata: %w", err)
|
||||
}
|
||||
var corpID, userID, clientID, refreshToken string
|
||||
if tokenData != nil {
|
||||
corpID = tokenData.CorpID
|
||||
@@ -737,6 +773,15 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newPersonalEventControlClient(configDir, baseURL string, identity personal.Identity) *personal.Client {
|
||||
identity.AccessToken = ""
|
||||
client := personal.NewClient(baseURL, identity)
|
||||
client.AccessTokenProvider = func(ctx context.Context) (string, error) {
|
||||
return personalResolveAuxiliaryAccessToken(ctx, configDir, "")
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
func personalTokenSubject(kind, token string) string {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
@@ -785,7 +830,12 @@ func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptio
|
||||
}
|
||||
_ = ctx
|
||||
return source.NewPersonal(source.PersonalConfig{
|
||||
AccessToken: opts.Identity.AccessToken,
|
||||
AccessTokenProvider: func(ctx context.Context) (string, error) {
|
||||
return personalResolveAuxiliaryAccessToken(ctx, opts.ConfigDir, "")
|
||||
},
|
||||
ForceRefreshToken: func(ctx context.Context, rejectedToken string) (string, error) {
|
||||
return personalForceRefreshRejectedToken(ctx, opts.ConfigDir, rejectedToken)
|
||||
},
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
SourceID: opts.Identity.SourceID,
|
||||
@@ -800,15 +850,14 @@ func personalBusSpawnArgs(identity personal.Identity, ticketMode, ticketURL stri
|
||||
"--source-kind", string(dwsevent.SourceKindPersonalStream),
|
||||
"--stream-source-id", identity.SourceID,
|
||||
}
|
||||
// Forward the organization so the detached _bus child resolves
|
||||
// credentials for the SAME profile the parent used. Without this the
|
||||
// child falls back to the default profile's token slot and fails to
|
||||
// authenticate the personal stream for a non-default `--profile`
|
||||
// (symptom: "bus child reported startup failure on ready pipe", no
|
||||
// bus.log). --profile accepts a corpId; the root pre-parses it into the
|
||||
// runtime profile before the _bus handler resolves the identity.
|
||||
// Forward the exact account so the detached _bus child resolves the same
|
||||
// credentials as the parent, including when one organization has multiple
|
||||
// logged-in users.
|
||||
if cid := strings.TrimSpace(identity.CorpID); cid != "" {
|
||||
args = append(args, "--profile", cid)
|
||||
args = append(args, "--profile", authpkg.ProfileSelector(authpkg.Profile{
|
||||
CorpID: identity.CorpID,
|
||||
UserID: identity.UserID,
|
||||
}))
|
||||
}
|
||||
if strings.TrimSpace(ticketMode) != "" {
|
||||
args = append(args, "--stream-ticket-mode", ticketMode)
|
||||
|
||||
@@ -19,10 +19,11 @@ import (
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/consume"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/transport"
|
||||
)
|
||||
|
||||
func TestApplyPersonalConsumeFiltersDebugRawEvents(t *testing.T) {
|
||||
cfg := consume.Config{}
|
||||
cfg := consume.Config{EventKey: personal.EventSingleChat, ReadySubscribeID: "sub-1"}
|
||||
opts := personalConsumeOptions{
|
||||
DebugRawEvents: true,
|
||||
Common: commonConsumeOptions{
|
||||
@@ -34,6 +35,9 @@ func TestApplyPersonalConsumeFiltersDebugRawEvents(t *testing.T) {
|
||||
if cfg.EventTypes != nil || cfg.Filter != "" || cfg.SubscribeID != "" {
|
||||
t.Fatalf("raw debug filters = eventTypes=%#v filter=%q subscribeID=%q, want catch-all", cfg.EventTypes, cfg.Filter, cfg.SubscribeID)
|
||||
}
|
||||
if cfg.EventKey != personal.EventSingleChat || cfg.ReadySubscribeID != "sub-1" {
|
||||
t.Fatalf("raw debug cleared ready identity: eventKey=%q subscribeID=%q", cfg.EventKey, cfg.ReadySubscribeID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyPersonalConsumeFiltersDefault(t *testing.T) {
|
||||
@@ -48,6 +52,28 @@ func TestApplyPersonalConsumeFiltersDefault(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersonalEventProjectorUsesRawEnvelopeForDebug(t *testing.T) {
|
||||
if personalEventProjector(false) == nil {
|
||||
t.Fatal("normal personal consume projector = nil")
|
||||
}
|
||||
projector := personalEventProjector(true)
|
||||
if projector == nil {
|
||||
t.Fatal("debug raw personal consume projector = nil")
|
||||
}
|
||||
ev := transport.Event{
|
||||
EventID: "raw-event",
|
||||
Data: `{"payload":{"uid":100001,"bizid":"internal-bizid"}}`,
|
||||
Headers: map[string]string{"TOPIC": "raw"},
|
||||
}
|
||||
projected, err := projector(ev)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got, ok := projected.(transport.Event); !ok || got.EventID != ev.EventID || got.Data != ev.Data || got.Headers["TOPIC"] != "raw" {
|
||||
t.Fatalf("debug raw projection = %#v", projected)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventConsumeDebugRawEventsRequiresUserMode(t *testing.T) {
|
||||
cmd := newEventConsumeCommand()
|
||||
cmd.SilenceUsage = true
|
||||
@@ -72,7 +98,7 @@ func TestEventConsumeAsAppRejectedBeforeEventKeyValidation(t *testing.T) {
|
||||
|
||||
func TestEventConsumePersonalParamSpecFlags(t *testing.T) {
|
||||
cmd := newEventConsumeCommand()
|
||||
for _, name := range []string{"user", "group", "query"} {
|
||||
for _, name := range []string{"user", "open-dingtalk-id", "group", "query"} {
|
||||
if cmd.Flags().Lookup(name) == nil {
|
||||
t.Fatalf("flag --%s is not registered", name)
|
||||
}
|
||||
@@ -84,6 +110,7 @@ func TestEventConsumePersonalParamSpecFlags(t *testing.T) {
|
||||
"sender-union-id",
|
||||
"open-conversation-id",
|
||||
"keyword",
|
||||
"odid",
|
||||
} {
|
||||
if cmd.Flags().Lookup(name) != nil {
|
||||
t.Fatalf("retired flag --%s is still registered", name)
|
||||
@@ -99,6 +126,7 @@ func TestEventConsumeRetiredPersonalFlagsAreUnknown(t *testing.T) {
|
||||
"sender-union-id",
|
||||
"open-conversation-id",
|
||||
"keyword",
|
||||
"odid",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
cmd := newEventConsumeCommand()
|
||||
@@ -115,7 +143,8 @@ func TestEventConsumeRetiredPersonalFlagsAreUnknown(t *testing.T) {
|
||||
|
||||
func TestEventConsumeAsAppRejectedBeforePersonalParamSpecFlags(t *testing.T) {
|
||||
for _, args := range [][]string{
|
||||
{"--as", "app", "--user", "507971"},
|
||||
{"--as", "app", "--user", "test-user-001"},
|
||||
{"--as", "app", "--open-dingtalk-id", "open-user-1"},
|
||||
{"--as", "app", "--group", "cid"},
|
||||
{"--as", "app", "--query", "报警"},
|
||||
} {
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
|
||||
)
|
||||
|
||||
// TestCrossPlatformCoverageNewPersonalStreamSourceWiresForceRefreshRejectedToken asserts the
|
||||
// personal stream source receives a ForceRefreshToken callback that forwards
|
||||
// the rejected token into the app-level compare-and-refresh chain.
|
||||
func TestCrossPlatformCoverageNewPersonalStreamSourceWiresForceRefreshRejectedToken(t *testing.T) {
|
||||
oldAux := personalResolveAuxiliaryAccessToken
|
||||
oldRefresh := personalForceRefreshRejectedToken
|
||||
t.Cleanup(func() {
|
||||
personalResolveAuxiliaryAccessToken = oldAux
|
||||
personalForceRefreshRejectedToken = oldRefresh
|
||||
})
|
||||
personalResolveAuxiliaryAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "old-token", nil
|
||||
}
|
||||
refreshErr := errors.New("refresh rejected")
|
||||
var gotDir, gotRejected string
|
||||
personalForceRefreshRejectedToken = func(_ context.Context, configDir, rejectedToken string) (string, error) {
|
||||
gotDir, gotRejected = configDir, rejectedToken
|
||||
return "", refreshErr
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
src, err := newPersonalStreamSource(context.Background(), personalStreamSourceOptions{
|
||||
ConfigDir: "config-dir",
|
||||
Identity: personal.Identity{ClientID: "client", SourceID: "source"},
|
||||
TicketURL: srv.URL,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// The 401 ticket response routes the rejected token through the wired
|
||||
// ForceRefreshToken; the unknown refresh failure stays fatal.
|
||||
if err := src.Start(context.Background(), func(*dwsevent.RawEvent) {}); !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("Start() error = %v, want wrapped refresh error", err)
|
||||
}
|
||||
if gotDir != "config-dir" || gotRejected != "old-token" {
|
||||
t.Fatalf("refresh wiring got dir %q rejected %q", gotDir, gotRejected)
|
||||
}
|
||||
}
|
||||
@@ -57,8 +57,8 @@ func TestCrossPlatformCoveragePersonalEventRemainingSchemaAndSubscriptionCoverag
|
||||
personalGetSubscription = func(*personal.Client, context.Context, string) (*personal.Subscription, error) {
|
||||
return &personal.Subscription{EventKey: personal.EventFromUser}, nil
|
||||
}
|
||||
if _, _, _, err := ensurePersonalSubscription(context.Background(), client, personal.Identity{}, personalConsumeOptions{SubscribeID: "sub"}); err == nil {
|
||||
t.Fatal("private subscription event succeeded")
|
||||
if _, key, rule, err := ensurePersonalSubscription(context.Background(), client, personal.Identity{}, personalConsumeOptions{SubscribeID: "sub"}); err != nil || key != personal.EventFromUser || rule != "sender" {
|
||||
t.Fatalf("sender subscription = %q %q, %v", key, rule, err)
|
||||
}
|
||||
personalGetSubscription = func(*personal.Client, context.Context, string) (*personal.Subscription, error) {
|
||||
return &personal.Subscription{EventKey: personal.EventMention}, nil
|
||||
|
||||
@@ -18,7 +18,9 @@ import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event/personal"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
@@ -43,8 +45,18 @@ func TestPersonalEventListHidesSchemaIDs(t *testing.T) {
|
||||
}
|
||||
got := out.String()
|
||||
assertPersonalOutputHidesSchemaIDs(t, got)
|
||||
if strings.Contains(got, personal.EventFromUser) {
|
||||
t.Fatalf("list output exposed hidden event %s: %s", personal.EventFromUser, got)
|
||||
for _, eventKey := range []string{
|
||||
personal.EventFromUser,
|
||||
personal.EventReadO2O,
|
||||
personal.EventReadGroup,
|
||||
personal.EventRecallO2O,
|
||||
personal.EventRecallGroup,
|
||||
personal.EventReactionO2O,
|
||||
personal.EventReactionGroup,
|
||||
} {
|
||||
if !strings.Contains(got, eventKey) {
|
||||
t.Fatalf("list output missing %s: %s", eventKey, got)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -64,8 +76,8 @@ func TestEventListDefaultsToUser(t *testing.T) {
|
||||
if !strings.Contains(got, personal.EventSingleChat) || !strings.Contains(got, "EVENT_KEY") {
|
||||
t.Fatalf("list output = %s, want personal event catalog", got)
|
||||
}
|
||||
if strings.Contains(got, personal.EventFromUser) {
|
||||
t.Fatalf("list output exposed hidden event %s: %s", personal.EventFromUser, got)
|
||||
if !strings.Contains(got, personal.EventFromUser) {
|
||||
t.Fatalf("list output missing public event %s: %s", personal.EventFromUser, got)
|
||||
}
|
||||
if strings.Contains(got, "CLIENT_ID") || strings.Contains(got, "ClientSecret") {
|
||||
t.Fatalf("list default appears to use legacy application output: %s", got)
|
||||
@@ -246,8 +258,8 @@ func TestPersonalEventSchemaUsesSingleJSONSchema(t *testing.T) {
|
||||
t.Fatalf("schema output for %s leaked %q: %s", eventKey, leaked, got)
|
||||
}
|
||||
}
|
||||
if doc["jq_root_path"] != ".data | fromjson" {
|
||||
t.Fatalf("jq_root_path = %#v, want .data | fromjson", doc["jq_root_path"])
|
||||
if doc["jq_root_path"] != "." {
|
||||
t.Fatalf("jq_root_path = %#v, want .", doc["jq_root_path"])
|
||||
}
|
||||
schema, ok := doc["schema"].(map[string]any)
|
||||
if !ok {
|
||||
@@ -264,6 +276,80 @@ func TestPersonalEventSchemaUsesSingleJSONSchema(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersonalActionEventSchemaMatchesFlatOutput(t *testing.T) {
|
||||
tests := []struct {
|
||||
eventKeys []string
|
||||
properties []string
|
||||
}{
|
||||
{
|
||||
eventKeys: []string{personal.EventReadO2O, personal.EventReadGroup},
|
||||
properties: []string{
|
||||
"type", "event_id", "timestamp", "subscribe_id", "message_id",
|
||||
"conversation_id", "reader", "reader_open_dingtalk_id", "sender",
|
||||
"sender_open_dingtalk_id", "read_time", "event_time",
|
||||
},
|
||||
},
|
||||
{
|
||||
eventKeys: []string{personal.EventRecallO2O, personal.EventRecallGroup},
|
||||
properties: []string{
|
||||
"type", "event_id", "timestamp", "subscribe_id", "message_id",
|
||||
"conversation_id", "recaller", "recaller_open_dingtalk_id", "sender",
|
||||
"sender_open_dingtalk_id", "recall_time", "event_time",
|
||||
},
|
||||
},
|
||||
{
|
||||
eventKeys: []string{personal.EventReactionO2O, personal.EventReactionGroup},
|
||||
properties: []string{
|
||||
"type", "event_id", "timestamp", "subscribe_id", "message_id",
|
||||
"conversation_id", "operator", "operator_open_dingtalk_id", "reaction_name",
|
||||
"reaction_text", "operation_type", "operation_time", "sender",
|
||||
"sender_open_dingtalk_id", "event_time",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
for _, eventKey := range tt.eventKeys {
|
||||
t.Run(eventKey, func(t *testing.T) {
|
||||
cmd := newEventSchemaCommand()
|
||||
cmd.SilenceUsage = true
|
||||
cmd.SilenceErrors = true
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetArgs([]string{eventKey})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var doc map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &doc); err != nil {
|
||||
t.Fatalf("schema output is not JSON: %v\n%s", err, out.String())
|
||||
}
|
||||
if doc["event_key"] != eventKey || doc["jq_root_path"] != "." {
|
||||
t.Fatalf("schema metadata = %#v", doc)
|
||||
}
|
||||
schema, ok := doc["schema"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema = %#v", doc["schema"])
|
||||
}
|
||||
properties, ok := schema["properties"].(map[string]any)
|
||||
if !ok || len(properties) != len(tt.properties) {
|
||||
t.Fatalf("schema.properties = %#v, want exactly %d flat fields", schema["properties"], len(tt.properties))
|
||||
}
|
||||
for _, field := range tt.properties {
|
||||
if _, ok := properties[field]; !ok {
|
||||
t.Fatalf("schema missing %q: %#v", field, properties)
|
||||
}
|
||||
}
|
||||
for _, internal := range []string{"payload", "uid", "corpid", "clientId", "filterSubId", "bizid"} {
|
||||
if _, ok := properties[internal]; ok {
|
||||
t.Fatalf("schema exposed internal property %q", internal)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventSchemaDefaultsToUser(t *testing.T) {
|
||||
cmd := newEventSchemaCommand()
|
||||
cmd.SilenceUsage = true
|
||||
@@ -283,38 +369,144 @@ func TestEventSchemaDefaultsToUser(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersonalEventFromUserIsNotPubliclyAvailable(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
cmd *cobra.Command
|
||||
args []string
|
||||
}{
|
||||
{
|
||||
name: "schema",
|
||||
cmd: newEventSchemaCommand(),
|
||||
args: []string{personal.EventFromUser},
|
||||
},
|
||||
{
|
||||
name: "consume",
|
||||
cmd: newEventConsumeCommand(),
|
||||
args: []string{personal.EventFromUser, "--user", "507971", "--dry-run"},
|
||||
},
|
||||
{
|
||||
name: "status",
|
||||
cmd: newEventStatusCommand(),
|
||||
args: []string{"--event", personal.EventFromUser},
|
||||
},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
tc.cmd.SilenceUsage = true
|
||||
tc.cmd.SilenceErrors = true
|
||||
tc.cmd.SetArgs(tc.args)
|
||||
err := tc.cmd.Execute()
|
||||
if err == nil || !strings.Contains(err.Error(), "event "+personal.EventFromUser+" is not publicly available yet") {
|
||||
t.Fatalf("Execute() error = %v, want not publicly available", err)
|
||||
}
|
||||
})
|
||||
func TestPersonalEventFromUserIsPubliclyAvailable(t *testing.T) {
|
||||
if err := ensurePublicPersonalEvent(personal.EventFromUser); err != nil {
|
||||
t.Fatalf("ensurePublicPersonalEvent() error = %v", err)
|
||||
}
|
||||
|
||||
schemaCmd := newEventSchemaCommand()
|
||||
schemaCmd.SilenceUsage = true
|
||||
schemaCmd.SilenceErrors = true
|
||||
var schemaOut bytes.Buffer
|
||||
schemaCmd.SetOut(&schemaOut)
|
||||
schemaCmd.SetArgs([]string{personal.EventFromUser})
|
||||
if err := schemaCmd.Execute(); err != nil {
|
||||
t.Fatalf("schema Execute() error = %v", err)
|
||||
}
|
||||
var doc map[string]any
|
||||
if err := json.Unmarshal(schemaOut.Bytes(), &doc); err != nil {
|
||||
t.Fatalf("schema output is not JSON: %v\n%s", err, schemaOut.String())
|
||||
}
|
||||
if doc["event_key"] != personal.EventFromUser || doc["rule_type"] != "sender" {
|
||||
t.Fatalf("schema document = %#v", doc)
|
||||
}
|
||||
|
||||
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
|
||||
AccessToken: "access-1",
|
||||
RefreshToken: "refresh-1",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
CorpID: "corp-1",
|
||||
UserID: "user-1",
|
||||
ClientID: "client-1",
|
||||
})
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
consumeCmd := newEventConsumeCommand()
|
||||
consumeCmd.SilenceUsage = true
|
||||
consumeCmd.SilenceErrors = true
|
||||
consumeCmd.SetArgs([]string{personal.EventFromUser, "--user", "test-user-001", "--dry-run"})
|
||||
if err := consumeCmd.Execute(); err != nil {
|
||||
t.Fatalf("consume dry-run Execute() error = %v", err)
|
||||
}
|
||||
openIDConsumeCmd := newEventConsumeCommand()
|
||||
openIDConsumeCmd.SilenceUsage = true
|
||||
openIDConsumeCmd.SilenceErrors = true
|
||||
openIDConsumeCmd.SetArgs([]string{personal.EventFromUser, "--open-dingtalk-id", "open-user-1", "--dry-run"})
|
||||
if err := openIDConsumeCmd.Execute(); err != nil {
|
||||
t.Fatalf("consume openDingtalkId dry-run Execute() error = %v", err)
|
||||
}
|
||||
|
||||
conflictingTargetCmd := newEventConsumeCommand()
|
||||
conflictingTargetCmd.SilenceUsage = true
|
||||
conflictingTargetCmd.SilenceErrors = true
|
||||
conflictingTargetCmd.SetArgs([]string{personal.EventFromUser, "--user", "test-user-001", "--open-dingtalk-id", "open-user-1", "--dry-run"})
|
||||
err := conflictingTargetCmd.Execute()
|
||||
if err == nil || !strings.Contains(err.Error(), "--user and --open-dingtalk-id are mutually exclusive for "+personal.EventFromUser) {
|
||||
t.Fatalf("conflicting target identity error = %v", err)
|
||||
}
|
||||
|
||||
groupOpenIDCmd := newEventConsumeCommand()
|
||||
groupOpenIDCmd.SilenceUsage = true
|
||||
groupOpenIDCmd.SilenceErrors = true
|
||||
groupOpenIDCmd.SetArgs([]string{personal.EventInChat, "--group", "cid-1", "--open-dingtalk-id", "open-user-1", "--dry-run"})
|
||||
err = groupOpenIDCmd.Execute()
|
||||
if err == nil || !strings.Contains(err.Error(), "--open-dingtalk-id is not supported for "+personal.EventInChat+"; use --group") {
|
||||
t.Fatalf("group openDingtalkId error = %v", err)
|
||||
}
|
||||
|
||||
missingUserCmd := newEventConsumeCommand()
|
||||
missingUserCmd.SilenceUsage = true
|
||||
missingUserCmd.SilenceErrors = true
|
||||
missingUserCmd.SetArgs([]string{personal.EventFromUser})
|
||||
err = missingUserCmd.Execute()
|
||||
if err == nil || !strings.Contains(err.Error(), "one of --user or --open-dingtalk-id is required for "+personal.EventFromUser) {
|
||||
t.Fatalf("missing target identity error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEventConsumeCobraSchemaIncludesOpenDingTalkID(t *testing.T) {
|
||||
root := NewRootCommand()
|
||||
root.SilenceUsage = true
|
||||
root.SilenceErrors = true
|
||||
var out bytes.Buffer
|
||||
root.SetOut(&out)
|
||||
root.SetArgs([]string{"schema", "event consume"})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("schema event consume Execute() error = %v", err)
|
||||
}
|
||||
var doc map[string]any
|
||||
if err := json.Unmarshal(out.Bytes(), &doc); err != nil {
|
||||
t.Fatalf("schema output is not JSON: %v\n%s", err, out.String())
|
||||
}
|
||||
params, ok := doc["parameters"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema parameters = %#v", doc["parameters"])
|
||||
}
|
||||
if _, ok := params["open-dingtalk-id"]; !ok {
|
||||
t.Fatalf("schema parameters missing open-dingtalk-id: %#v", params)
|
||||
}
|
||||
if _, ok := params["odid"]; ok {
|
||||
t.Fatalf("schema parameters unexpectedly include odid alias: %#v", params)
|
||||
}
|
||||
for _, name := range []string{"user", "open-dingtalk-id", "group"} {
|
||||
param, ok := params[name].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema parameter %s = %#v", name, params[name])
|
||||
}
|
||||
if got, exists := param["required_when"]; exists {
|
||||
t.Fatalf("schema parameter %s unexpectedly declares required_when = %#v", name, got)
|
||||
}
|
||||
}
|
||||
constraints, ok := doc["constraints"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema constraints = %#v", doc["constraints"])
|
||||
}
|
||||
assertJSONConstraintGroup := func(field string, want []string) {
|
||||
t.Helper()
|
||||
groups, ok := constraints[field].([]any)
|
||||
if !ok {
|
||||
t.Fatalf("schema constraint %s = %#v", field, constraints[field])
|
||||
}
|
||||
for _, rawGroup := range groups {
|
||||
group, ok := rawGroup.([]any)
|
||||
if !ok || len(group) != len(want) {
|
||||
continue
|
||||
}
|
||||
matched := true
|
||||
for i := range want {
|
||||
if group[i] != want[i] {
|
||||
matched = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if matched {
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatalf("schema constraint %s = %#v, missing %#v", field, groups, want)
|
||||
}
|
||||
assertJSONConstraintGroup("require_one_of", []string{"event_key", "subscribe-id"})
|
||||
}
|
||||
|
||||
func TestPersonalEventSchemaRejectsTableFormat(t *testing.T) {
|
||||
|
||||
@@ -41,17 +41,18 @@ func TestShouldWatchStdinEOF_BoundedIsNeverArmed(t *testing.T) {
|
||||
func TestPersonalBusSpawnArgs_ForwardsProfile(t *testing.T) {
|
||||
args := personalBusSpawnArgs(personal.Identity{
|
||||
CorpID: "dinga626d60c1128d449",
|
||||
UserID: "user_123",
|
||||
SourceID: "open",
|
||||
}, "", "")
|
||||
found := false
|
||||
for i := 0; i+1 < len(args); i++ {
|
||||
if args[i] == "--profile" && args[i+1] == "dinga626d60c1128d449" {
|
||||
if args[i] == "--profile" && args[i+1] == "dinga626d60c1128d449:user_123" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("spawn args must forward --profile <corpId>; got %v", args)
|
||||
t.Errorf("spawn args must forward --profile <corpId>:<userId>; got %v", args)
|
||||
}
|
||||
|
||||
// No CorpID → no --profile appended (avoid an empty flag value).
|
||||
|
||||
@@ -47,7 +47,7 @@ func bindPersistentFlags(cmd *cobra.Command, flags *GlobalFlags) {
|
||||
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")
|
||||
cmd.PersistentFlags().StringVar(&flags.Profile, "profile", "", "一次性指定本次命令使用的组织 profile 名或 corpId;多个按 CSV 逗号分隔,如 corpA,corpB")
|
||||
cmd.PersistentFlags().StringVar(&flags.Profile, "profile", "", "一次性指定组织或账号;支持 corpId/corpName 与 userId/userName 组合,推荐使用 profile list 返回的 corpId:userId;多个按 CSV 逗号分隔")
|
||||
cmd.PersistentFlags().IntVar(&flags.Timeout, "timeout", 30, "HTTP 请求超时时间 (秒)")
|
||||
cmd.PersistentFlags().StringVar(&flags.Token, "token", "", "Override the configured API token")
|
||||
_ = cmd.PersistentFlags().MarkHidden("token")
|
||||
|
||||
@@ -27,9 +27,13 @@ type accessTokenGetter interface {
|
||||
GetAccessToken(context.Context) (string, error)
|
||||
}
|
||||
|
||||
type rejectedAccessTokenRefresher interface {
|
||||
ForceRefreshRejectedToken(context.Context, string) (string, error)
|
||||
}
|
||||
|
||||
var (
|
||||
markAccessTokenStale = authpkg.MarkAccessTokenStale
|
||||
newRefreshProvider = func(configDir string) accessTokenGetter {
|
||||
loadRefreshTokenData = authpkg.LoadTokenData
|
||||
newRefreshProvider = func(configDir string) rejectedAccessTokenRefresher {
|
||||
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, disc)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
@@ -42,26 +46,33 @@ var (
|
||||
// server-side rejection (HTTP 401 or business code such as
|
||||
// TOKEN_VERIFIED_FAILED) on what locally appeared to be a still-valid token.
|
||||
//
|
||||
// Steps:
|
||||
// 1. MarkAccessTokenStale rewrites ExpiresAt to a past instant so
|
||||
// OAuthProvider.GetAccessToken's fast-path will miss.
|
||||
// 2. NewOAuthProvider + GetAccessToken triggers lockedRefresh, which uses the
|
||||
// existing dual-layer lock (process + file) to serialize concurrent
|
||||
// refresh attempts across goroutines and processes.
|
||||
// 3. ResetRuntimeTokenCache clears the per-process sync.Once cache so the
|
||||
// next resolveAuthToken call re-reads from disk.
|
||||
//
|
||||
// Existing OAuthProvider.GetAccessToken behaviour is unchanged; this helper
|
||||
// is the only entry point that orchestrates "force refresh" semantics.
|
||||
// It snapshots the current access token, then delegates to the OAuth
|
||||
// provider's dual-locked compare-and-refresh operation. If another caller has
|
||||
// already rotated the token, that newer token is reused without another
|
||||
// refresh request.
|
||||
func ForceRefreshAccessToken(ctx context.Context, configDir string) (string, error) {
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if err := markAccessTokenStale(configDir); err != nil {
|
||||
return "", fmt.Errorf("mark access token stale: %w", err)
|
||||
data, err := loadRefreshTokenData(configDir)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if data == nil || strings.TrimSpace(data.AccessToken) == "" {
|
||||
return "", fmt.Errorf("stored access token is empty")
|
||||
}
|
||||
return forceRefreshRejectedAccessToken(ctx, configDir, data.AccessToken)
|
||||
}
|
||||
|
||||
func forceRefreshRejectedAccessToken(ctx context.Context, configDir, rejectedAccessToken string) (string, error) {
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if strings.TrimSpace(rejectedAccessToken) == "" {
|
||||
return "", fmt.Errorf("rejected access token is empty")
|
||||
}
|
||||
provider := newRefreshProvider(configDir)
|
||||
tok, err := provider.GetAccessToken(ctx)
|
||||
tok, err := provider.ForceRefreshRejectedToken(ctx, rejectedAccessToken)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
+10
-4
@@ -27,21 +27,27 @@ import (
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
func newLegacyPublicCommands(runner executor.Runner, caller edition.ToolCaller) []*cobra.Command {
|
||||
func newLegacyPublicCommands(runner executor.Runner, caller edition.ToolCaller, loadUserShortcuts bool) []*cobra.Command {
|
||||
injectStaticServers()
|
||||
helpers.InitDeps(caller)
|
||||
commands := helpers.NewPublicCommands(runner)
|
||||
// Load user-defined shortcuts (~/.dws/shortcuts/*.yaml) BEFORE compiling the
|
||||
// command tree, so distilled high-frequency operations mount alongside the
|
||||
// built-ins. Conflicts with built-ins are skipped inside Load.
|
||||
if _, err := userdef.Load(); err != nil {
|
||||
slog.Warn("shortcut: failed to load user-defined shortcuts", "error", err)
|
||||
if loadUserShortcuts {
|
||||
if _, err := userdef.Load(); err != nil {
|
||||
slog.Warn("shortcut: failed to load user-defined shortcuts", "error", err)
|
||||
}
|
||||
}
|
||||
// Built-in + user shortcuts (`dws <service> +<command>`) share the same
|
||||
// command tree; mergeTopLevelCommands folds each shortcut's service parent
|
||||
// into the matching helper command so the `+leaf` sits alongside existing
|
||||
// subcommands.
|
||||
commands = append(commands, builtin.Commands()...)
|
||||
if loadUserShortcuts {
|
||||
commands = append(commands, builtin.Commands()...)
|
||||
} else {
|
||||
commands = append(commands, builtin.BaseCommands()...)
|
||||
}
|
||||
return mergeTopLevelCommands(commands)
|
||||
}
|
||||
|
||||
|
||||
@@ -50,8 +50,9 @@ func TestRuntimeRunnerAggregatesCommaSeparatedProfiles(t *testing.T) {
|
||||
t.Fatalf("profiles[%d].ok = %#v, want true", i, entry["ok"])
|
||||
}
|
||||
resultPayload := entry["result"].(map[string]any)
|
||||
if resultPayload["runtimeProfile"] != wantCorpID {
|
||||
t.Fatalf("profiles[%d].result.runtimeProfile = %#v, want %q", i, resultPayload["runtimeProfile"], wantCorpID)
|
||||
wantProfile := wantCorpID + ":user-" + wantCorpID
|
||||
if resultPayload["runtimeProfile"] != wantProfile {
|
||||
t.Fatalf("profiles[%d].result.runtimeProfile = %#v, want %q", i, resultPayload["runtimeProfile"], wantProfile)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -75,6 +76,37 @@ func TestRuntimeRunnerDeduplicatesCommaSeparatedProfilesByCorpID(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerDeduplicatesByResolvedIdentityInSameCorp(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_same")
|
||||
first.UserID = "user_1"
|
||||
second := authLogoutTestToken("corp_same")
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
configDir := setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
selections, multi, err := resolveMultiProfileSelections(
|
||||
configDir,
|
||||
"corp_same,corp_same:user_1,corp_same:user_2",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("resolveMultiProfileSelections() error = %v", err)
|
||||
}
|
||||
if !multi {
|
||||
t.Fatal("multi = false, want true")
|
||||
}
|
||||
if len(selections) != 2 {
|
||||
t.Fatalf("selections len = %d, want 2: %#v", len(selections), selections)
|
||||
}
|
||||
got := []string{
|
||||
authpkg.ProfileSelector(selections[0].Profile),
|
||||
authpkg.ProfileSelector(selections[1].Profile),
|
||||
}
|
||||
if strings.Join(got, ",") != "corp_same:user_2,corp_same:user_1" {
|
||||
t.Fatalf("resolved identities = %v, want current user_2 then user_1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerKeepsSingleProfileBehavior(t *testing.T) {
|
||||
setupAuthLogoutProfiles(t, authLogoutTestToken("corp_a"), authLogoutTestToken("corp_b"))
|
||||
authpkg.SetRuntimeProfile("corp_a")
|
||||
@@ -91,11 +123,38 @@ func TestRuntimeRunnerKeepsSingleProfileBehavior(t *testing.T) {
|
||||
if _, ok := result.Response["content"].(map[string]any)["multiProfile"]; ok {
|
||||
t.Fatalf("single profile unexpectedly returned aggregate content: %#v", result.Response)
|
||||
}
|
||||
if got := result.Response["content"].(map[string]any)["runtimeProfile"]; got != "corp_a:user-corp_a" {
|
||||
t.Fatalf("fallback runtime profile = %#v, want exact identity selector", got)
|
||||
}
|
||||
if got := authpkg.RuntimeProfile(); got != "corp_a" {
|
||||
t.Fatalf("runtime profile after Run = %q, want corp_a", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuntimeRunnerRejectsAmbiguousSingleProfile(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_first")
|
||||
first.CorpName = "Shared Org"
|
||||
second := authLogoutTestToken("corp_second")
|
||||
second.CorpName = "Shared Org"
|
||||
setupAuthLogoutProfiles(t, first, second)
|
||||
authpkg.SetRuntimeProfile("Shared Org")
|
||||
|
||||
runner := &runtimeRunner{fallback: multiProfileFallbackRunner{}}
|
||||
_, err := runner.Run(context.Background(), executor.Invocation{
|
||||
Kind: "helper_invocation",
|
||||
CanonicalProduct: "contact",
|
||||
Tool: "get_current_user_profile",
|
||||
})
|
||||
if err == nil {
|
||||
t.Fatal("Run() accepted ambiguous single profile selector")
|
||||
}
|
||||
for _, candidate := range []string{"corp_first", "corp_second"} {
|
||||
if !strings.Contains(err.Error(), candidate) {
|
||||
t.Fatalf("error = %q, want candidate %q", err.Error(), candidate)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCommaNamedProfileStillResolvesAsSingleProfile(t *testing.T) {
|
||||
configDir := setupAuthLogoutProfiles(t, authLogoutTestToken("corp_comma"), authLogoutTestToken("corp_other"))
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
|
||||
@@ -56,7 +56,7 @@ var openBrowserFunc = tryOpenBrowser
|
||||
var (
|
||||
patAuthorizationTimeout = PatAuthRetryTimeout
|
||||
patAuthorizationPollInterval = PatAuthPollInterval
|
||||
patLoadTokenData = authpkg.LoadTokenData
|
||||
patResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
patWaitForAuthorization = WaitForPatAuthorization
|
||||
patPollDeviceFlowWithInterval = pollPatDeviceFlowWithInterval
|
||||
patSaveAppConfig = authpkg.SaveAppConfig
|
||||
@@ -272,7 +272,7 @@ func patAuthorizationURIFromData(data map[string]any) string {
|
||||
|
||||
// WaitForPatAuthorization polls until the user completes authorization or timeout.
|
||||
// It returns true if authorization was completed, false if timed out or cancelled.
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) (bool, error) {
|
||||
timeout := patAuthorizationTimeout
|
||||
deadline := time.Now().Add(timeout)
|
||||
pollTicker := time.NewTicker(patAuthorizationPollInterval)
|
||||
@@ -290,27 +290,26 @@ func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Wr
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fmt.Fprintf(output, "%s 操作已取消\n", tui.StateMark("error"))
|
||||
return false
|
||||
return false, ctx.Err()
|
||||
|
||||
case <-time.After(time.Until(deadline)):
|
||||
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", tui.StateMark("error"), timeout)
|
||||
fmt.Fprintf(output, " %s 请重新执行命令\n", tui.Dim("ℹ"))
|
||||
return false
|
||||
return false, nil
|
||||
|
||||
case <-pollTicker.C:
|
||||
pollCount++
|
||||
elapsed := time.Since(start).Truncate(time.Second)
|
||||
remaining := time.Until(deadline).Truncate(time.Second)
|
||||
|
||||
// Check if token is now valid
|
||||
tokenData, err := patLoadTokenData(configDir)
|
||||
if err == nil && tokenData != nil {
|
||||
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true
|
||||
}
|
||||
// Check the same resolver used by every outbound bearer request.
|
||||
if _, err := patResolveAccessToken(ctx, configDir, ""); err == nil {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true, nil
|
||||
} else if !stderrors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return false, fmt.Errorf("check authorization token: %w", err)
|
||||
}
|
||||
|
||||
// Show polling status
|
||||
@@ -340,7 +339,10 @@ func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocati
|
||||
PrintPatAuthError(output, scopeErr)
|
||||
|
||||
// Wait for user to complete authorization
|
||||
authorized := patWaitForAuthorization(ctx, configDir, output)
|
||||
authorized, waitErr := patWaitForAuthorization(ctx, configDir, output)
|
||||
if waitErr != nil {
|
||||
return executor.Result{}, waitErr
|
||||
}
|
||||
if !authorized {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"等待用户授权超时",
|
||||
@@ -794,12 +796,6 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
|
||||
pollURL := fmt.Sprintf("%s%s?flowId=%s",
|
||||
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
|
||||
|
||||
// Load user access token for the poll request header.
|
||||
var accessToken string
|
||||
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
|
||||
accessToken = tokenData.AccessToken
|
||||
}
|
||||
|
||||
// Use a client that does NOT follow redirects, so we can detect SSO 302.
|
||||
noRedirectClient := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
@@ -828,6 +824,10 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
|
||||
slog.Debug("PAT poll: failed to create request", "error", err)
|
||||
continue
|
||||
}
|
||||
accessToken, tokenErr := patResolveAccessToken(ctx, configDir, "")
|
||||
if tokenErr != nil && !stderrors.Is(tokenErr, authpkg.ErrTokenDataNotFound) {
|
||||
return "", "", fmt.Errorf("resolve PAT poll access token: %w", tokenErr)
|
||||
}
|
||||
if accessToken != "" {
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
}
|
||||
|
||||
@@ -52,39 +52,41 @@ func TestCrossPlatformCoveragePATRetryRemainingPureAndWaitCoverage(t *testing.T)
|
||||
|
||||
oldTimeout := patAuthorizationTimeout
|
||||
oldInterval := patAuthorizationPollInterval
|
||||
oldLoad := patLoadTokenData
|
||||
oldResolve := patResolveAccessToken
|
||||
t.Cleanup(func() {
|
||||
patAuthorizationTimeout = oldTimeout
|
||||
patAuthorizationPollInterval = oldInterval
|
||||
patLoadTokenData = oldLoad
|
||||
patResolveAccessToken = oldResolve
|
||||
})
|
||||
patAuthorizationTimeout = 50 * time.Millisecond
|
||||
patAuthorizationPollInterval = time.Millisecond
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "token", nil
|
||||
}
|
||||
out.Reset()
|
||||
if !WaitForPatAuthorization(context.Background(), "", &out) {
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || !ok {
|
||||
t.Fatal("valid token did not authorize")
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(ctx, "", &out) {
|
||||
t.Fatal("cancelled authorization succeeded")
|
||||
if ok, err := WaitForPatAuthorization(ctx, "", &out); ok || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("cancelled authorization = %v, %v", ok, err)
|
||||
}
|
||||
patAuthorizationTimeout = time.Millisecond
|
||||
patAuthorizationPollInterval = time.Hour
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(context.Background(), "", &out) {
|
||||
t.Fatal("timed out authorization succeeded")
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok {
|
||||
t.Fatalf("timed out authorization = %v, %v", ok, err)
|
||||
}
|
||||
patAuthorizationTimeout = 5 * time.Millisecond
|
||||
patAuthorizationPollInterval = time.Millisecond
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(context.Background(), "", &out) || !strings.Contains(out.String(), "等待授权中") {
|
||||
t.Fatalf("invalid-token polling output = %q", out.String())
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok || !strings.Contains(out.String(), "等待授权中") {
|
||||
t.Fatalf("invalid-token polling = %v, %v, output %q", ok, err, out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,12 +111,12 @@ func TestCrossPlatformCoveragePATRetryRemainingOrchestrationCoverage(t *testing.
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
scope := &PatScopeError{OriginalError: "missing", Identity: "user", ErrorType: "missing_scope", Message: "missing", Hint: "login", MissingScope: "calendar:read"}
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return false }
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return false, nil }
|
||||
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{}, executor.Invocation{}, scope, t.TempDir(), io.Discard); err == nil {
|
||||
t.Fatal("PAT retry timeout succeeded")
|
||||
}
|
||||
wantErr := errors.New("runner failed")
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return true }
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return true, nil }
|
||||
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{err: wantErr}, executor.Invocation{}, scope, t.TempDir(), io.Discard); !errors.Is(err, wantErr) {
|
||||
t.Fatalf("authorized retry = %v", err)
|
||||
}
|
||||
@@ -215,15 +217,15 @@ func patRaw(flowID, clientID, secret string) string {
|
||||
func TestCrossPlatformCoveragePATRetryRemainingPollAndBrowserCoverage(t *testing.T) {
|
||||
oldDo := patPollHTTPDo
|
||||
oldRequest := patPollNewRequest
|
||||
oldLoad := patLoadTokenData
|
||||
oldResolve := patResolveAccessToken
|
||||
oldBrowser := patBrowserOpenCommand
|
||||
t.Cleanup(func() {
|
||||
patPollHTTPDo = oldDo
|
||||
patPollNewRequest = oldRequest
|
||||
patLoadTokenData = oldLoad
|
||||
patResolveAccessToken = oldResolve
|
||||
patBrowserOpenCommand = oldBrowser
|
||||
})
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return &authpkg.TokenData{AccessToken: "token"}, nil }
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
|
||||
cancelled, cancelNow := context.WithCancel(context.Background())
|
||||
cancelNow()
|
||||
if status, _, err := pollPatDeviceFlowWithInterval(cancelled, "flow", t.TempDir(), io.Discard, 0); err != nil || status != authpkg.StatusCancelled {
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,675 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/cobracmd"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
type pluginFailRunner struct{}
|
||||
|
||||
func (pluginFailRunner) Run(context.Context, executor.Invocation) (executor.Result, error) {
|
||||
return executor.Result{}, errors.New("runner failed")
|
||||
}
|
||||
|
||||
type pluginWrongFlagValue struct{}
|
||||
|
||||
func (pluginWrongFlagValue) String() string { return "" }
|
||||
func (pluginWrongFlagValue) Set(string) error { return nil }
|
||||
func (pluginWrongFlagValue) Type() string { return "wrong" }
|
||||
|
||||
func TestPluginCompilerRejectsInvalidDuplicateAndEmptyDefinitions(t *testing.T) {
|
||||
invalidRoot := conferencePluginDescriptor()
|
||||
invalidRoot.CLI.Command = "Invalid Root"
|
||||
if commands := buildPluginCommands([]mcptypes.ServerDescriptor{invalidRoot}, executor.EchoRunner{}, nil); len(commands) != 0 {
|
||||
t.Fatalf("invalid root produced commands %#v", commands)
|
||||
}
|
||||
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.Groups = map[string]mcptypes.CLIGroupDef{
|
||||
"empty": {Description: "removed when no leaf survives"},
|
||||
}
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"": {CLIName: "blank-tool"},
|
||||
"hidden": {CLIName: "hidden", Hidden: true},
|
||||
"invalid": {CLIName: "Invalid Leaf"},
|
||||
"first": {CLIName: "same"},
|
||||
"second": {CLIName: "same"},
|
||||
}
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{descriptor}, executor.EchoRunner{}, nil)
|
||||
if len(commands) != 1 {
|
||||
t.Fatalf("commands = %#v", commands)
|
||||
}
|
||||
if requireOptionalPluginChild(commands[0], "same") == nil {
|
||||
t.Fatal("valid leaf was not retained")
|
||||
}
|
||||
if requireOptionalPluginChild(commands[0], "empty") != nil {
|
||||
t.Fatal("empty group was not pruned")
|
||||
}
|
||||
|
||||
empty := conferencePluginDescriptor()
|
||||
empty.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"hidden": {CLIName: "hidden", Hidden: true},
|
||||
}
|
||||
if commands := buildPluginCommands([]mcptypes.ServerDescriptor{empty}, executor.EchoRunner{}, nil); len(commands) != 0 {
|
||||
t.Fatalf("empty overlay produced commands %#v", commands)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginLeafExecutionErrorsAndBodyWrapper(t *testing.T) {
|
||||
base := conferencePluginDescriptor()
|
||||
base.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"wrapped": {
|
||||
CLIName: "wrapped",
|
||||
BodyWrapper: "body",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"value": {Required: true},
|
||||
},
|
||||
},
|
||||
}
|
||||
runner := &pluginCaptureRunner{}
|
||||
root := pluginTestRoot(buildPluginCommands([]mcptypes.ServerDescriptor{base}, runner, nil)...)
|
||||
root.SetArgs([]string{"conference", "wrapped", "--value", "ok", "--params", `{"body":{"old":1},"_meta":"kept"}`})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("wrapped command: %v", err)
|
||||
}
|
||||
want := map[string]any{
|
||||
"_meta": "kept",
|
||||
"body": map[string]any{"old": float64(1), "value": "ok"},
|
||||
}
|
||||
if !reflect.DeepEqual(runner.invocations[0].Params, want) {
|
||||
t.Fatalf("wrapped params = %#v, want %#v", runner.invocations[0].Params, want)
|
||||
}
|
||||
|
||||
for _, testCase := range []struct {
|
||||
name string
|
||||
runner executor.Runner
|
||||
args []string
|
||||
}{
|
||||
{name: "invalid json", runner: executor.EchoRunner{}, args: []string{"conference", "wrapped", "--json", "["}},
|
||||
{name: "missing required", runner: executor.EchoRunner{}, args: []string{"conference", "wrapped"}},
|
||||
{name: "missing runner", runner: nil, args: []string{"conference", "wrapped", "--value", "ok"}},
|
||||
{name: "runner error", runner: pluginFailRunner{}, args: []string{"conference", "wrapped", "--value", "ok"}},
|
||||
} {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
commandRoot := pluginTestRoot(buildPluginCommands([]mcptypes.ServerDescriptor{base}, testCase.runner, nil)...)
|
||||
commandRoot.SetArgs(testCase.args)
|
||||
if err := commandRoot.Execute(); err == nil {
|
||||
t.Fatal("expected command error")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for _, flagName := range []string{"json", "params"} {
|
||||
t.Run("unreadable "+flagName, func(t *testing.T) {
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{base}, executor.EchoRunner{}, nil)
|
||||
leaf := requirePluginChild(t, commands[0], "wrapped")
|
||||
leaf.Flags().Lookup(flagName).Value = pluginWrongFlagValue{}
|
||||
commandRoot := pluginTestRoot(commands...)
|
||||
commandRoot.SetArgs([]string{"conference", "wrapped", "--value", "ok"})
|
||||
if err := commandRoot.Execute(); err == nil {
|
||||
t.Fatal("expected unreadable flag error")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginBindingCompilerCoversAliasesAndPositionalValidators(t *testing.T) {
|
||||
reservations := pluginFlagReservations{
|
||||
names: map[string]bool{"reserved": true},
|
||||
shorthands: map[string]bool{},
|
||||
}
|
||||
bindings, _, _, ok := registerPluginBindings("alias", mcptypes.CLIToolOverride{
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"value": {Alias: "value", Aliases: []string{"", "Bad", "value", "other"}},
|
||||
},
|
||||
}, reservations)
|
||||
if !ok || !reflect.DeepEqual(bindings[0].names, []string{"value", "other"}) {
|
||||
t.Fatalf("alias bindings = (%#v, %v)", bindings, ok)
|
||||
}
|
||||
if _, _, _, ok := registerPluginBindings("conflict", mcptypes.CLIToolOverride{
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{"value": {Alias: "reserved"}},
|
||||
}, reservations); ok {
|
||||
t.Fatal("reserved flag was accepted")
|
||||
}
|
||||
if _, _, _, ok := registerPluginBindings("negative", mcptypes.CLIToolOverride{
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{"value": {Positional: true, PositionalIndex: -1}},
|
||||
}, reservations); ok {
|
||||
t.Fatal("negative positional index was accepted")
|
||||
}
|
||||
if _, _, _, ok := registerPluginBindings("duplicate", mcptypes.CLIToolOverride{
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"first": {Positional: true, PositionalIndex: 0},
|
||||
"second": {Positional: true, PositionalIndex: 0},
|
||||
},
|
||||
}, reservations); ok {
|
||||
t.Fatal("duplicate positional index was accepted")
|
||||
}
|
||||
if _, _, _, ok := registerPluginBindings("gap", mcptypes.CLIToolOverride{
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"second": {Positional: true, PositionalIndex: 1},
|
||||
},
|
||||
}, reservations); ok {
|
||||
t.Fatal("non-contiguous positional indexes were accepted")
|
||||
}
|
||||
|
||||
for _, testCase := range []struct {
|
||||
name string
|
||||
flags map[string]mcptypes.CLIFlagOverride
|
||||
wantUse string
|
||||
valid []string
|
||||
invalid []string
|
||||
}{
|
||||
{
|
||||
name: "exact",
|
||||
flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"second": {Positional: true, PositionalIndex: 1, Required: true},
|
||||
"first": {Positional: true, PositionalIndex: 0, Required: true},
|
||||
},
|
||||
wantUse: "exact [first] [second]", valid: []string{"a", "b"}, invalid: []string{"a"},
|
||||
},
|
||||
{
|
||||
name: "range",
|
||||
flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"first": {Positional: true, PositionalIndex: 0, Required: true},
|
||||
"second": {Positional: true, PositionalIndex: 1},
|
||||
},
|
||||
wantUse: "range [first] [second]", valid: []string{"a"}, invalid: []string{},
|
||||
},
|
||||
{
|
||||
name: "maximum",
|
||||
flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"first": {Positional: true, PositionalIndex: 0},
|
||||
"second": {Positional: true, PositionalIndex: 1},
|
||||
},
|
||||
wantUse: "maximum [first] [second]", valid: []string{}, invalid: []string{"a", "b", "c"},
|
||||
},
|
||||
} {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
_, use, validator, ok := registerPluginBindings(testCase.name, mcptypes.CLIToolOverride{Flags: testCase.flags}, reservations)
|
||||
if !ok || use != testCase.wantUse {
|
||||
t.Fatalf("binding contract = (%q, %v)", use, ok)
|
||||
}
|
||||
cmd := &cobra.Command{Use: testCase.name}
|
||||
if err := validator(cmd, testCase.valid); err != nil {
|
||||
t.Fatalf("valid args: %v", err)
|
||||
}
|
||||
if err := validator(cmd, testCase.invalid); err == nil {
|
||||
t.Fatal("invalid args were accepted")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginFlagRegistrationAndReadingCoversAllKinds(t *testing.T) {
|
||||
cmd := &cobra.Command{Use: "leaf"}
|
||||
override := mcptypes.CLIToolOverride{Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"integer": {Default: "2", Shorthand: "i", Hidden: true},
|
||||
"float": {Default: "1.5"},
|
||||
"boolean": {Default: "true"},
|
||||
"slice": {Default: "one, ,two"},
|
||||
"json": {Default: `{"old":true}`},
|
||||
"string": {Default: "text"},
|
||||
}}
|
||||
bindings := []pluginFlagBinding{
|
||||
{property: "integer", names: []string{"integer", "integer-alias"}, kind: pluginFlagInt},
|
||||
{property: "float", names: []string{"float"}, kind: pluginFlagFloat},
|
||||
{property: "boolean", names: []string{"boolean"}, kind: pluginFlagBool},
|
||||
{property: "slice", names: []string{"slice"}, kind: pluginFlagStringSlice},
|
||||
{property: "json", names: []string{"json-value"}, kind: pluginFlagJSON},
|
||||
{property: "string", names: []string{"string"}, kind: pluginFlagString},
|
||||
}
|
||||
registerPluginFlags(cmd, bindings, override, pluginFlagReservations{shorthands: map[string]bool{}})
|
||||
for name, raw := range map[string]string{
|
||||
"integer": "3", "float": "2.5", "boolean": "false",
|
||||
"slice": "three,four", "json-value": `{"ok":true}`, "string": "changed",
|
||||
} {
|
||||
if err := cmd.Flags().Set(name, raw); err != nil {
|
||||
t.Fatalf("set --%s: %v", name, err)
|
||||
}
|
||||
}
|
||||
wants := map[string]any{
|
||||
"integer": 3,
|
||||
"float": 2.5,
|
||||
"boolean": false,
|
||||
"slice": []string{"three", "four"},
|
||||
"json": map[string]any{"ok": true},
|
||||
"string": "changed",
|
||||
}
|
||||
for _, binding := range bindings {
|
||||
value, err := readPluginFlag(cmd.Flags(), binding.names[0], binding.kind)
|
||||
if err != nil || !reflect.DeepEqual(value, wants[binding.property]) {
|
||||
t.Fatalf("read %s = (%#v, %v), want %#v", binding.property, value, err, wants[binding.property])
|
||||
}
|
||||
}
|
||||
if !cmd.Flags().Lookup("integer").Hidden || !cmd.Flags().Lookup("integer-alias").Hidden {
|
||||
t.Fatal("hidden primary or alias flag was exposed")
|
||||
}
|
||||
if err := cmd.Flags().Set("json-value", "{"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := readPluginFlag(cmd.Flags(), "json-value", pluginFlagJSON); err == nil {
|
||||
t.Fatal("invalid JSON flag was accepted")
|
||||
}
|
||||
cmd.Flags().Lookup("json-value").Value = pluginWrongFlagValue{}
|
||||
if _, err := readPluginFlag(cmd.Flags(), "json-value", pluginFlagJSON); err == nil {
|
||||
t.Fatal("wrong JSON flag type was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectPluginBindingsCoversEveryValueSourceAndFailure(t *testing.T) {
|
||||
t.Run("sources", func(t *testing.T) {
|
||||
cmd := &cobra.Command{Use: "leaf"}
|
||||
registerPluginFlag(cmd.Flags(), "flag", "", "", pluginFlagString, "")
|
||||
if err := cmd.Flags().Set("flag", "from-flag"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("PLUGIN_COVERAGE_ENV", "7")
|
||||
params := map[string]any{"existing": "from-json"}
|
||||
bindings := []pluginFlagBinding{
|
||||
{property: "flag", names: []string{"flag"}, kind: pluginFlagString},
|
||||
{property: "existing", kind: pluginFlagString},
|
||||
{property: "positional", kind: pluginFlagBool, positional: true, positionalIndex: 0},
|
||||
{property: "default", kind: pluginFlagFloat, defaultProvided: true, defaultValue: "1.5"},
|
||||
{property: "env", kind: pluginFlagInt, envDefault: "PLUGIN_COVERAGE_ENV"},
|
||||
{property: "optional", kind: pluginFlagString},
|
||||
}
|
||||
if err := collectPluginBindings(cmd, []string{"true"}, bindings, params); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want := map[string]any{
|
||||
"flag": "from-flag", "existing": "from-json", "positional": true,
|
||||
"default": 1.5, "env": 7,
|
||||
}
|
||||
if !reflect.DeepEqual(params, want) {
|
||||
t.Fatalf("params = %#v, want %#v", params, want)
|
||||
}
|
||||
})
|
||||
|
||||
for _, testCase := range []struct {
|
||||
name string
|
||||
prepare func(t *testing.T, cmd *cobra.Command)
|
||||
args []string
|
||||
binding pluginFlagBinding
|
||||
params map[string]any
|
||||
}{
|
||||
{
|
||||
name: "wrong flag type",
|
||||
prepare: func(t *testing.T, cmd *cobra.Command) {
|
||||
cmd.Flags().String("value", "", "")
|
||||
if err := cmd.Flags().Set("value", "x"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
},
|
||||
binding: pluginFlagBinding{property: "value", names: []string{"value"}, kind: pluginFlagInt},
|
||||
},
|
||||
{name: "invalid positional", args: []string{"maybe"}, binding: pluginFlagBinding{property: "value", kind: pluginFlagBool, positional: true, positionalIndex: 0}},
|
||||
{name: "invalid default", binding: pluginFlagBinding{property: "value", kind: pluginFlagInt, defaultProvided: true, defaultValue: "bad"}},
|
||||
{
|
||||
name: "invalid env",
|
||||
prepare: func(t *testing.T, _ *cobra.Command) { t.Setenv("PLUGIN_COVERAGE_BAD_ENV", "bad") },
|
||||
binding: pluginFlagBinding{property: "value", kind: pluginFlagInt, envDefault: "PLUGIN_COVERAGE_BAD_ENV"},
|
||||
},
|
||||
{name: "missing named required", binding: pluginFlagBinding{property: "value", names: []string{"value"}, required: true}},
|
||||
{name: "missing positional required", binding: pluginFlagBinding{property: "value", required: true, positional: true, positionalIndex: 0}},
|
||||
{name: "required omitted", binding: pluginFlagBinding{property: "value", required: true, defaultProvided: true, defaultValue: "", omitWhen: "empty"}},
|
||||
} {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
cmd := &cobra.Command{Use: "leaf"}
|
||||
if testCase.prepare != nil {
|
||||
testCase.prepare(t, cmd)
|
||||
}
|
||||
if err := collectPluginBindings(cmd, testCase.args, []pluginFlagBinding{testCase.binding}, testCase.params); err == nil {
|
||||
t.Fatal("expected binding error")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
params := map[string]any{"value": ""}
|
||||
if err := collectPluginBindings(&cobra.Command{Use: "leaf"}, nil, []pluginFlagBinding{{
|
||||
property: "value", kind: pluginFlagString, omitWhen: "empty",
|
||||
}}, params); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, exists := params["value"]; exists {
|
||||
t.Fatal("optional empty value was not omitted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginValueAndNamingHelpers(t *testing.T) {
|
||||
parseCases := []struct {
|
||||
kind pluginFlagKind
|
||||
raw string
|
||||
want any
|
||||
}{
|
||||
{pluginFlagInt, " 2 ", 2},
|
||||
{pluginFlagFloat, " 2.5 ", 2.5},
|
||||
{pluginFlagBool, "true", true},
|
||||
{pluginFlagStringSlice, "one, ,two", []string{"one", "two"}},
|
||||
{pluginFlagJSON, `{"ok":true}`, map[string]any{"ok": true}},
|
||||
{pluginFlagString, " raw ", " raw "},
|
||||
}
|
||||
for _, testCase := range parseCases {
|
||||
got, err := parsePluginValue(testCase.raw, testCase.kind)
|
||||
if err != nil || !reflect.DeepEqual(got, testCase.want) {
|
||||
t.Fatalf("parse %q = (%#v, %v), want %#v", testCase.raw, got, err, testCase.want)
|
||||
}
|
||||
}
|
||||
for _, testCase := range []struct {
|
||||
kind pluginFlagKind
|
||||
raw string
|
||||
}{
|
||||
{pluginFlagInt, "bad"}, {pluginFlagFloat, "bad"}, {pluginFlagBool, "bad"}, {pluginFlagJSON, "{"},
|
||||
} {
|
||||
if _, err := parsePluginValue(testCase.raw, testCase.kind); err == nil {
|
||||
t.Fatalf("invalid %q was accepted", testCase.raw)
|
||||
}
|
||||
}
|
||||
|
||||
omitCases := []struct {
|
||||
value any
|
||||
mode string
|
||||
want bool
|
||||
}{
|
||||
{nil, "", true}, {" ", "", true}, {[]string{}, "", true},
|
||||
{"", "never", false}, {false, "zero", true}, {0, "zero", true},
|
||||
{float64(0), "zero", true}, {true, "zero", false}, {1, "zero", false},
|
||||
{float64(1), "zero", false}, {[]any{}, "zero", true}, {map[string]any{}, "zero", true},
|
||||
{[]any{"value"}, "zero", false}, {map[string]any{"value": true}, "zero", false},
|
||||
{struct{}{}, "zero", false}, {false, "", false},
|
||||
}
|
||||
for _, testCase := range omitCases {
|
||||
if got := shouldOmitPluginValue(testCase.value, testCase.mode); got != testCase.want {
|
||||
t.Fatalf("omit (%#v, %q) = %v, want %v", testCase.value, testCase.mode, got, testCase.want)
|
||||
}
|
||||
}
|
||||
|
||||
wrapPluginParams(nil, "body")
|
||||
untouched := map[string]any{"value": 1}
|
||||
wrapPluginParams(untouched, " ")
|
||||
wrapped := map[string]any{"body": map[string]any{"old": 1}, "value": 2, "_meta": 3}
|
||||
wrapPluginParams(wrapped, "body")
|
||||
wantWrapped := map[string]any{"body": map[string]any{"old": 1, "value": 2}, "_meta": 3}
|
||||
if !reflect.DeepEqual(wrapped, wantWrapped) {
|
||||
t.Fatalf("wrapped = %#v, want %#v", wrapped, wantWrapped)
|
||||
}
|
||||
|
||||
kinds := map[string]pluginFlagKind{
|
||||
"int": pluginFlagInt, "integer": pluginFlagInt,
|
||||
"float": pluginFlagFloat, "float64": pluginFlagFloat, "number": pluginFlagFloat,
|
||||
"bool": pluginFlagBool, "boolean": pluginFlagBool,
|
||||
"stringSlice": pluginFlagStringSlice, "string_slice": pluginFlagStringSlice,
|
||||
"array": pluginFlagStringSlice, "[]string": pluginFlagStringSlice,
|
||||
"json": pluginFlagJSON, "object": pluginFlagJSON, "unknown": pluginFlagString,
|
||||
}
|
||||
for raw, want := range kinds {
|
||||
if got := pluginFlagKindFromString(raw); got != want {
|
||||
t.Fatalf("kind %q = %v, want %v", raw, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
used := map[string]bool{}
|
||||
reserved := map[string]bool{"r": true}
|
||||
if got := safePluginShorthand(" x ", used, reserved); got != "x" || !used["x"] {
|
||||
t.Fatalf("safe shorthand = %q / %#v", got, used)
|
||||
}
|
||||
for _, raw := range []string{"", "xy", "x", "r"} {
|
||||
if got := safePluginShorthand(raw, used, reserved); got != "" {
|
||||
t.Fatalf("unsafe shorthand %q = %q", raw, got)
|
||||
}
|
||||
}
|
||||
|
||||
baseReservations := pluginReservedFlags(nil)
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.PersistentFlags().StringP("custom", "c", "", "")
|
||||
rootReservations := pluginReservedFlags(root)
|
||||
if !baseReservations.names["yes"] || !rootReservations.names["custom"] || !rootReservations.shorthands["c"] {
|
||||
t.Fatalf("reservations = %#v / %#v", baseReservations, rootReservations)
|
||||
}
|
||||
|
||||
if got := safePluginAliases([]string{"", "help", "auth", "cmd", "cmd", "ok", "Bad"}, "cmd"); !reflect.DeepEqual(got, []string{"ok"}) {
|
||||
t.Fatalf("aliases = %#v", got)
|
||||
}
|
||||
if got := derivePluginCommandName("conference_getCurrent2Status", []string{"other", "conference"}); got != "get-current2-status" {
|
||||
t.Fatalf("derived name = %q", got)
|
||||
}
|
||||
if got := pluginKebabName(" HTTP2.Foo_bar baz@ "); got != "http2-foo-bar-baz@" {
|
||||
t.Fatalf("kebab name = %q", got)
|
||||
}
|
||||
for _, name := range []string{"", "1bad", "bad-", "bad--name", "bad_name", "bad@name"} {
|
||||
if validPluginKebabName(name) {
|
||||
t.Fatalf("invalid kebab name %q was accepted", name)
|
||||
}
|
||||
}
|
||||
if !validPluginKebabName("good-name2") || validPluginCommandName("help") || validPluginFlagName("json") || validPluginFlagName("params") {
|
||||
t.Fatal("name validation contract failed")
|
||||
}
|
||||
if got := firstNonEmptyPluginString(" ", " value "); got != "value" || firstNonEmptyPluginString("", " ") != "" {
|
||||
t.Fatal("first non-empty string contract failed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginConstraintGroupAndRootHelpers(t *testing.T) {
|
||||
cmd := &cobra.Command{Use: "leaf"}
|
||||
for _, name := range []string{"a", "b", "c"} {
|
||||
cmd.Flags().String(name, "", "")
|
||||
}
|
||||
applyPluginFlagConstraints(cmd, mcptypes.CLIToolOverride{
|
||||
MutuallyExclusive: [][]string{{"a", "b"}, {"a", "missing"}},
|
||||
RequireOneOf: [][]string{{"a", "b"}, {"missing"}},
|
||||
RequireTogether: [][]string{{"b", "c"}, {"c", "missing"}},
|
||||
})
|
||||
bindings := []pluginFlagBinding{{names: []string{"a"}}, {names: []string{"b"}}, {names: []string{"c"}}}
|
||||
if !validPluginFlagConstraints(bindings, mcptypes.CLIToolOverride{
|
||||
MutuallyExclusive: [][]string{{"a", "b"}},
|
||||
RequireOneOf: [][]string{{"a"}},
|
||||
RequireTogether: [][]string{{"b", "c"}},
|
||||
}) {
|
||||
t.Fatal("valid plugin constraints were rejected")
|
||||
}
|
||||
for _, invalid := range []mcptypes.CLIToolOverride{
|
||||
{MutuallyExclusive: [][]string{{"a"}}},
|
||||
{RequireOneOf: [][]string{{"missing"}}},
|
||||
{RequireTogether: [][]string{{"a", "a"}}},
|
||||
} {
|
||||
if validPluginFlagConstraints(bindings, invalid) {
|
||||
t.Fatalf("invalid plugin constraints were accepted: %#v", invalid)
|
||||
}
|
||||
}
|
||||
|
||||
groups := map[string]*cobra.Command{}
|
||||
root := &cobra.Command{Use: "root"}
|
||||
group := ensurePluginGroup(root, "parent.child", "child description", groups)
|
||||
if group.Name() != "child" || group.Short != "child description" || !cmdutil.IsPluginSourced(group) {
|
||||
t.Fatalf("group = %#v", group)
|
||||
}
|
||||
if again := ensurePluginGroup(root, "parent.child", "ignored", groups); again != group {
|
||||
t.Fatal("existing group was not reused")
|
||||
}
|
||||
for _, invalid := range []string{"safe.bad_name", "_bad", ".parent", "parent."} {
|
||||
if got := ensurePluginGroup(root, invalid, "invalid", groups); got != nil {
|
||||
t.Fatalf("invalid group path %q produced %#v", invalid, got)
|
||||
}
|
||||
}
|
||||
|
||||
mergePluginRoot(nil, root)
|
||||
mergePluginRoot(root, nil)
|
||||
destination := &cobra.Command{Use: "plugin", Aliases: []string{"one"}}
|
||||
source := &cobra.Command{Use: "plugin", Aliases: []string{"one", "two"}}
|
||||
source.AddCommand(&cobra.Command{Use: "leaf"})
|
||||
mergePluginRoot(destination, source)
|
||||
if !reflect.DeepEqual(destination.Aliases, []string{"one", "two"}) || requireOptionalPluginChild(destination, "leaf") == nil {
|
||||
t.Fatalf("merged root = %#v", destination)
|
||||
}
|
||||
|
||||
pruneEmptyPluginGroups(nil)
|
||||
pruneRoot := &cobra.Command{Use: "root"}
|
||||
empty := cobracmd.NewGroupCommand("empty", "empty")
|
||||
nonEmpty := cobracmd.NewGroupCommand("non-empty", "non-empty")
|
||||
nonEmpty.AddCommand(&cobra.Command{Use: "leaf"})
|
||||
pruneRoot.AddCommand(empty, nonEmpty)
|
||||
pruneEmptyPluginGroups(pruneRoot)
|
||||
if requireOptionalPluginChild(pruneRoot, "empty") != nil || requireOptionalPluginChild(pruneRoot, "non-empty") == nil {
|
||||
t.Fatal("empty plugin groups were not pruned correctly")
|
||||
}
|
||||
|
||||
if pluginRootBoolFlag(nil, "yes") {
|
||||
t.Fatal("nil command reported a root flag")
|
||||
}
|
||||
noFlag := &cobra.Command{Use: "root"}
|
||||
if pluginRootBoolFlag(noFlag, "yes") {
|
||||
t.Fatal("missing flag reported true")
|
||||
}
|
||||
wrongType := &cobra.Command{Use: "root"}
|
||||
wrongType.PersistentFlags().String("yes", "true", "")
|
||||
if pluginRootBoolFlag(wrongType, "yes") {
|
||||
t.Fatal("wrong flag type reported true")
|
||||
}
|
||||
boolRoot := &cobra.Command{Use: "root"}
|
||||
boolRoot.PersistentFlags().Bool("yes", false, "")
|
||||
if err := boolRoot.PersistentFlags().Set("yes", "true"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !pluginRootBoolFlag(boolRoot, "yes") {
|
||||
t.Fatal("true root flag was not observed")
|
||||
}
|
||||
if err := pluginConfirmationRequired("dws plugin"); err == nil || !strings.Contains(err.Error(), "sensitive") {
|
||||
t.Fatalf("confirmation error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsupportedPluginSemanticsReportEveryField(t *testing.T) {
|
||||
overlays := []struct {
|
||||
value mcptypes.CLIOverlay
|
||||
want string
|
||||
}{
|
||||
{mcptypes.CLIOverlay{Parent: "root"}, "parent"},
|
||||
{mcptypes.CLIOverlay{Group: "group"}, "group"},
|
||||
{mcptypes.CLIOverlay{ServerDeps: []string{"other"}}, "serverDeps"},
|
||||
{mcptypes.CLIOverlay{Hints: map[string]json.RawMessage{"x": json.RawMessage(`{}`)}}, "hintCommands"},
|
||||
{mcptypes.CLIOverlay{RedirectTo: "other"}, "redirectTo"},
|
||||
{mcptypes.CLIOverlay{}, ""},
|
||||
}
|
||||
for _, testCase := range overlays {
|
||||
if got := unsupportedPluginOverlay(testCase.value); got != testCase.want {
|
||||
t.Fatalf("unsupported overlay = %q, want %q", got, testCase.want)
|
||||
}
|
||||
}
|
||||
|
||||
tools := []struct {
|
||||
value mcptypes.CLIToolOverride
|
||||
want string
|
||||
}{
|
||||
{mcptypes.CLIToolOverride{CLIAliases: []string{"x"}}, "cliAliases"},
|
||||
{mcptypes.CLIToolOverride{OutputFormat: map[string]any{"x": true}}, "outputFormat"},
|
||||
{mcptypes.CLIToolOverride{ServerOverride: "other"}, "serverOverride"},
|
||||
{mcptypes.CLIToolOverride{RedirectTo: "x"}, "redirectTo"},
|
||||
{mcptypes.CLIToolOverride{Pipeline: []json.RawMessage{json.RawMessage(`{}`)}}, "pipeline"},
|
||||
{mcptypes.CLIToolOverride{}, ""},
|
||||
}
|
||||
for _, testCase := range tools {
|
||||
if got := unsupportedPluginToolOverride(testCase.value); got != testCase.want {
|
||||
t.Fatalf("unsupported tool = %q, want %q", got, testCase.want)
|
||||
}
|
||||
}
|
||||
|
||||
flags := []struct {
|
||||
value mcptypes.CLIFlagOverride
|
||||
want string
|
||||
}{
|
||||
{mcptypes.CLIFlagOverride{MapsTo: "x"}, "mapsTo"},
|
||||
{mcptypes.CLIFlagOverride{Transform: "x"}, "transform"},
|
||||
{mcptypes.CLIFlagOverride{TransformArgs: map[string]any{"x": true}}, "transformArgs"},
|
||||
{mcptypes.CLIFlagOverride{RuntimeDefault: "x"}, "runtimeDefault"},
|
||||
{mcptypes.CLIFlagOverride{PipelineLocal: true}, "pipelineLocal"},
|
||||
{mcptypes.CLIFlagOverride{Type: "mystery"}, "type"},
|
||||
{mcptypes.CLIFlagOverride{OmitWhen: "sometimes"}, "omitWhen"},
|
||||
{mcptypes.CLIFlagOverride{}, ""},
|
||||
}
|
||||
for _, testCase := range flags {
|
||||
if got := unsupportedPluginFlagOverride(testCase.value); got != testCase.want {
|
||||
t.Fatalf("unsupported flag = %q, want %q", got, testCase.want)
|
||||
}
|
||||
}
|
||||
for _, value := range []string{"", "string", "integer", "float64", "boolean", "stringSlice", "array", "json", "object"} {
|
||||
if !supportedPluginFlagType(value) {
|
||||
t.Fatalf("supported plugin flag type %q was rejected", value)
|
||||
}
|
||||
}
|
||||
for _, value := range []string{"", "empty", "zero", "never"} {
|
||||
if !supportedPluginOmitMode(value) {
|
||||
t.Fatalf("supported plugin omit mode %q was rejected", value)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsupportedPluginDescriptorRejectsEveryInvalidLayer(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
mutate func(*mcptypes.ServerDescriptor)
|
||||
want string
|
||||
}{
|
||||
{name: "overlay", mutate: func(value *mcptypes.ServerDescriptor) { value.CLI.Parent = "root" }, want: "parent"},
|
||||
{name: "no tools", mutate: func(value *mcptypes.ServerDescriptor) { value.CLI.ToolOverrides = nil }, want: ""},
|
||||
{name: "root", mutate: func(value *mcptypes.ServerDescriptor) { value.CLI.Command = "Bad" }, want: "command"},
|
||||
{name: "declared group", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.Groups = map[string]mcptypes.CLIGroupDef{"bad_name": {}}
|
||||
}, want: "groups"},
|
||||
{name: "blank tool", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"": {}}
|
||||
}, want: "tool"},
|
||||
{name: "tool semantics", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {ServerOverride: "drive"}}
|
||||
}, want: "serverOverride"},
|
||||
{name: "hidden tool", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {Hidden: true, ServerOverride: "drive"}}
|
||||
}, want: ""},
|
||||
{name: "derived leaf", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"conference_derived_tool": {}}
|
||||
}, want: ""},
|
||||
{name: "leaf", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {CLIName: "Bad"}}
|
||||
}, want: "cliName"},
|
||||
{name: "leaf group", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {CLIName: "leaf", Group: "bad_name"}}
|
||||
}, want: "group"},
|
||||
{name: "flags", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {
|
||||
CLIName: "leaf",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{"value": {Alias: "yes"}},
|
||||
}}
|
||||
}, want: "flags"},
|
||||
{name: "constraints", mutate: func(value *mcptypes.ServerDescriptor) {
|
||||
value.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{"tool": {
|
||||
CLIName: "leaf",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{"value": {}},
|
||||
RequireTogether: [][]string{{"value", "missing"}},
|
||||
}}
|
||||
}, want: "constraints"},
|
||||
{name: "valid", want: ""},
|
||||
}
|
||||
root := pluginTestRoot()
|
||||
for _, testCase := range testCases {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
if testCase.mutate != nil {
|
||||
testCase.mutate(&descriptor)
|
||||
}
|
||||
if got := unsupportedPluginDescriptor(root, descriptor); got != testCase.want {
|
||||
t.Fatalf("unsupported descriptor = %q, want %q", got, testCase.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,809 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
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/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
type pluginCaptureRunner struct {
|
||||
invocations []executor.Invocation
|
||||
}
|
||||
|
||||
func (r *pluginCaptureRunner) Run(_ context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
r.invocations = append(r.invocations, invocation)
|
||||
return executor.Result{Invocation: invocation}, nil
|
||||
}
|
||||
|
||||
func conferencePluginDescriptor() mcptypes.ServerDescriptor {
|
||||
return mcptypes.ServerDescriptor{
|
||||
Key: "conference-local",
|
||||
DisplayName: "conference/conference-local",
|
||||
Description: "conference plugin",
|
||||
Endpoint: "stdio://conference/conference-local",
|
||||
Source: "plugin",
|
||||
HasCLIMeta: true,
|
||||
CLI: mcptypes.CLIOverlay{
|
||||
ID: "conference-local",
|
||||
Command: "conference",
|
||||
Description: "视频会议:发起/邀请入会/会中控制",
|
||||
Prefixes: []string{"conference"},
|
||||
Groups: map[string]mcptypes.CLIGroupDef{
|
||||
"camera": {Description: "摄像头控制"},
|
||||
"mic": {Description: "麦克风控制"},
|
||||
"share": {Description: "屏幕共享"},
|
||||
},
|
||||
ToolOverrides: map[string]mcptypes.CLIToolOverride{
|
||||
"create_conference": {
|
||||
CLIName: "start",
|
||||
Description: "发起即时会议",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"title": {Description: "会议标题"},
|
||||
},
|
||||
},
|
||||
"get_conference_status": {
|
||||
CLIName: "status",
|
||||
Description: "查询当前会议状态",
|
||||
},
|
||||
"ai_end_meeting_for_all": {
|
||||
CLIName: "end",
|
||||
Description: "结束会议(所有人)",
|
||||
IsSensitive: true,
|
||||
},
|
||||
"ai_open_camera": {
|
||||
CLIName: "open",
|
||||
Group: "camera",
|
||||
Description: "打开摄像头",
|
||||
},
|
||||
"ai_mute_mic": {
|
||||
CLIName: "mute",
|
||||
Group: "mic",
|
||||
Description: "静音自己",
|
||||
},
|
||||
"ai_share_desktop": {
|
||||
CLIName: "start",
|
||||
Group: "share",
|
||||
Description: "开始共享桌面",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"capture_speaker": {Description: "是否共享电脑音频"},
|
||||
},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func pluginTestRoot(commands ...*cobra.Command) *cobra.Command {
|
||||
root := &cobra.Command{
|
||||
Use: "dws",
|
||||
SilenceErrors: true,
|
||||
SilenceUsage: true,
|
||||
}
|
||||
root.PersistentFlags().Bool("dry-run", false, "")
|
||||
root.PersistentFlags().Bool("yes", false, "")
|
||||
root.PersistentFlags().StringP("format", "f", "json", "")
|
||||
root.SetOut(io.Discard)
|
||||
root.SetErr(io.Discard)
|
||||
root.AddCommand(commands...)
|
||||
return root
|
||||
}
|
||||
|
||||
func requirePluginChild(t *testing.T, parent *cobra.Command, names ...string) *cobra.Command {
|
||||
t.Helper()
|
||||
current := parent
|
||||
for _, name := range names {
|
||||
var next *cobra.Command
|
||||
for _, child := range current.Commands() {
|
||||
if child.Name() == name {
|
||||
next = child
|
||||
break
|
||||
}
|
||||
}
|
||||
if next == nil {
|
||||
t.Fatalf("missing plugin command %q below %q", name, current.CommandPath())
|
||||
}
|
||||
current = next
|
||||
}
|
||||
return current
|
||||
}
|
||||
|
||||
func TestPluginOverlayBuildsConferenceTreeAndDispatchesOriginalProperties(t *testing.T) {
|
||||
runner := &pluginCaptureRunner{}
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{conferencePluginDescriptor()}, runner, nil)
|
||||
if len(commands) != 1 {
|
||||
t.Fatalf("plugin roots = %d, want 1", len(commands))
|
||||
}
|
||||
conference := commands[0]
|
||||
if conference.Name() != "conference" || conference.Short != "视频会议:发起/邀请入会/会中控制" {
|
||||
t.Fatalf("conference root = %q / %q", conference.Name(), conference.Short)
|
||||
}
|
||||
if !cmdutil.IsPluginSourced(conference) {
|
||||
t.Fatal("conference root is missing plugin provenance")
|
||||
}
|
||||
if got := requirePluginChild(t, conference, "camera").Short; got != "摄像头控制" {
|
||||
t.Fatalf("camera group short = %q", got)
|
||||
}
|
||||
if got := requirePluginChild(t, conference, "camera", "open").Short; got != "打开摄像头" {
|
||||
t.Fatalf("camera open short = %q", got)
|
||||
}
|
||||
requirePluginChild(t, conference, "mic", "mute")
|
||||
requirePluginChild(t, conference, "status")
|
||||
share := requirePluginChild(t, conference, "share", "start")
|
||||
flag := share.Flags().Lookup("capture-speaker")
|
||||
if flag == nil || flag.Usage != "是否共享电脑音频" {
|
||||
t.Fatalf("capture-speaker flag = %#v", flag)
|
||||
}
|
||||
|
||||
root := pluginTestRoot(commands...)
|
||||
root.SetArgs([]string{
|
||||
"conference", "start",
|
||||
"--json", `{"from_json":"kept","title":"json"}`,
|
||||
"--params", `{"from_params":2,"title":"params"}`,
|
||||
"--title", "验证会议",
|
||||
"--dry-run",
|
||||
})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("conference start: %v", err)
|
||||
}
|
||||
if len(runner.invocations) != 1 {
|
||||
t.Fatalf("runner calls = %d, want 1", len(runner.invocations))
|
||||
}
|
||||
invocation := runner.invocations[0]
|
||||
if invocation.Kind != "compat_invocation" ||
|
||||
invocation.CanonicalProduct != "conference-local" ||
|
||||
invocation.Tool != "create_conference" ||
|
||||
!invocation.DryRun {
|
||||
t.Fatalf("conference invocation = %#v", invocation)
|
||||
}
|
||||
wantParams := map[string]any{
|
||||
"from_json": "kept",
|
||||
"from_params": float64(2),
|
||||
"title": "验证会议",
|
||||
}
|
||||
if !reflect.DeepEqual(invocation.Params, wantParams) {
|
||||
t.Fatalf("conference params = %#v, want %#v", invocation.Params, wantParams)
|
||||
}
|
||||
|
||||
precedenceRunner := &pluginCaptureRunner{}
|
||||
precedenceRoot := pluginTestRoot(buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{conferencePluginDescriptor()},
|
||||
precedenceRunner,
|
||||
nil,
|
||||
)...)
|
||||
precedenceRoot.SetArgs([]string{
|
||||
"conference", "start",
|
||||
"--json", `{"title":"json"}`,
|
||||
"--params", `{"title":"params"}`,
|
||||
"--dry-run",
|
||||
})
|
||||
if err := precedenceRoot.Execute(); err != nil {
|
||||
t.Fatalf("conference payload precedence: %v", err)
|
||||
}
|
||||
if got := precedenceRunner.invocations[0].Params["title"]; got != "params" {
|
||||
t.Fatalf("conference payload title = %#v, want --params value", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginOverlayTypedFlags(t *testing.T) {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"typed_tool": {
|
||||
CLIName: "typed",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"conversationId": {Required: true, Description: "conversation"},
|
||||
"enabled": {Type: "bool"},
|
||||
"limit": {Type: "int"},
|
||||
"tags": {Type: "stringSlice"},
|
||||
},
|
||||
},
|
||||
}
|
||||
runner := &pluginCaptureRunner{}
|
||||
root := pluginTestRoot(buildPluginCommands([]mcptypes.ServerDescriptor{descriptor}, runner, nil)...)
|
||||
root.SetArgs([]string{
|
||||
"conference", "typed",
|
||||
"--conversation-id", "cid",
|
||||
"--enabled=false",
|
||||
"--limit", "3",
|
||||
"--tags", "one,two",
|
||||
"--dry-run",
|
||||
})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("typed plugin command: %v", err)
|
||||
}
|
||||
if len(runner.invocations) != 1 {
|
||||
t.Fatalf("runner calls = %d", len(runner.invocations))
|
||||
}
|
||||
invocation := runner.invocations[0]
|
||||
if invocation.CanonicalProduct != "conference-local" {
|
||||
t.Fatalf("canonical product = %q", invocation.CanonicalProduct)
|
||||
}
|
||||
want := map[string]any{
|
||||
"conversationId": "cid",
|
||||
"enabled": false,
|
||||
"limit": 3,
|
||||
"tags": []string{"one", "two"},
|
||||
}
|
||||
if !reflect.DeepEqual(invocation.Params, want) {
|
||||
t.Fatalf("typed params = %#v, want %#v", invocation.Params, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginSensitiveCommandRequiresConfirmation(t *testing.T) {
|
||||
for _, testCase := range []struct {
|
||||
name string
|
||||
args []string
|
||||
wantCalls int
|
||||
wantDry bool
|
||||
wantError bool
|
||||
}{
|
||||
{name: "blocked", args: []string{"conference", "end"}, wantError: true},
|
||||
{name: "preview", args: []string{"conference", "end", "--dry-run"}, wantCalls: 1, wantDry: true},
|
||||
{name: "confirmed", args: []string{"conference", "end", "--yes"}, wantCalls: 1},
|
||||
} {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
runner := &pluginCaptureRunner{}
|
||||
root := pluginTestRoot(buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{conferencePluginDescriptor()}, runner, nil)...)
|
||||
root.SetArgs(testCase.args)
|
||||
err := root.Execute()
|
||||
if testCase.wantError {
|
||||
var appErr *apperrors.Error
|
||||
if !errors.As(err, &appErr) ||
|
||||
appErr.Category != apperrors.CategoryValidation ||
|
||||
appErr.Reason != "confirmation_required" {
|
||||
t.Fatalf("sensitive error = %#v", err)
|
||||
}
|
||||
} else if err != nil {
|
||||
t.Fatalf("sensitive command: %v", err)
|
||||
}
|
||||
if len(runner.invocations) != testCase.wantCalls {
|
||||
t.Fatalf("runner calls = %d, want %d", len(runner.invocations), testCase.wantCalls)
|
||||
}
|
||||
if testCase.wantCalls == 1 && runner.invocations[0].DryRun != testCase.wantDry {
|
||||
t.Fatalf("dry-run = %v, want %v", runner.invocations[0].DryRun, testCase.wantDry)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginOverlayMergesServersWithoutProbingHTTP(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
var calls atomic.Int32
|
||||
server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
||||
calls.Add(1)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
first := conferencePluginDescriptor()
|
||||
first.Endpoint = server.URL
|
||||
first.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"one": {CLIName: "one"},
|
||||
}
|
||||
second := first
|
||||
second.Key = "conference-extra"
|
||||
second.DisplayName = "conference/conference-extra"
|
||||
second.Endpoint = server.URL + "/extra"
|
||||
second.CLI.ID = "conference-extra"
|
||||
second.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"two": {CLIName: "two"},
|
||||
}
|
||||
registerPluginHTTPServer(first)
|
||||
registerPluginHTTPServer(second)
|
||||
runner := &pluginCaptureRunner{}
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{second, first}, runner, nil)
|
||||
if len(commands) != 1 {
|
||||
t.Fatalf("merged roots = %d, want 1", len(commands))
|
||||
}
|
||||
requirePluginChild(t, commands[0], "one")
|
||||
requirePluginChild(t, commands[0], "two")
|
||||
root := pluginTestRoot(commands...)
|
||||
root.SetArgs([]string{"conference", "--help"})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("conference help: %v", err)
|
||||
}
|
||||
if got := calls.Load(); got != 0 {
|
||||
t.Fatalf("HTTP calls while building help = %d, want 0", got)
|
||||
}
|
||||
for _, command := range []string{"one", "two"} {
|
||||
root.SetArgs([]string{"conference", command, "--dry-run"})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("conference %s: %v", command, err)
|
||||
}
|
||||
}
|
||||
if len(runner.invocations) != 2 ||
|
||||
runner.invocations[0].CanonicalProduct != "conference-local" ||
|
||||
runner.invocations[1].CanonicalProduct != "conference-extra" {
|
||||
t.Fatalf("merged routes = %#v", runner.invocations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginCanReplaceHiddenFallbackButNotVisibleDistributionCommand(t *testing.T) {
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
fallback := &cobra.Command{Use: "conference", Hidden: true}
|
||||
fallback.AddCommand(&cobra.Command{Use: "meeting"})
|
||||
distribution := &cobra.Command{Use: "drive"}
|
||||
root.AddCommand(fallback, distribution)
|
||||
|
||||
conference := buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{conferencePluginDescriptor()},
|
||||
executor.EchoRunner{},
|
||||
nil,
|
||||
)[0]
|
||||
drive := &cobra.Command{Use: "drive"}
|
||||
cmdutil.MarkPluginSource(drive)
|
||||
addPluginCommandsSafe(root, []*cobra.Command{conference, drive})
|
||||
|
||||
gotConference := requirePluginChild(t, root, "conference")
|
||||
if gotConference == fallback || gotConference.Hidden {
|
||||
t.Fatalf("conference fallback was not replaced: %#v", gotConference)
|
||||
}
|
||||
requirePluginChild(t, gotConference, "status")
|
||||
if gotDrive := requirePluginChild(t, root, "drive"); gotDrive != distribution {
|
||||
t.Fatal("visible distribution command was replaced by a plugin")
|
||||
}
|
||||
}
|
||||
|
||||
func TestConflictingPluginDescriptorCannotReplaceDistributionEndpoint(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
pluginDir := filepath.Join(configDir, "plugins", "user", "drive-hijack")
|
||||
if err := os.MkdirAll(pluginDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
manifest := `{
|
||||
"name":"drive-hijack",
|
||||
"version":"1.0.0",
|
||||
"mcpServers":{
|
||||
"drive":{
|
||||
"type":"streamable-http",
|
||||
"endpoint":"https://plugin.invalid/mcp",
|
||||
"cli":{
|
||||
"id":"drive-service",
|
||||
"command":"drive-hijack",
|
||||
"toolOverrides":{"plugin_tool":{"cliName":"plugin-tool"}}
|
||||
}
|
||||
}
|
||||
}
|
||||
}`
|
||||
if err := os.WriteFile(filepath.Join(pluginDir, "plugin.json"), []byte(manifest), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
AppendDynamicServer(mcptypes.ServerDescriptor{
|
||||
Key: "drive",
|
||||
Endpoint: "https://distribution.invalid/mcp",
|
||||
CLI: mcptypes.CLIOverlay{ID: "drive-service", Command: "drive"},
|
||||
})
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.AddCommand(&cobra.Command{Use: "drive"})
|
||||
if commands := loadPlugins(root, nil, executor.EchoRunner{}); len(commands) != 0 {
|
||||
t.Fatalf("conflicting plugin commands = %#v", commands)
|
||||
}
|
||||
if endpoint, ok := directRuntimeEndpoint("drive-service", "plugin_tool"); !ok ||
|
||||
endpoint != "https://distribution.invalid/mcp" {
|
||||
t.Fatalf("drive endpoint after rejected plugin = (%q, %v)", endpoint, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSchemaSourceRootDoesNotLoadRuntimePlugins(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
previous := rootLoadPlugins
|
||||
t.Cleanup(func() { rootLoadPlugins = previous })
|
||||
var calls atomic.Int32
|
||||
rootLoadPlugins = func(*cobra.Command, *pipeline.Engine, executor.Runner) []*cobra.Command {
|
||||
calls.Add(1)
|
||||
AppendDynamicServer(conferencePluginDescriptor())
|
||||
return buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{conferencePluginDescriptor()},
|
||||
executor.EchoRunner{},
|
||||
nil,
|
||||
)
|
||||
}
|
||||
|
||||
base := NewSchemaSourceRootCommand()
|
||||
if calls.Load() != 0 {
|
||||
t.Fatalf("Schema source root loaded plugins %d times", calls.Load())
|
||||
}
|
||||
baseConference := requirePluginChild(t, base, "conference")
|
||||
if !baseConference.Hidden || requireOptionalPluginChild(baseConference, "status") != nil {
|
||||
t.Fatal("Schema source root contains installed conference plugin commands")
|
||||
}
|
||||
|
||||
runtime := NewRootCommand()
|
||||
if calls.Load() != 1 {
|
||||
t.Fatalf("runtime root plugin loads = %d, want 1", calls.Load())
|
||||
}
|
||||
runtimeConference := requirePluginChild(t, runtime, "conference")
|
||||
if runtimeConference.Hidden {
|
||||
t.Fatal("runtime conference plugin is hidden")
|
||||
}
|
||||
requirePluginChild(t, runtimeConference, "status")
|
||||
}
|
||||
|
||||
func requireOptionalPluginChild(parent *cobra.Command, name string) *cobra.Command {
|
||||
for _, child := range parent.Commands() {
|
||||
if child.Name() == name {
|
||||
return child
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func TestPluginDerivedNamesAndReservedAliases(t *testing.T) {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.Aliases = []string{"auth", "conf", "conf"}
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"conference_getCurrentStatus": {},
|
||||
}
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{descriptor}, executor.EchoRunner{}, nil)
|
||||
if len(commands) != 1 || !reflect.DeepEqual(commands[0].Aliases, []string{"conf"}) {
|
||||
t.Fatalf("plugin aliases = %#v", commands)
|
||||
}
|
||||
if requireOptionalPluginChild(commands[0], "get-current-status") == nil {
|
||||
var names []string
|
||||
for _, command := range commands[0].Commands() {
|
||||
names = append(names, command.Name())
|
||||
}
|
||||
t.Fatalf("derived command missing, got %s", strings.Join(names, ", "))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginFlagsCannotShadowHostControls(t *testing.T) {
|
||||
host := pluginTestRoot()
|
||||
host.PersistentFlags().StringP("host-extra", "x", "", "")
|
||||
reservations := pluginReservedFlags(host)
|
||||
for name := range reservations.names {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {
|
||||
CLIName: "unsafe",
|
||||
IsSensitive: true,
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"value": {Alias: name},
|
||||
},
|
||||
},
|
||||
}
|
||||
if commands := buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{descriptor},
|
||||
executor.EchoRunner{},
|
||||
host,
|
||||
); len(commands) != 0 {
|
||||
t.Fatalf("reserved host flag %q produced commands %#v", name, commands)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginShorthandsCannotShadowHostOrHelp(t *testing.T) {
|
||||
host := pluginTestRoot()
|
||||
host.PersistentFlags().StringP("host-extra", "x", "", "")
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"safe": {
|
||||
CLIName: "safe",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"alpha": {Shorthand: "f"},
|
||||
"bravo": {Shorthand: "h"},
|
||||
"charlie": {Shorthand: "o"},
|
||||
"delta": {Shorthand: "v"},
|
||||
"echo": {Shorthand: "x"},
|
||||
"foxtrot": {Shorthand: "y"},
|
||||
},
|
||||
},
|
||||
}
|
||||
runner := &pluginCaptureRunner{}
|
||||
commands := buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{descriptor},
|
||||
runner,
|
||||
host,
|
||||
)
|
||||
if len(commands) != 1 {
|
||||
t.Fatalf("plugin commands = %#v", commands)
|
||||
}
|
||||
host.AddCommand(commands...)
|
||||
leaf := requirePluginChild(t, commands[0], "safe")
|
||||
for _, name := range []string{"alpha", "bravo", "charlie", "delta", "echo", "foxtrot"} {
|
||||
if shorthand := leaf.Flags().Lookup(name).Shorthand; shorthand != "" {
|
||||
t.Fatalf("--%s shorthand = %q, want empty", name, shorthand)
|
||||
}
|
||||
}
|
||||
host.SetArgs([]string{"conference", "safe", "-h"})
|
||||
if err := host.Execute(); err != nil {
|
||||
t.Fatalf("plugin help: %v", err)
|
||||
}
|
||||
if len(runner.invocations) != 0 {
|
||||
t.Fatalf("help executed plugin: %#v", runner.invocations)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginPayloadPrecedenceRequiredAndTypedPositionals(t *testing.T) {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"payload": {
|
||||
CLIName: "payload",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"title": {Required: true},
|
||||
"mode": {Default: "fallback"},
|
||||
"enabled": {Positional: true, PositionalIndex: 0, Alias: "enabled-value", Required: true, Type: "bool"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range []struct {
|
||||
name string
|
||||
args []string
|
||||
wantEnabled bool
|
||||
}{
|
||||
{
|
||||
name: "flag satisfies dual positional",
|
||||
args: []string{
|
||||
"conference", "payload",
|
||||
"--params", `{"title":"from-json","mode":"from-json"}`,
|
||||
"--enabled-value=true",
|
||||
"--dry-run",
|
||||
},
|
||||
wantEnabled: true,
|
||||
},
|
||||
{
|
||||
name: "json beats positional",
|
||||
args: []string{
|
||||
"conference", "payload", "true",
|
||||
"--params", `{"title":"from-json","mode":"from-json","enabled":false}`,
|
||||
"--dry-run",
|
||||
},
|
||||
wantEnabled: false,
|
||||
},
|
||||
} {
|
||||
t.Run(testCase.name, func(t *testing.T) {
|
||||
runner := &pluginCaptureRunner{}
|
||||
root := pluginTestRoot(buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{descriptor},
|
||||
runner,
|
||||
nil,
|
||||
)...)
|
||||
root.SetArgs(testCase.args)
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("payload command: %v", err)
|
||||
}
|
||||
if len(runner.invocations) != 1 {
|
||||
t.Fatalf("runner calls = %d", len(runner.invocations))
|
||||
}
|
||||
params := runner.invocations[0].Params
|
||||
if params["title"] != "from-json" ||
|
||||
params["mode"] != "from-json" ||
|
||||
params["enabled"] != testCase.wantEnabled {
|
||||
t.Fatalf("payload params = %#v", params)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginDescriptorWinnerKeepsRouteAuthAndClientAtomic(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
writeManifest := func(name, manifest string) {
|
||||
t.Helper()
|
||||
directory := filepath.Join(configDir, "plugins", "user", name)
|
||||
if err := os.MkdirAll(directory, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(directory, "plugin.json"), []byte(manifest), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
writeManifest("alpha-plugin", `{
|
||||
"name":"alpha-plugin",
|
||||
"version":"1.0.0",
|
||||
"mcpServers":{
|
||||
"alpha":{
|
||||
"type":"streamable-http",
|
||||
"endpoint":"https://alpha.invalid/mcp",
|
||||
"headers":{"Authorization":"Bearer alpha-secret"},
|
||||
"cli":{
|
||||
"id":"shared-plugin-id",
|
||||
"command":"alpha-command",
|
||||
"toolOverrides":{"alpha_tool":{"cliName":"alpha"}}
|
||||
}
|
||||
},
|
||||
"alpha-extra":{
|
||||
"type":"streamable-http",
|
||||
"endpoint":"https://alpha-extra.invalid/mcp",
|
||||
"cli":{
|
||||
"id":"alpha-extra-id",
|
||||
"command":"alpha-command",
|
||||
"toolOverrides":{"extra_tool":{"cliName":"extra"}}
|
||||
}
|
||||
}
|
||||
}
|
||||
}`)
|
||||
writeManifest("beta-plugin", `{
|
||||
"name":"beta-plugin",
|
||||
"version":"1.0.0",
|
||||
"mcpServers":{
|
||||
"beta":{
|
||||
"type":"stdio",
|
||||
"command":"bin/beta",
|
||||
"cli":{
|
||||
"id":"shared-plugin-id",
|
||||
"command":"beta-command",
|
||||
"toolOverrides":{"beta_tool":{"cliName":"beta"}}
|
||||
}
|
||||
}
|
||||
}
|
||||
}`)
|
||||
|
||||
root := pluginTestRoot()
|
||||
commands := loadPlugins(root, nil, executor.EchoRunner{})
|
||||
if len(commands) != 1 || commands[0].Name() != "alpha-command" {
|
||||
t.Fatalf("plugin winner commands = %#v", commands)
|
||||
}
|
||||
requirePluginChild(t, commands[0], "alpha")
|
||||
requirePluginChild(t, commands[0], "extra")
|
||||
endpoint, ok := directRuntimeEndpoint("shared-plugin-id", "alpha_tool")
|
||||
if !ok || endpoint != "https://alpha.invalid/mcp" {
|
||||
t.Fatalf("winner endpoint = (%q, %v)", endpoint, ok)
|
||||
}
|
||||
extraEndpoint, ok := directRuntimeEndpoint("alpha-extra-id", "extra_tool")
|
||||
if !ok || extraEndpoint != "https://alpha-extra.invalid/mcp" {
|
||||
t.Fatalf("merged server endpoint = (%q, %v)", extraEndpoint, ok)
|
||||
}
|
||||
auth, ok := LookupPluginAuth("shared-plugin-id")
|
||||
if !ok || auth.Token != "alpha-secret" {
|
||||
t.Fatalf("winner auth = (%#v, %v)", auth, ok)
|
||||
}
|
||||
if _, ok := LookupStdioClient("beta-plugin/beta"); ok {
|
||||
t.Fatal("losing stdio client was registered")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsupportedPluginOverlaySemanticsFailClosed(t *testing.T) {
|
||||
for _, mutate := range []func(*mcptypes.ServerDescriptor){
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.RedirectTo = "drive"
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {
|
||||
CLIName: "unsafe",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"source": {MapsTo: "target"},
|
||||
},
|
||||
},
|
||||
}
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {
|
||||
CLIName: "unsafe",
|
||||
Pipeline: []json.RawMessage{json.RawMessage(`{"tool":"one"}`)},
|
||||
},
|
||||
}
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {CLIName: "unsafe", ServerOverride: "drive"},
|
||||
}
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {
|
||||
CLIName: "unsafe",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{
|
||||
"Body.query": {},
|
||||
},
|
||||
},
|
||||
}
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {CLIName: "unsafe", Group: "safe.bad_name"},
|
||||
}
|
||||
},
|
||||
func(descriptor *mcptypes.ServerDescriptor) {
|
||||
descriptor.CLI.ToolOverrides = map[string]mcptypes.CLIToolOverride{
|
||||
"unsafe": {
|
||||
CLIName: "unsafe",
|
||||
Flags: map[string]mcptypes.CLIFlagOverride{"value": {}},
|
||||
RequireTogether: [][]string{{"value", "missing"}},
|
||||
},
|
||||
}
|
||||
},
|
||||
} {
|
||||
descriptor := conferencePluginDescriptor()
|
||||
mutate(&descriptor)
|
||||
if commands := buildPluginCommands(
|
||||
[]mcptypes.ServerDescriptor{descriptor},
|
||||
executor.EchoRunner{},
|
||||
nil,
|
||||
); len(commands) != 0 {
|
||||
t.Fatalf("unsupported overlay produced commands %#v", commands)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsupportedPluginDescriptorsDoNotRegisterRuntimeState(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
writeManifest := func(name, manifest string) {
|
||||
t.Helper()
|
||||
directory := filepath.Join(configDir, "plugins", "user", name)
|
||||
if err := os.MkdirAll(directory, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(directory, "plugin.json"), []byte(manifest), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
writeManifest("unsafe-http", `{
|
||||
"name":"unsafe-http",
|
||||
"version":"1.0.0",
|
||||
"mcpServers":{"unsafe":{
|
||||
"type":"streamable-http",
|
||||
"endpoint":"https://unsafe.invalid/mcp",
|
||||
"headers":{"Authorization":"Bearer unsafe-secret"},
|
||||
"cli":{"id":"unsafe-http-id","command":"unsafe-http","toolOverrides":{
|
||||
"unsafe_tool":{"cliName":"run","serverOverride":"drive"}
|
||||
}}
|
||||
}}
|
||||
}`)
|
||||
writeManifest("unsafe-stdio", `{
|
||||
"name":"unsafe-stdio",
|
||||
"version":"1.0.0",
|
||||
"mcpServers":{"unsafe":{
|
||||
"type":"stdio",
|
||||
"command":"bin/unsafe",
|
||||
"cli":{"id":"unsafe-stdio-id","command":"unsafe-stdio","toolOverrides":{
|
||||
"unsafe_tool":{"cliName":"run","flags":{"value":{"mapsTo":"target"}}}
|
||||
}}
|
||||
}}
|
||||
}`)
|
||||
|
||||
root := pluginTestRoot()
|
||||
if commands := loadPlugins(root, nil, executor.EchoRunner{}); len(commands) != 0 {
|
||||
t.Fatalf("unsupported plugin descriptors produced commands %#v", commands)
|
||||
}
|
||||
if endpoint, ok := directRuntimeEndpoint("unsafe-http-id", "unsafe_tool"); ok {
|
||||
t.Fatalf("unsupported HTTP descriptor registered endpoint %q", endpoint)
|
||||
}
|
||||
if _, ok := LookupPluginAuth("unsafe-http-id"); ok {
|
||||
t.Fatal("unsupported HTTP descriptor registered plugin auth")
|
||||
}
|
||||
if _, ok := LookupStdioClient("unsafe-stdio/unsafe"); ok {
|
||||
t.Fatal("unsupported stdio descriptor registered a client")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
// 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 (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"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/transport"
|
||||
)
|
||||
|
||||
func pluginToolInputSchema(
|
||||
tools transport.ToolsListResult,
|
||||
toolName string,
|
||||
) (map[string]any, bool) {
|
||||
for _, tool := range tools.Tools {
|
||||
if strings.TrimSpace(tool.Name) == strings.TrimSpace(toolName) {
|
||||
return tool.InputSchema, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func normalizePluginInputParams(
|
||||
params map[string]any,
|
||||
schema map[string]any,
|
||||
) (map[string]any, error) {
|
||||
schema = canonicalPluginInputSchema(schema)
|
||||
normalized := make(map[string]any, len(params))
|
||||
for key, value := range params {
|
||||
normalized[key] = value
|
||||
}
|
||||
if _, err := coercePluginSchemaValue(normalized, schema); err != nil {
|
||||
return nil, cliInputValidationError(err)
|
||||
}
|
||||
if err := cli.ValidateInputSchema(normalized, schema); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func canonicalPluginInputSchema(schema map[string]any) map[string]any {
|
||||
if len(schema) == 0 {
|
||||
return schema
|
||||
}
|
||||
cloned := make(map[string]any, len(schema))
|
||||
for key, value := range schema {
|
||||
cloned[key] = clonePluginSchemaValue(key, value)
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func clonePluginSchemaValue(key string, value any) any {
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
cloned := make(map[string]any, len(typed))
|
||||
for childKey, childValue := range typed {
|
||||
cloned[childKey] = clonePluginSchemaValue(childKey, childValue)
|
||||
}
|
||||
return cloned
|
||||
case []any:
|
||||
cloned := make([]any, len(typed))
|
||||
for index, item := range typed {
|
||||
cloned[index] = clonePluginSchemaValue(key, item)
|
||||
}
|
||||
return cloned
|
||||
case []string:
|
||||
cloned := make([]string, len(typed))
|
||||
for index, item := range typed {
|
||||
if key == "type" {
|
||||
item = canonicalPluginSchemaType(item)
|
||||
}
|
||||
cloned[index] = item
|
||||
}
|
||||
return cloned
|
||||
case string:
|
||||
if key == "type" {
|
||||
return canonicalPluginSchemaType(typed)
|
||||
}
|
||||
return typed
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
func canonicalPluginSchemaType(value string) string {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "bool":
|
||||
return "boolean"
|
||||
case "int":
|
||||
return "integer"
|
||||
case "float":
|
||||
return "number"
|
||||
default:
|
||||
return value
|
||||
}
|
||||
}
|
||||
|
||||
func cliInputValidationError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
return apperrors.NewValidation(
|
||||
fmt.Sprintf("input schema normalization failed: %v", err),
|
||||
apperrors.WithReason("plugin_input_schema_invalid"),
|
||||
)
|
||||
}
|
||||
|
||||
func coercePluginSchemaValue(value any, schema map[string]any) (any, error) {
|
||||
target := singlePluginSchemaType(schema)
|
||||
if raw, ok := value.(string); ok {
|
||||
trimmed := strings.TrimSpace(raw)
|
||||
switch target {
|
||||
case "bool", "boolean":
|
||||
parsed, err := strconv.ParseBool(trimmed)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot convert %q to boolean: %w", raw, err)
|
||||
}
|
||||
value = parsed
|
||||
case "int", "integer":
|
||||
parsed, err := strconv.Atoi(trimmed)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot convert %q to integer: %w", raw, err)
|
||||
}
|
||||
value = parsed
|
||||
case "float", "number":
|
||||
parsed, err := strconv.ParseFloat(trimmed, 64)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("cannot convert %q to number: %w", raw, err)
|
||||
}
|
||||
value = parsed
|
||||
case "object":
|
||||
var parsed map[string]any
|
||||
if err := json.Unmarshal([]byte(trimmed), &parsed); err != nil {
|
||||
return nil, fmt.Errorf("cannot convert plugin parameter to object: %w", err)
|
||||
}
|
||||
if parsed == nil {
|
||||
return nil, fmt.Errorf("cannot convert plugin parameter to object: expected a JSON object")
|
||||
}
|
||||
value = parsed
|
||||
case "array":
|
||||
var parsed []any
|
||||
if strings.HasPrefix(trimmed, "[") {
|
||||
if err := json.Unmarshal([]byte(trimmed), &parsed); err != nil {
|
||||
return nil, fmt.Errorf("cannot convert plugin parameter to array: %w", err)
|
||||
}
|
||||
} else if trimmed != "" {
|
||||
for _, item := range strings.Split(trimmed, ",") {
|
||||
if item = strings.TrimSpace(item); item != "" {
|
||||
parsed = append(parsed, item)
|
||||
}
|
||||
}
|
||||
}
|
||||
value = parsed
|
||||
}
|
||||
}
|
||||
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
properties, _ := schema["properties"].(map[string]any)
|
||||
for key, propertyValue := range typed {
|
||||
propertySchema, _ := properties[key].(map[string]any)
|
||||
if len(propertySchema) == 0 {
|
||||
continue
|
||||
}
|
||||
coerced, err := coercePluginSchemaValue(propertyValue, propertySchema)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", key, err)
|
||||
}
|
||||
typed[key] = coerced
|
||||
}
|
||||
return typed, nil
|
||||
case []string:
|
||||
items := make([]any, len(typed))
|
||||
for index, item := range typed {
|
||||
items[index] = item
|
||||
}
|
||||
value = items
|
||||
}
|
||||
|
||||
if items, ok := value.([]any); ok {
|
||||
itemSchema, _ := schema["items"].(map[string]any)
|
||||
if len(itemSchema) == 0 {
|
||||
return items, nil
|
||||
}
|
||||
for index, item := range items {
|
||||
coerced, err := coercePluginSchemaValue(item, itemSchema)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("item %d: %w", index, err)
|
||||
}
|
||||
items[index] = coerced
|
||||
}
|
||||
return items, nil
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func singlePluginSchemaType(schema map[string]any) string {
|
||||
var types []string
|
||||
switch typed := schema["type"].(type) {
|
||||
case string:
|
||||
types = []string{typed}
|
||||
case []string:
|
||||
types = typed
|
||||
case []any:
|
||||
for _, value := range typed {
|
||||
if text, ok := value.(string); ok {
|
||||
types = append(types, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
var target string
|
||||
for _, candidate := range types {
|
||||
candidate = strings.TrimSpace(candidate)
|
||||
if candidate == "" || candidate == "null" {
|
||||
continue
|
||||
}
|
||||
if target != "" && target != candidate {
|
||||
return ""
|
||||
}
|
||||
target = candidate
|
||||
}
|
||||
return target
|
||||
}
|
||||
@@ -0,0 +1,229 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
)
|
||||
|
||||
func TestPluginToolInputSchemaMatchesTrimmedName(t *testing.T) {
|
||||
want := map[string]any{"type": "object"}
|
||||
tools := transport.ToolsListResult{Tools: []transport.ToolDescriptor{
|
||||
{Name: "other", InputSchema: map[string]any{"type": "string"}},
|
||||
{Name: " create_conference ", InputSchema: want},
|
||||
}}
|
||||
|
||||
got, ok := pluginToolInputSchema(tools, " create_conference ")
|
||||
if !ok || !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("pluginToolInputSchema() = (%#v, %v), want (%#v, true)", got, ok, want)
|
||||
}
|
||||
if got, ok := pluginToolInputSchema(tools, "missing"); ok || got != nil {
|
||||
t.Fatalf("missing pluginToolInputSchema() = (%#v, %v), want (nil, false)", got, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizePluginInputParamsCoercesNestedValues(t *testing.T) {
|
||||
schema := map[string]any{
|
||||
"type": "object",
|
||||
"required": []string{"enabled"},
|
||||
"properties": map[string]any{
|
||||
"enabled": map[string]any{"type": []any{"null", "bool"}},
|
||||
"count": map[string]any{"type": "int"},
|
||||
"ratio": map[string]any{"type": "float"},
|
||||
"settings": map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"active": map[string]any{"type": "bool"},
|
||||
},
|
||||
},
|
||||
"ids": map[string]any{
|
||||
"type": []string{"array", "null"},
|
||||
"items": map[string]any{"type": "int"},
|
||||
},
|
||||
"labels": map[string]any{
|
||||
"type": "array",
|
||||
"items": map[string]any{"type": "string"},
|
||||
},
|
||||
"booleans": map[string]any{
|
||||
"type": "array",
|
||||
"items": map[string]any{"type": "bool"},
|
||||
},
|
||||
"ambiguous": map[string]any{"type": []string{"string", "int"}},
|
||||
},
|
||||
}
|
||||
params := map[string]any{
|
||||
"enabled": " true ",
|
||||
"count": " 7 ",
|
||||
"ratio": " 2.5 ",
|
||||
"settings": `{"active":"false"}`,
|
||||
"ids": `["1", "2"]`,
|
||||
"labels": "alpha, , beta",
|
||||
"booleans": []string{"true", "false"},
|
||||
"ambiguous": "9",
|
||||
}
|
||||
|
||||
got, err := normalizePluginInputParams(params, schema)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizePluginInputParams() error = %v", err)
|
||||
}
|
||||
want := map[string]any{
|
||||
"enabled": true,
|
||||
"count": 7,
|
||||
"ratio": 2.5,
|
||||
"settings": map[string]any{"active": false},
|
||||
"ids": []any{1, 2},
|
||||
"labels": []any{"alpha", "beta"},
|
||||
"booleans": []any{true, false},
|
||||
"ambiguous": "9",
|
||||
}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("normalizePluginInputParams() = %#v, want %#v", got, want)
|
||||
}
|
||||
|
||||
properties := schema["properties"].(map[string]any)
|
||||
if gotType := properties["enabled"].(map[string]any)["type"].([]any)[1]; gotType != "bool" {
|
||||
t.Fatalf("normalization mutated source schema type to %#v", gotType)
|
||||
}
|
||||
if gotValue := params["enabled"]; gotValue != " true " {
|
||||
t.Fatalf("normalization mutated source params to %#v", gotValue)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizePluginInputParamsReportsConversionPath(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
value any
|
||||
fieldSchema map[string]any
|
||||
wantText string
|
||||
}{
|
||||
{name: "boolean", value: "sometimes", fieldSchema: map[string]any{"type": "bool"}, wantText: "cannot convert"},
|
||||
{name: "integer", value: "1.5", fieldSchema: map[string]any{"type": "int"}, wantText: "integer"},
|
||||
{name: "number", value: "many", fieldSchema: map[string]any{"type": "float"}, wantText: "number"},
|
||||
{name: "object", value: "{", fieldSchema: map[string]any{"type": "object"}, wantText: "object"},
|
||||
{name: "null object", value: "null", fieldSchema: map[string]any{"type": "object"}, wantText: "expected a JSON object"},
|
||||
{name: "array", value: "[", fieldSchema: map[string]any{"type": "array"}, wantText: "array"},
|
||||
{
|
||||
name: "nested property",
|
||||
value: `{"active":"sometimes"}`,
|
||||
fieldSchema: map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{
|
||||
"active": map[string]any{"type": "bool"},
|
||||
},
|
||||
},
|
||||
wantText: "field: active:",
|
||||
},
|
||||
{
|
||||
name: "array item",
|
||||
value: "1,not-an-int",
|
||||
fieldSchema: map[string]any{
|
||||
"type": "array",
|
||||
"items": map[string]any{"type": "int"},
|
||||
},
|
||||
wantText: "item 1",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
schema := map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{"field": tt.fieldSchema},
|
||||
}
|
||||
_, err := normalizePluginInputParams(map[string]any{"field": tt.value}, schema)
|
||||
if err == nil {
|
||||
t.Fatal("normalizePluginInputParams() error = nil, want conversion error")
|
||||
}
|
||||
var appError *apperrors.Error
|
||||
if !errors.As(err, &appError) ||
|
||||
appError.Category != apperrors.CategoryValidation ||
|
||||
appError.Reason != "plugin_input_schema_invalid" {
|
||||
t.Fatalf("conversion error = %#v, want categorized plugin schema validation error", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tt.wantText) {
|
||||
t.Fatalf("conversion error = %q, want text %q", err, tt.wantText)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizePluginInputParamsRunsSchemaValidation(t *testing.T) {
|
||||
schema := map[string]any{
|
||||
"type": "object",
|
||||
"required": []any{"name"},
|
||||
"properties": map[string]any{
|
||||
"name": map[string]any{"type": "string"},
|
||||
},
|
||||
}
|
||||
if _, err := normalizePluginInputParams(map[string]any{}, schema); err == nil ||
|
||||
!strings.Contains(err.Error(), "$.name is required") {
|
||||
t.Fatalf("required-field validation error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginInputSchemaHelperEdges(t *testing.T) {
|
||||
if got := canonicalPluginInputSchema(nil); got != nil {
|
||||
t.Fatalf("canonicalPluginInputSchema(nil) = %#v, want nil", got)
|
||||
}
|
||||
if got := clonePluginSchemaValue("minimum", 1); got != 1 {
|
||||
t.Fatalf("clonePluginSchemaValue(scalar) = %#v, want 1", got)
|
||||
}
|
||||
if got := cliInputValidationError(nil); got != nil {
|
||||
t.Fatalf("cliInputValidationError(nil) = %v, want nil", got)
|
||||
}
|
||||
|
||||
if got, err := coercePluginSchemaValue("", map[string]any{"type": "array"}); err != nil || !reflect.DeepEqual(got, []any(nil)) {
|
||||
t.Fatalf("empty array coercion = (%#v, %v), want nil slice", got, err)
|
||||
}
|
||||
items := []any{"unchanged"}
|
||||
if got, err := coercePluginSchemaValue(items, map[string]any{"type": "array"}); err != nil || !reflect.DeepEqual(got, items) {
|
||||
t.Fatalf("array without item schema = (%#v, %v)", got, err)
|
||||
}
|
||||
if got, err := coercePluginSchemaValue(12, map[string]any{"type": "integer"}); err != nil || got != 12 {
|
||||
t.Fatalf("non-string scalar coercion = (%#v, %v), want (12, nil)", got, err)
|
||||
}
|
||||
unknown := map[string]any{"unknown": "unchanged"}
|
||||
if got, err := coercePluginSchemaValue(unknown, map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
}); err != nil || !reflect.DeepEqual(got, unknown) {
|
||||
t.Fatalf("unknown property coercion = (%#v, %v), want unchanged map", got, err)
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
schema map[string]any
|
||||
want string
|
||||
}{
|
||||
{name: "missing", schema: map[string]any{}, want: ""},
|
||||
{name: "single string", schema: map[string]any{"type": "integer"}, want: "integer"},
|
||||
{name: "single string slice", schema: map[string]any{"type": []string{"null", "number"}}, want: "number"},
|
||||
{name: "any slice", schema: map[string]any{"type": []any{nil, 3, "", "null", "boolean"}}, want: "boolean"},
|
||||
{name: "ambiguous", schema: map[string]any{"type": []any{"string", "integer"}}, want: ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := singlePluginSchemaType(tt.schema); got != tt.want {
|
||||
t.Fatalf("singlePluginSchemaType(%#v) = %q, want %q", tt.schema, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
for raw, want := range map[string]string{
|
||||
" BOOL ": "boolean",
|
||||
"Int": "integer",
|
||||
"FLOAT": "number",
|
||||
"custom": "custom",
|
||||
} {
|
||||
if got := canonicalPluginSchemaType(raw); got != want {
|
||||
t.Errorf("canonicalPluginSchemaType(%q) = %q, want %q", raw, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
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/transport"
|
||||
)
|
||||
|
||||
func TestPluginStdioExecutionNormalizesAndValidatesLiveSchema(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
previousInit := runnerStdioEnsureInitialized
|
||||
previousList := runnerStdioListTools
|
||||
previousCall := runnerStdioCallTool
|
||||
t.Cleanup(func() {
|
||||
runnerStdioEnsureInitialized = previousInit
|
||||
runnerStdioListTools = previousList
|
||||
runnerStdioCallTool = previousCall
|
||||
})
|
||||
|
||||
client := transport.NewStdioClient("unused", nil, nil)
|
||||
RegisterStdioClient("conference/local", client)
|
||||
runnerStdioEnsureInitialized = func(*transport.StdioClient, context.Context) error {
|
||||
return nil
|
||||
}
|
||||
runnerStdioListTools = func(*transport.StdioClient, context.Context) (transport.ToolsListResult, error) {
|
||||
return transport.ToolsListResult{
|
||||
Tools: []transport.ToolDescriptor{{
|
||||
Name: "create_conference",
|
||||
InputSchema: map[string]any{
|
||||
"type": "object",
|
||||
"required": []any{"title"},
|
||||
"properties": map[string]any{
|
||||
"title": map[string]any{"type": "string"},
|
||||
"capture_speaker": map[string]any{"type": "bool"},
|
||||
},
|
||||
"additionalProperties": false,
|
||||
},
|
||||
}},
|
||||
}, nil
|
||||
}
|
||||
var calledParams map[string]any
|
||||
runnerStdioCallTool = func(
|
||||
_ *transport.StdioClient,
|
||||
_ context.Context,
|
||||
_ string,
|
||||
params map[string]any,
|
||||
) (transport.ToolCallResult, error) {
|
||||
calledParams = params
|
||||
return transport.ToolCallResult{Content: map[string]any{"ok": true}}, nil
|
||||
}
|
||||
|
||||
runner := &runtimeRunner{}
|
||||
invocation := executor.Invocation{
|
||||
CanonicalProduct: "conference-local",
|
||||
Tool: "create_conference",
|
||||
Params: map[string]any{
|
||||
"title": "schema validation",
|
||||
"capture_speaker": "true",
|
||||
},
|
||||
}
|
||||
result, err := runner.executeInvocation(
|
||||
context.Background(),
|
||||
"stdio://conference/local",
|
||||
invocation,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("stdio plugin execution: %v", err)
|
||||
}
|
||||
wantParams := map[string]any{
|
||||
"title": "schema validation",
|
||||
"capture_speaker": true,
|
||||
}
|
||||
if !reflect.DeepEqual(calledParams, wantParams) ||
|
||||
!reflect.DeepEqual(result.Invocation.Params, wantParams) {
|
||||
t.Fatalf("normalized wire params = %#v, result = %#v", calledParams, result.Invocation.Params)
|
||||
}
|
||||
|
||||
calledParams = nil
|
||||
invocation.Params = map[string]any{"capture_speaker": "true"}
|
||||
_, err = runner.executeInvocation(
|
||||
context.Background(),
|
||||
"stdio://conference/local",
|
||||
invocation,
|
||||
)
|
||||
var appError *apperrors.Error
|
||||
if !errors.As(err, &appError) ||
|
||||
appError.Category != apperrors.CategoryValidation ||
|
||||
calledParams != nil {
|
||||
t.Fatalf("missing required schema validation = %#v, call params = %#v", err, calledParams)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,369 @@
|
||||
// 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.
|
||||
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"testing"
|
||||
|
||||
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/pipeline"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/shortcut/userdef"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/cmdutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
"github.com/spf13/cobra"
|
||||
)
|
||||
|
||||
type schemaSourceContextKey struct{}
|
||||
|
||||
func TestSchemaSourceRootPropagatesContextWithoutLoadingPlugins(t *testing.T) {
|
||||
previous := rootLoadPlugins
|
||||
t.Cleanup(func() { rootLoadPlugins = previous })
|
||||
|
||||
pluginLoads := 0
|
||||
rootLoadPlugins = func(*cobra.Command, *pipeline.Engine, executor.Runner) []*cobra.Command {
|
||||
pluginLoads++
|
||||
return nil
|
||||
}
|
||||
wantContext := context.WithValue(context.Background(), schemaSourceContextKey{}, "schema")
|
||||
root := NewSchemaSourceRootCommand(wantContext)
|
||||
if root.Context() != wantContext {
|
||||
t.Fatal("Schema source root did not retain the caller context")
|
||||
}
|
||||
if pluginLoads != 0 {
|
||||
t.Fatalf("Schema source root loaded runtime plugins %d times", pluginLoads)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCollectPluginServerCandidatesSortsAndSkipsInvalidStdio(t *testing.T) {
|
||||
previousDescriptors := rootPluginDescriptors
|
||||
previousClients := rootPluginStdioClients
|
||||
previousDescriptor := rootPluginStdioDescriptor
|
||||
t.Cleanup(func() {
|
||||
rootPluginDescriptors = previousDescriptors
|
||||
rootPluginStdioClients = previousClients
|
||||
rootPluginStdioDescriptor = previousDescriptor
|
||||
})
|
||||
|
||||
first := &plugin.Plugin{Manifest: plugin.Manifest{Name: "first"}}
|
||||
second := &plugin.Plugin{Manifest: plugin.Manifest{Name: "second"}}
|
||||
wantContext := &plugin.UserContext{UserID: "user", CorpID: "corp"}
|
||||
client := transport.NewStdioClient("unused", nil, nil)
|
||||
|
||||
rootPluginDescriptors = func(owner *plugin.Plugin) []mcptypes.ServerDescriptor {
|
||||
if owner == first {
|
||||
return []mcptypes.ServerDescriptor{{Key: "same"}, {Key: " beta "}}
|
||||
}
|
||||
return []mcptypes.ServerDescriptor{{Key: "aardvark"}}
|
||||
}
|
||||
rootPluginStdioClients = func(owner *plugin.Plugin, gotContext *plugin.UserContext) []plugin.StdioServerClient {
|
||||
if gotContext != wantContext {
|
||||
t.Fatalf("stdio user context = %#v, want %#v", gotContext, wantContext)
|
||||
}
|
||||
if owner != first {
|
||||
return nil
|
||||
}
|
||||
return []plugin.StdioServerClient{
|
||||
{Key: "same", Client: client},
|
||||
{Key: " alpha ", Client: client},
|
||||
{Key: "invalid", Client: client},
|
||||
}
|
||||
}
|
||||
rootPluginStdioDescriptor = func(_ *plugin.Plugin, stdio plugin.StdioServerClient) (mcptypes.ServerDescriptor, bool) {
|
||||
if stdio.Key == "invalid" {
|
||||
return mcptypes.ServerDescriptor{}, false
|
||||
}
|
||||
return mcptypes.ServerDescriptor{Key: stdio.Key}, true
|
||||
}
|
||||
|
||||
candidates := collectPluginServerCandidates([]*plugin.Plugin{first, second}, wantContext)
|
||||
if len(candidates) != 5 {
|
||||
t.Fatalf("candidate count = %d, want 5", len(candidates))
|
||||
}
|
||||
gotKeys := make([]string, 0, len(candidates))
|
||||
gotKinds := make([]string, 0, len(candidates))
|
||||
for _, candidate := range candidates {
|
||||
gotKeys = append(gotKeys, candidate.descriptor.Key)
|
||||
if candidate.stdioClient == nil {
|
||||
gotKinds = append(gotKinds, "http")
|
||||
} else {
|
||||
gotKinds = append(gotKinds, "stdio")
|
||||
if candidate.stdioClient.Client != client {
|
||||
t.Fatal("stdio candidate did not retain its client")
|
||||
}
|
||||
}
|
||||
}
|
||||
if want := []string{" alpha ", " beta ", "same", "same", "aardvark"}; !reflect.DeepEqual(gotKeys, want) {
|
||||
t.Fatalf("candidate keys = %#v, want %#v", gotKeys, want)
|
||||
}
|
||||
if want := []string{"stdio", "http", "http", "stdio", "http"}; !reflect.DeepEqual(gotKinds, want) {
|
||||
t.Fatalf("candidate transports = %#v, want %#v", gotKinds, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginDescriptorBlankIdentityAndDistributionOwnership(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
|
||||
blank := mcptypes.ServerDescriptor{
|
||||
Key: " ",
|
||||
CLI: mcptypes.CLIOverlay{
|
||||
ID: " ",
|
||||
Command: " ",
|
||||
Aliases: []string{"", " "},
|
||||
},
|
||||
}
|
||||
if claims := pluginDescriptorIdentityClaims(blank); len(claims) != 0 {
|
||||
t.Fatalf("blank descriptor claims = %#v, want none", claims)
|
||||
}
|
||||
if rootName := pluginDescriptorRootName(blank); rootName != "" {
|
||||
t.Fatalf("blank descriptor root = %q", rootName)
|
||||
}
|
||||
owner := &plugin.Plugin{Manifest: plugin.Manifest{Name: "blank"}}
|
||||
accepted := selectPluginServerCandidates(
|
||||
&cobra.Command{Use: "dws"},
|
||||
[]pluginServerCandidate{
|
||||
{owner: owner, descriptor: mcptypes.ServerDescriptor{CLI: mcptypes.CLIOverlay{Skip: true}}},
|
||||
{owner: owner, descriptor: blank},
|
||||
},
|
||||
)
|
||||
if len(accepted) != 1 {
|
||||
t.Fatalf("blank descriptor candidates = %#v, want one accepted candidate", accepted)
|
||||
}
|
||||
|
||||
if distributionRootOwns(nil, "visible") {
|
||||
t.Fatal("nil root claimed a command")
|
||||
}
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
visible := &cobra.Command{Use: "visible", Aliases: []string{" visible-alias "}}
|
||||
hiddenFallback := &cobra.Command{Use: "conference", Hidden: true}
|
||||
hiddenOwned := &cobra.Command{Use: "hidden-owned", Hidden: true}
|
||||
pluginOwned := &cobra.Command{Use: "plugin-owned", Aliases: []string{"plugin-alias"}}
|
||||
cmdutil.MarkPluginSource(pluginOwned)
|
||||
root.AddCommand(visible, hiddenFallback, hiddenOwned, pluginOwned)
|
||||
|
||||
for _, name := range []string{"visible", "visible-alias", "hidden-owned"} {
|
||||
if !distributionRootOwns(root, name) {
|
||||
t.Errorf("distribution root did not claim %q", name)
|
||||
}
|
||||
}
|
||||
for _, name := range []string{"conference", "plugin-owned", "plugin-alias", "missing"} {
|
||||
if distributionRootOwns(root, name) {
|
||||
t.Errorf("distribution root unexpectedly claimed %q", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReplaceableFallbackIdentitySurvivesDistributionConflictChecks(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
SetDynamicServers([]mcptypes.ServerDescriptor{
|
||||
{
|
||||
Key: "conference",
|
||||
Endpoint: "https://example.com/conference/mcp",
|
||||
CLI: mcptypes.CLIOverlay{ID: "conference"},
|
||||
},
|
||||
{
|
||||
Key: "chat",
|
||||
Endpoint: "https://example.com/chat/mcp",
|
||||
CLI: mcptypes.CLIOverlay{ID: "chat"},
|
||||
},
|
||||
})
|
||||
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.AddCommand(&cobra.Command{Use: "conference", Hidden: true})
|
||||
distributionProducts := DirectRuntimeProductIDs()
|
||||
|
||||
conferenceDescriptor := mcptypes.ServerDescriptor{
|
||||
Key: "conference-local",
|
||||
DisplayName: "conference/conference-local",
|
||||
CLI: mcptypes.CLIOverlay{ID: "conference-local", Command: "conference"},
|
||||
}
|
||||
if pluginDescriptorConflictsWithDistribution(root, conferenceDescriptor, distributionProducts) {
|
||||
t.Fatal("replaceable fallback identity blocked plugin server selection")
|
||||
}
|
||||
|
||||
chatDescriptor := mcptypes.ServerDescriptor{
|
||||
Key: "chat-local",
|
||||
DisplayName: "chat/chat-local",
|
||||
CLI: mcptypes.CLIOverlay{ID: "chat-local", Command: "chat"},
|
||||
}
|
||||
if !pluginDescriptorConflictsWithDistribution(root, chatDescriptor, distributionProducts) {
|
||||
t.Fatal("non-replaceable distribution product no longer conflicts")
|
||||
}
|
||||
|
||||
reservedDescriptor := mcptypes.ServerDescriptor{
|
||||
Key: "auth-local",
|
||||
DisplayName: "auth/auth-local",
|
||||
CLI: mcptypes.CLIOverlay{ID: "auth-local", Command: "auth"},
|
||||
}
|
||||
if !pluginDescriptorConflictsWithDistribution(root, reservedDescriptor, distributionProducts) {
|
||||
t.Fatal("reserved command name no longer conflicts")
|
||||
}
|
||||
|
||||
first := &plugin.Plugin{Manifest: plugin.Manifest{Name: "conference"}}
|
||||
second := &plugin.Plugin{Manifest: plugin.Manifest{Name: "other"}}
|
||||
accepted := selectPluginServerCandidates(root, []pluginServerCandidate{
|
||||
{owner: first, descriptor: conferenceDescriptor},
|
||||
{
|
||||
owner: second,
|
||||
descriptor: mcptypes.ServerDescriptor{
|
||||
Key: "conference-other",
|
||||
DisplayName: "other/conference-other",
|
||||
CLI: mcptypes.CLIOverlay{ID: "conference-other", Command: "conference"},
|
||||
},
|
||||
},
|
||||
})
|
||||
if len(accepted) != 1 {
|
||||
t.Fatalf("accepted candidates = %d, want the first conference plugin only", len(accepted))
|
||||
}
|
||||
if accepted[0].owner != first {
|
||||
t.Fatalf("accepted owner = %q, want the first conference plugin", accepted[0].owner.Manifest.Name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddPluginCommandsSafeFiltersConflictingAliases(t *testing.T) {
|
||||
root := &cobra.Command{Use: "dws"}
|
||||
root.AddCommand(&cobra.Command{Use: "taken"})
|
||||
command := &cobra.Command{
|
||||
Use: "extension",
|
||||
Aliases: []string{"", "extension", "auth", "taken", "shared", " shared ", " okay "},
|
||||
}
|
||||
addPluginCommandsSafe(root, []*cobra.Command{
|
||||
command,
|
||||
{Use: "shared"},
|
||||
{Use: "other", Aliases: []string{"extension"}},
|
||||
})
|
||||
|
||||
if want := []string{"shared", "okay"}; !reflect.DeepEqual(command.Aliases, want) {
|
||||
t.Fatalf("filtered aliases = %#v, want %#v", command.Aliases, want)
|
||||
}
|
||||
if child := findDirectChild(root, "shared"); child != nil {
|
||||
t.Fatal("an accepted alias was also registered as a plugin primary command")
|
||||
}
|
||||
other := findDirectChild(root, "other")
|
||||
if other == nil || len(other.Aliases) != 0 {
|
||||
t.Fatalf("later plugin aliases = %#v", other)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioRunnerReportsToolsListFailureAndMissingTool(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
previousInit := runnerStdioEnsureInitialized
|
||||
previousList := runnerStdioListTools
|
||||
previousCall := runnerStdioCallTool
|
||||
t.Cleanup(func() {
|
||||
runnerStdioEnsureInitialized = previousInit
|
||||
runnerStdioListTools = previousList
|
||||
runnerStdioCallTool = previousCall
|
||||
})
|
||||
|
||||
client := transport.NewStdioClient("unused", nil, nil)
|
||||
RegisterStdioClient("plugin/server", client)
|
||||
runnerStdioEnsureInitialized = func(*transport.StdioClient, context.Context) error { return nil }
|
||||
toolCalls := 0
|
||||
runnerStdioCallTool = func(*transport.StdioClient, context.Context, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
toolCalls++
|
||||
return transport.ToolCallResult{}, nil
|
||||
}
|
||||
runner := &runtimeRunner{}
|
||||
invocation := executor.Invocation{CanonicalProduct: "overlay-id", Tool: "wanted"}
|
||||
|
||||
listFailure := errors.New("list failed")
|
||||
runnerStdioListTools = func(*transport.StdioClient, context.Context) (transport.ToolsListResult, error) {
|
||||
return transport.ToolsListResult{}, listFailure
|
||||
}
|
||||
_, err := runner.executeStdioInvocationAtEndpoint(context.Background(), "stdio://plugin/server", invocation)
|
||||
assertPluginRuntimeError(t, err, apperrors.CategoryAPI, "tools/list", "stdio_tools_list_error")
|
||||
|
||||
runnerStdioListTools = func(*transport.StdioClient, context.Context) (transport.ToolsListResult, error) {
|
||||
return transport.ToolsListResult{Tools: []transport.ToolDescriptor{{Name: "other"}}}, nil
|
||||
}
|
||||
_, err = runner.executeStdioInvocationAtEndpoint(context.Background(), "stdio://plugin/server", invocation)
|
||||
assertPluginRuntimeError(t, err, apperrors.CategoryValidation, "", "plugin_tool_not_found")
|
||||
if toolCalls != 0 {
|
||||
t.Fatalf("tools/call attempts after tools/list failures = %d", toolCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStdioManifestDescriptorAndRegistrationFailClosed(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
p := &plugin.Plugin{
|
||||
Manifest: plugin.Manifest{
|
||||
Name: "broken-plugin",
|
||||
MCPServers: map[string]*plugin.MCPServer{
|
||||
"local": {CLI: json.RawMessage(`{`)},
|
||||
},
|
||||
},
|
||||
}
|
||||
server := plugin.StdioServerClient{
|
||||
Key: "local",
|
||||
Client: transport.NewStdioClient("unused", nil, nil),
|
||||
}
|
||||
if descriptor, ok := stdioServerDescriptorFromManifest(p, server); ok || !reflect.ValueOf(descriptor).IsZero() {
|
||||
t.Fatalf("invalid descriptor = (%#v, %v), want zero, false", descriptor, ok)
|
||||
}
|
||||
if descriptor := registerStdioServerFromManifest(p, server); !reflect.ValueOf(descriptor).IsZero() {
|
||||
t.Fatalf("invalid registered descriptor = %#v, want zero", descriptor)
|
||||
}
|
||||
if _, ok := LookupStdioClient("broken-plugin/local"); ok {
|
||||
t.Fatal("invalid stdio manifest registered a client")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLegacyCommandsContinueWhenUserShortcutLoadFails(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
shortcutDir := filepath.Join(configDir, "shortcuts")
|
||||
if err := os.MkdirAll(shortcutDir, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(shortcutDir, "broken.yaml"), []byte("version: ["), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, loadErrors := userdef.Load(); len(loadErrors) == 0 {
|
||||
t.Fatal("malformed shortcut fixture did not fail to load")
|
||||
}
|
||||
|
||||
runner := executor.EchoRunner{}
|
||||
caller := newToolCallerAdapter(runner, &GlobalFlags{})
|
||||
if commands := newLegacyPublicCommands(runner, caller, true); len(commands) == 0 {
|
||||
t.Fatal("legacy commands were dropped after a user shortcut load error")
|
||||
}
|
||||
}
|
||||
|
||||
func findDirectChild(root *cobra.Command, name string) *cobra.Command {
|
||||
for _, command := range root.Commands() {
|
||||
if command.Name() == name {
|
||||
return command
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func assertPluginRuntimeError(
|
||||
t *testing.T,
|
||||
err error,
|
||||
wantCategory apperrors.Category,
|
||||
wantOperation string,
|
||||
wantReason string,
|
||||
) {
|
||||
t.Helper()
|
||||
var appError *apperrors.Error
|
||||
if !errors.As(err, &appError) {
|
||||
t.Fatalf("runtime error = %#v, want structured app error", err)
|
||||
}
|
||||
if appError.Category != wantCategory ||
|
||||
appError.Operation != wantOperation ||
|
||||
appError.Reason != wantReason {
|
||||
t.Fatalf("runtime error = %#v, want category=%q operation=%q reason=%q", appError, wantCategory, wantOperation, wantReason)
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -12,6 +13,7 @@ import (
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
@@ -35,6 +37,11 @@ func isolatePluginRuntime(t *testing.T) {
|
||||
stdioClients = make(map[string]*transport.StdioClient)
|
||||
stdioMu.Unlock()
|
||||
|
||||
pluginAuthMu.Lock()
|
||||
previousPluginAuth := pluginAuthRegistry
|
||||
pluginAuthRegistry = make(map[string]*PluginAuth)
|
||||
pluginAuthMu.Unlock()
|
||||
|
||||
t.Cleanup(func() {
|
||||
StopAllStdioClients()
|
||||
dynamicMu.Lock()
|
||||
@@ -46,6 +53,9 @@ func isolatePluginRuntime(t *testing.T) {
|
||||
stdioMu.Lock()
|
||||
stdioClients = previousStdio
|
||||
stdioMu.Unlock()
|
||||
pluginAuthMu.Lock()
|
||||
pluginAuthRegistry = previousPluginAuth
|
||||
pluginAuthMu.Unlock()
|
||||
})
|
||||
}
|
||||
|
||||
@@ -77,14 +87,40 @@ func TestRegisterPluginHTTPServerDoesNotProbeEndpoint(t *testing.T) {
|
||||
func TestRegisterStdioServerFromManifestDoesNotStartProcess(t *testing.T) {
|
||||
isolatePluginRuntime(t)
|
||||
marker := t.TempDir() + "/started"
|
||||
pluginRoot := t.TempDir()
|
||||
if err := os.WriteFile(pluginRoot+"/overlay.json", []byte(`{
|
||||
"id":"local",
|
||||
"command":"lazy-stdio",
|
||||
"groups":{"health":{"description":"health checks"}},
|
||||
"toolOverrides":{"ping":{"cliName":"ping","group":"health"}}
|
||||
}`), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := transport.NewStdioClient("/bin/sh", []string{
|
||||
"-c", fmt.Sprintf("printf started > %q", marker),
|
||||
}, nil)
|
||||
p := &plugin.Plugin{
|
||||
Manifest: plugin.Manifest{Name: "lazy-stdio", Description: "lazy stdio test"},
|
||||
Root: t.TempDir(),
|
||||
Manifest: plugin.Manifest{
|
||||
Name: "lazy-stdio",
|
||||
Description: "lazy stdio test",
|
||||
MCPServers: map[string]*plugin.MCPServer{
|
||||
"local": {
|
||||
Type: "stdio",
|
||||
Command: "unused",
|
||||
CLI: json.RawMessage(`"overlay.json"`),
|
||||
},
|
||||
},
|
||||
},
|
||||
Root: pluginRoot,
|
||||
}
|
||||
descriptor := registerStdioServerFromManifest(p, plugin.StdioServerClient{Key: "local", Client: client})
|
||||
commands := buildPluginCommands([]mcptypes.ServerDescriptor{descriptor}, executor.EchoRunner{}, nil)
|
||||
root := pluginTestRoot(commands...)
|
||||
root.SetArgs([]string{"lazy-stdio", "--help"})
|
||||
if err := root.Execute(); err != nil {
|
||||
t.Fatalf("lazy stdio help: %v", err)
|
||||
}
|
||||
requirePluginChild(t, commands[0], "health", "ping")
|
||||
|
||||
if _, err := os.Stat(marker); !os.IsNotExist(err) {
|
||||
t.Fatalf("stdio process started during registration: stat error = %v", err)
|
||||
|
||||
@@ -14,10 +14,7 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/plugin"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
@@ -33,50 +30,26 @@ import (
|
||||
// When no CLI metadata is present, a minimal overlay keyed by the server
|
||||
// name is returned so callers can still build an identity descriptor.
|
||||
func resolveStdioOverlay(p *plugin.Plugin, sc plugin.StdioServerClient) mcptypes.CLIOverlay {
|
||||
serverID := sc.Key
|
||||
overlay := mcptypes.CLIOverlay{
|
||||
ID: serverID,
|
||||
Command: serverID,
|
||||
}
|
||||
srv, ok := p.Manifest.MCPServers[sc.Key]
|
||||
if !ok || len(srv.CLI) == 0 {
|
||||
return overlay
|
||||
}
|
||||
|
||||
cliData := srv.CLI
|
||||
// A JSON string is interpreted as a relative path to an external
|
||||
// overlay file (e.g. "overlay.json") anchored at the plugin root.
|
||||
if len(cliData) > 0 && cliData[0] == '"' {
|
||||
var cliPath string
|
||||
if err := json.Unmarshal(cliData, &cliPath); err == nil && cliPath != "" {
|
||||
absPath := filepath.Join(p.Root, cliPath)
|
||||
if fileData, readErr := os.ReadFile(absPath); readErr == nil {
|
||||
cliData = fileData
|
||||
} else {
|
||||
slog.Warn("plugin: failed to read CLI overlay file",
|
||||
"plugin", p.Manifest.Name, "path", absPath, "error", readErr)
|
||||
}
|
||||
overlay, ok := p.ResolveCLIOverlay(sc.Key)
|
||||
if !ok {
|
||||
return mcptypes.CLIOverlay{
|
||||
ID: sc.Key,
|
||||
Command: sc.Key,
|
||||
Skip: true,
|
||||
}
|
||||
}
|
||||
if err := json.Unmarshal(cliData, &overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLI overlay for stdio server",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key, "error", err)
|
||||
}
|
||||
if overlay.ID == "" {
|
||||
overlay.ID = serverID
|
||||
}
|
||||
if overlay.Command == "" {
|
||||
overlay.Command = serverID
|
||||
}
|
||||
return overlay
|
||||
}
|
||||
|
||||
// registerStdioServerFromManifest registers an endpoint descriptor and an
|
||||
// unstarted client from versioned plugin metadata. Tool discovery is not part
|
||||
// of command-tree construction; execution starts and initializes the client.
|
||||
func registerStdioServerFromManifest(p *plugin.Plugin, sc plugin.StdioServerClient) mcptypes.ServerDescriptor {
|
||||
overlay := resolveStdioOverlay(p, sc)
|
||||
descriptor := mcptypes.ServerDescriptor{
|
||||
func stdioServerDescriptorFromManifest(
|
||||
p *plugin.Plugin,
|
||||
sc plugin.StdioServerClient,
|
||||
) (mcptypes.ServerDescriptor, bool) {
|
||||
overlay, ok := p.ResolveCLIOverlay(sc.Key)
|
||||
if !ok {
|
||||
return mcptypes.ServerDescriptor{}, false
|
||||
}
|
||||
return mcptypes.ServerDescriptor{
|
||||
Key: sc.Key,
|
||||
DisplayName: p.Manifest.Name + "/" + sc.Key,
|
||||
Description: p.Manifest.Description,
|
||||
@@ -84,13 +57,30 @@ func registerStdioServerFromManifest(p *plugin.Plugin, sc plugin.StdioServerClie
|
||||
Source: "plugin",
|
||||
CLI: overlay,
|
||||
HasCLIMeta: true,
|
||||
}
|
||||
}, true
|
||||
}
|
||||
|
||||
func registerResolvedStdioServer(
|
||||
p *plugin.Plugin,
|
||||
sc plugin.StdioServerClient,
|
||||
descriptor mcptypes.ServerDescriptor,
|
||||
) {
|
||||
AppendDynamicServer(descriptor)
|
||||
RegisterStdioClient(p.Manifest.Name+"/"+sc.Key, sc.Client)
|
||||
|
||||
slog.Debug("plugin: stdio server registered from manifest",
|
||||
"plugin", p.Manifest.Name, "server", sc.Key,
|
||||
"toolOverrides", len(overlay.ToolOverrides))
|
||||
"toolOverrides", len(descriptor.CLI.ToolOverrides))
|
||||
}
|
||||
|
||||
// registerStdioServerFromManifest registers an endpoint descriptor and an
|
||||
// unstarted client from versioned plugin metadata. Tool discovery is not part
|
||||
// of command-tree construction; execution starts and initializes the client.
|
||||
func registerStdioServerFromManifest(p *plugin.Plugin, sc plugin.StdioServerClient) mcptypes.ServerDescriptor {
|
||||
descriptor, ok := stdioServerDescriptorFromManifest(p, sc)
|
||||
if !ok {
|
||||
return mcptypes.ServerDescriptor{}
|
||||
}
|
||||
registerResolvedStdioServer(p, sc, descriptor)
|
||||
return descriptor
|
||||
}
|
||||
|
||||
+184
-89
@@ -18,7 +18,6 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -34,15 +33,17 @@ func newProfileCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "profile",
|
||||
Short: "组织 profile 管理",
|
||||
Long: `管理本机已登录的钉钉组织 profile。
|
||||
Long: `管理本机已登录的钉钉账号 profile。
|
||||
|
||||
每个 profile 对应一个已授权组织。业务命令可通过全局 --profile 临时指定组织,
|
||||
profile switch/use 才会持久修改默认组织上下文。`,
|
||||
每个 profile 由 corpId + userId 唯一确定,同一组织可保存多个账号。业务命令可通过
|
||||
全局 --profile 临时指定组织或账号,profile switch/use 才会持久修改默认账号。`,
|
||||
Example: ` dws profile list
|
||||
dws profile switch
|
||||
dws profile switch <corpId>
|
||||
dws profile switch <corpId>:<userId>
|
||||
dws profile switch "<corpName>:<userName>"
|
||||
dws profile switch -
|
||||
dws --profile <corpId> contact user get-self`,
|
||||
dws --profile <corpId>:<userId> contact user get-self`,
|
||||
Args: cobra.NoArgs,
|
||||
TraverseChildren: true,
|
||||
DisableAutoGenTag: true,
|
||||
@@ -58,8 +59,8 @@ func newProfileListCommand() *cobra.Command {
|
||||
return &cobra.Command{
|
||||
Use: "list",
|
||||
Aliases: []string{"ls"},
|
||||
Short: "列出已登录组织 profile",
|
||||
Long: "列出本机已登录的所有组织 profile,包含当前组织、主组织、组织名、corpId、状态和用户信息。",
|
||||
Short: "列出全部已登录账号 profile",
|
||||
Long: "列出本机全部已登录账号。状态和到期时间直接读取各身份 Token,列表本身不会刷新 Token。",
|
||||
Example: ` dws profile list
|
||||
dws profile list --format json`,
|
||||
Args: cobra.NoArgs,
|
||||
@@ -75,9 +76,9 @@ func newProfileListCommand() *cobra.Command {
|
||||
}
|
||||
format, _ := cmd.Root().PersistentFlags().GetString("format")
|
||||
if strings.EqualFold(strings.TrimSpace(format), "json") {
|
||||
return writeProfileListJSON(cmd.OutOrStdout(), cfg)
|
||||
return writeProfileListJSON(cmd.OutOrStdout(), configDir, cfg)
|
||||
}
|
||||
writeProfileListTable(cmd.OutOrStdout(), cfg)
|
||||
writeProfileListTable(cmd.OutOrStdout(), configDir, cfg)
|
||||
return nil
|
||||
},
|
||||
}
|
||||
@@ -85,9 +86,9 @@ func newProfileListCommand() *cobra.Command {
|
||||
|
||||
func newProfileUseCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "use [name|corpId|-]",
|
||||
Short: "切换当前组织 profile(兼容 profile switch)",
|
||||
Long: "兼容命令,语义等同于 dws profile switch。可用组织名、profile 名、corpId 或 - 切回上一个组织。",
|
||||
Use: "use [profile-selector|-]",
|
||||
Short: "切换当前账号 profile(兼容 profile switch)",
|
||||
Long: "兼容命令,语义等同于 dws profile switch。选择器支持组织 ID/名称、账号 ID/名称组合或本地 profile 名;- 切回上一个账号。",
|
||||
Example: ` dws profile use <corpId>
|
||||
dws profile use --name "钉钉"
|
||||
dws profile use -`,
|
||||
@@ -103,18 +104,21 @@ func newProfileUseCommand() *cobra.Command {
|
||||
|
||||
func newProfileSwitchCommand() *cobra.Command {
|
||||
cmd := &cobra.Command{
|
||||
Use: "switch [name|corpId|-]",
|
||||
Short: "切换当前组织 profile",
|
||||
Long: `切换默认组织 profile,并记录 previousProfile 以支持 dws profile switch - 快速切回。
|
||||
Use: "switch [profile-selector|-]",
|
||||
Short: "切换当前账号 profile",
|
||||
Long: `切换默认账号 profile,并记录 previousProfile 以支持 dws profile switch - 快速切回。
|
||||
|
||||
不带参数时,交互终端会展示组织选择器;非交互环境请显式传入组织名、profile 名或 corpId。
|
||||
需要只影响单次业务命令时,请使用全局 --profile。`,
|
||||
选择器支持 corpId:userId、corpId:userName、corpName:userId、corpName:userName,
|
||||
也兼容单独的 corpId、唯一 corpName 和本地 profile 名。组织或账号名称重名时会报错,
|
||||
要求改用稳定的 corpId:userId。不带参数时交互选择;单次执行请使用全局 --profile。`,
|
||||
Example: ` dws profile switch
|
||||
dws profile switch <corpId>
|
||||
dws profile switch <corpId>:<userId>
|
||||
dws profile switch "<corpName>:<userName>"
|
||||
dws profile switch --corpId <corpId>
|
||||
dws profile switch --name "钉钉"
|
||||
dws profile switch -
|
||||
dws --profile <corpId> contact user get-self`,
|
||||
dws --profile <corpId>:<userId> contact user get-self`,
|
||||
Args: cobra.MaximumNArgs(1),
|
||||
DisableAutoGenTag: true,
|
||||
RunE: func(cmd *cobra.Command, args []string) error {
|
||||
@@ -142,6 +146,7 @@ var (
|
||||
profileSwitchTUIRunner = runProfileSwitchTUI
|
||||
profileEnsureProfilesMigration = authpkg.EnsureProfilesMigration
|
||||
profileLoadProfiles = authpkg.LoadProfiles
|
||||
profileLoadTokenData = authpkg.LoadTokenDataForProfile
|
||||
profileUsePrevious = authpkg.UsePreviousProfile
|
||||
profileSetCurrent = authpkg.SetCurrentProfile
|
||||
profileRunTeaProgram = (*tea.Program).Run
|
||||
@@ -247,7 +252,7 @@ func switchProfileAndWrite(cmd *cobra.Command, configDir, selector string, usedT
|
||||
|
||||
func selectProfileSwitchProfile(cmd *cobra.Command, configDir string) (string, error) {
|
||||
if !profileSwitchInteractiveTerminal() {
|
||||
return "", apperrors.NewValidation("profile selector required in non-interactive mode; use dws profile switch <name|corpId>")
|
||||
return "", apperrors.NewValidation("profile selector required in non-interactive mode; use dws profile switch <corpId|corpId:userId|corpName:userName>")
|
||||
}
|
||||
if err := profileEnsureProfilesMigration(configDir); err != nil {
|
||||
return "", apperrors.NewInternal(fmt.Sprintf("failed to migrate profiles: %v", err))
|
||||
@@ -261,10 +266,7 @@ func selectProfileSwitchProfile(cmd *cobra.Command, configDir string) (string, e
|
||||
}
|
||||
choice := strings.TrimSpace(cfg.CurrentProfile)
|
||||
if choice == "" {
|
||||
choice = strings.TrimSpace(cfg.PrimaryProfile)
|
||||
}
|
||||
if choice == "" {
|
||||
choice = cfg.Profiles[0].CorpID
|
||||
choice = authpkg.ProfileSelector(cfg.Profiles[0])
|
||||
}
|
||||
return profileSwitchTUIRunner(cmd, cfg, choice)
|
||||
}
|
||||
@@ -306,7 +308,7 @@ func newProfileSwitchTUIModel(cfg *authpkg.ProfilesConfig, selectedCorpID string
|
||||
if cfg != nil {
|
||||
model.profiles = profileSwitchSortedProfiles(cfg.Profiles)
|
||||
}
|
||||
model.selected = profileSwitchProfileIndex(model.profiles, selectedCorpID)
|
||||
model.selected = profileSwitchProfileIndex(model.profiles, selectedCorpID, cfg)
|
||||
if model.selected < 0 {
|
||||
model.selected = 0
|
||||
}
|
||||
@@ -315,40 +317,7 @@ func newProfileSwitchTUIModel(cfg *authpkg.ProfilesConfig, selectedCorpID string
|
||||
}
|
||||
|
||||
func profileSwitchSortedProfiles(profiles []authpkg.Profile) []authpkg.Profile {
|
||||
sorted := append([]authpkg.Profile(nil), profiles...)
|
||||
sort.SliceStable(sorted, func(i, j int) bool {
|
||||
left, leftOK := profileSwitchSortTime(sorted[i])
|
||||
right, rightOK := profileSwitchSortTime(sorted[j])
|
||||
if leftOK && rightOK && !left.Equal(right) {
|
||||
return left.After(right)
|
||||
}
|
||||
if leftOK != rightOK {
|
||||
return leftOK
|
||||
}
|
||||
return false
|
||||
})
|
||||
return sorted
|
||||
}
|
||||
|
||||
func profileSwitchSortTime(p authpkg.Profile) (time.Time, bool) {
|
||||
for _, raw := range []string{p.LastLoginAt, p.UpdatedAt, p.LastUsedAt} {
|
||||
if t, ok := parseProfileSwitchTime(raw); ok {
|
||||
return t, true
|
||||
}
|
||||
}
|
||||
return time.Time{}, false
|
||||
}
|
||||
|
||||
func parseProfileSwitchTime(raw string) (time.Time, bool) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" {
|
||||
return time.Time{}, false
|
||||
}
|
||||
t, err := time.Parse(time.RFC3339, raw)
|
||||
if err != nil {
|
||||
return time.Time{}, false
|
||||
}
|
||||
return t, true
|
||||
return append([]authpkg.Profile(nil), profiles...)
|
||||
}
|
||||
|
||||
func (m profileSwitchTUIModel) Init() tea.Cmd {
|
||||
@@ -459,17 +428,31 @@ func (m profileSwitchTUIModel) selectedCorpID() string {
|
||||
if m.selected < 0 || m.selected >= len(m.profiles) {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(m.profiles[m.selected].CorpID)
|
||||
return authpkg.ProfileSelector(m.profiles[m.selected])
|
||||
}
|
||||
|
||||
func profileSwitchProfileIndex(profiles []authpkg.Profile, corpID string) int {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
func profileSwitchProfileIndex(profiles []authpkg.Profile, selector string, cfg *authpkg.ProfilesConfig) int {
|
||||
selector = strings.TrimSpace(selector)
|
||||
if corpID, userID, exact := authpkg.ParseIdentitySelector(selector); exact {
|
||||
for i, p := range profiles {
|
||||
if strings.TrimSpace(p.CorpID) == corpID && strings.TrimSpace(p.UserID) == userID {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
fallback := -1
|
||||
for i, p := range profiles {
|
||||
if strings.TrimSpace(p.CorpID) == corpID {
|
||||
return i
|
||||
if strings.TrimSpace(p.CorpID) == selector {
|
||||
if fallback < 0 {
|
||||
fallback = i
|
||||
}
|
||||
if profileIsOrgCurrent(p, cfg) {
|
||||
return i
|
||||
}
|
||||
}
|
||||
}
|
||||
return -1
|
||||
return fallback
|
||||
}
|
||||
|
||||
func profileSwitchOptionLabel(p authpkg.Profile, cfg *authpkg.ProfilesConfig) string {
|
||||
@@ -481,11 +464,29 @@ func profileSwitchOptionLabel(p authpkg.Profile, cfg *authpkg.ProfilesConfig) st
|
||||
}
|
||||
|
||||
func profileSwitchProfileCells(p authpkg.Profile, cfg *authpkg.ProfilesConfig) (string, string) {
|
||||
return profileOrgName(p), profileSwitchProfileStatus(p, cfg)
|
||||
orgName := profileOrgName(p)
|
||||
if cfg != nil {
|
||||
sameCorp := 0
|
||||
for _, candidate := range cfg.Profiles {
|
||||
if candidate.CorpID == p.CorpID {
|
||||
sameCorp++
|
||||
}
|
||||
}
|
||||
if sameCorp > 1 {
|
||||
user := strings.TrimSpace(p.UserName)
|
||||
if user == "" {
|
||||
user = strings.TrimSpace(p.UserID)
|
||||
}
|
||||
if user != "" {
|
||||
orgName += " / " + user
|
||||
}
|
||||
}
|
||||
}
|
||||
return orgName, profileSwitchProfileStatus(p, cfg)
|
||||
}
|
||||
|
||||
func profileSwitchProfileStatus(p authpkg.Profile, cfg *authpkg.ProfilesConfig) string {
|
||||
if cfg != nil && p.CorpID == cfg.CurrentProfile {
|
||||
if cfg != nil && profileSelectorSelectsProfile(cfg.CurrentProfile, p, profileIsOrgCurrent(p, cfg), profileCountForCorp(cfg, p.CorpID) <= 1) {
|
||||
return "当前组织"
|
||||
}
|
||||
return ""
|
||||
@@ -575,6 +576,7 @@ type profileUseResponse struct {
|
||||
}
|
||||
|
||||
type profileView struct {
|
||||
Profile string `json:"profile"`
|
||||
CorpID string `json:"corpId"`
|
||||
CorpName string `json:"corpName"`
|
||||
UserID string `json:"userId,omitempty"`
|
||||
@@ -588,15 +590,16 @@ type profileView struct {
|
||||
LastUsedAt string `json:"lastUsedAt,omitempty"`
|
||||
IsPrimary bool `json:"isPrimary"`
|
||||
IsCurrent bool `json:"isCurrent"`
|
||||
IsOrgCurrent bool `json:"isOrgCurrent"`
|
||||
}
|
||||
|
||||
func writeProfileListJSON(w io.Writer, cfg *authpkg.ProfilesConfig) error {
|
||||
func writeProfileListJSON(w io.Writer, configDir string, cfg *authpkg.ProfilesConfig) error {
|
||||
resp := profileListResponse{
|
||||
Success: true,
|
||||
PrimaryProfile: cfg.PrimaryProfile,
|
||||
CurrentProfile: cfg.CurrentProfile,
|
||||
PreviousProfile: cfg.PreviousProfile,
|
||||
Profiles: profileViews(cfg),
|
||||
Profiles: profileViews(configDir, cfg),
|
||||
}
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
@@ -612,44 +615,50 @@ func writeProfileUseJSON(w io.Writer, profile *authpkg.Profile, cfg *authpkg.Pro
|
||||
primaryProfile = cfg.PrimaryProfile
|
||||
currentProfile = cfg.CurrentProfile
|
||||
}
|
||||
resp.Profile = profileViewFromProfile(*profile, primaryProfile, currentProfile)
|
||||
resp.Profile = profileViewFromProfile(
|
||||
*profile,
|
||||
cfg,
|
||||
primaryProfile,
|
||||
currentProfile,
|
||||
profileCountForCorp(cfg, profile.CorpID) <= 1,
|
||||
nil,
|
||||
)
|
||||
}
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetIndent("", " ")
|
||||
return enc.Encode(resp)
|
||||
}
|
||||
|
||||
func writeProfileListTable(w io.Writer, cfg *authpkg.ProfilesConfig) {
|
||||
func writeProfileListTable(w io.Writer, configDir string, cfg *authpkg.ProfilesConfig) {
|
||||
if cfg == nil || len(cfg.Profiles) == 0 {
|
||||
fmt.Fprintln(w, "未找到已登录 profile")
|
||||
return
|
||||
}
|
||||
fmt.Fprintf(w, "%-3s %-3s %-28s %-34s %-10s %s\n", "CUR", "PRI", "ORG_NAME", "CORP_ID", "STATUS", "USER")
|
||||
fmt.Fprintf(w, "%-3s %-28s %-34s %-10s %s\n", "CUR", "ORG_NAME", "CORP_ID", "STATUS", "USER")
|
||||
for _, p := range cfg.Profiles {
|
||||
view := profileViewFromProfile(
|
||||
p,
|
||||
cfg,
|
||||
cfg.PrimaryProfile,
|
||||
cfg.CurrentProfile,
|
||||
profileCountForCorp(cfg, p.CorpID) == 1,
|
||||
loadProfileTokenState(configDir, p),
|
||||
)
|
||||
current := ""
|
||||
if p.CorpID == cfg.CurrentProfile {
|
||||
if view.IsCurrent {
|
||||
current = "*"
|
||||
}
|
||||
primary := ""
|
||||
if p.CorpID == cfg.PrimaryProfile {
|
||||
primary = "*"
|
||||
}
|
||||
user := p.UserName
|
||||
if user == "" {
|
||||
user = p.UserID
|
||||
}
|
||||
status := p.Status
|
||||
if status == "" {
|
||||
status = authpkg.ProfileStatusActive
|
||||
}
|
||||
fmt.Fprintf(
|
||||
w,
|
||||
"%-3s %-3s %-28s %-34s %-10s %s\n",
|
||||
"%-3s %-28s %-34s %-10s %s\n",
|
||||
current,
|
||||
primary,
|
||||
clipProfileCell(profileOrgName(p), 28),
|
||||
clipProfileCell(p.CorpID, 34),
|
||||
status,
|
||||
view.Status,
|
||||
user,
|
||||
)
|
||||
}
|
||||
@@ -677,19 +686,40 @@ func profileOrgName(p authpkg.Profile) string {
|
||||
return strings.TrimSpace(p.CorpID)
|
||||
}
|
||||
|
||||
func profileViews(cfg *authpkg.ProfilesConfig) []profileView {
|
||||
type profileTokenState struct {
|
||||
Status string
|
||||
ExpiresAt string
|
||||
RefreshExpAt string
|
||||
}
|
||||
|
||||
func profileViews(configDir string, cfg *authpkg.ProfilesConfig) []profileView {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
views := make([]profileView, 0, len(cfg.Profiles))
|
||||
for _, p := range cfg.Profiles {
|
||||
views = append(views, profileViewFromProfile(p, cfg.PrimaryProfile, cfg.CurrentProfile))
|
||||
views = append(views, profileViewFromProfile(
|
||||
p,
|
||||
cfg,
|
||||
cfg.PrimaryProfile,
|
||||
cfg.CurrentProfile,
|
||||
profileCountForCorp(cfg, p.CorpID) == 1,
|
||||
loadProfileTokenState(configDir, p),
|
||||
))
|
||||
}
|
||||
return views
|
||||
}
|
||||
|
||||
func profileViewFromProfile(p authpkg.Profile, primaryProfile, currentProfile string) profileView {
|
||||
return profileView{
|
||||
func profileViewFromProfile(
|
||||
p authpkg.Profile,
|
||||
cfg *authpkg.ProfilesConfig,
|
||||
primaryProfile, currentProfile string,
|
||||
onlyAccountInOrg bool,
|
||||
tokenState *profileTokenState,
|
||||
) profileView {
|
||||
isOrgCurrent := profileIsOrgCurrent(p, cfg)
|
||||
view := profileView{
|
||||
Profile: authpkg.ProfileSelector(p),
|
||||
CorpID: p.CorpID,
|
||||
CorpName: profileOrgName(p),
|
||||
UserID: p.UserID,
|
||||
@@ -701,9 +731,74 @@ func profileViewFromProfile(p authpkg.Profile, primaryProfile, currentProfile st
|
||||
RefreshExpAt: p.RefreshExpAt,
|
||||
LastLoginAt: p.LastLoginAt,
|
||||
LastUsedAt: p.LastUsedAt,
|
||||
IsPrimary: p.CorpID == primaryProfile,
|
||||
IsCurrent: p.CorpID == currentProfile,
|
||||
IsPrimary: profileSelectorSelectsProfile(primaryProfile, p, isOrgCurrent, onlyAccountInOrg),
|
||||
IsCurrent: profileSelectorSelectsProfile(currentProfile, p, isOrgCurrent, onlyAccountInOrg),
|
||||
IsOrgCurrent: isOrgCurrent,
|
||||
}
|
||||
if tokenState != nil {
|
||||
view.Status = tokenState.Status
|
||||
view.ExpiresAt = tokenState.ExpiresAt
|
||||
view.RefreshExpAt = tokenState.RefreshExpAt
|
||||
}
|
||||
return view
|
||||
}
|
||||
|
||||
func loadProfileTokenState(configDir string, profile authpkg.Profile) *profileTokenState {
|
||||
data, err := profileLoadTokenData(configDir, authpkg.ProfileSelector(profile))
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) || (err == nil && data == nil) {
|
||||
return &profileTokenState{Status: authpkg.ProfileStatusRevoked}
|
||||
}
|
||||
if err != nil {
|
||||
return &profileTokenState{Status: authpkg.ProfileStatusUnavailable}
|
||||
}
|
||||
status := authpkg.ProfileStatusExpired
|
||||
if data.IsAccessTokenValid() {
|
||||
status = authpkg.ProfileStatusActive
|
||||
}
|
||||
return &profileTokenState{
|
||||
Status: status,
|
||||
ExpiresAt: profileTokenTime(data.ExpiresAt),
|
||||
RefreshExpAt: profileTokenTime(data.RefreshExpAt),
|
||||
}
|
||||
}
|
||||
|
||||
func profileTokenTime(value time.Time) string {
|
||||
if value.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return value.Format(time.RFC3339)
|
||||
}
|
||||
|
||||
func profileSelectorSelectsProfile(selector string, profile authpkg.Profile, isOrgCurrent, onlyAccountInOrg bool) bool {
|
||||
selector = strings.TrimSpace(selector)
|
||||
if corpID, userID, exact := authpkg.ParseIdentitySelector(selector); exact {
|
||||
return corpID == strings.TrimSpace(profile.CorpID) && userID == strings.TrimSpace(profile.UserID)
|
||||
}
|
||||
return selector == strings.TrimSpace(profile.CorpID) && (isOrgCurrent || onlyAccountInOrg)
|
||||
}
|
||||
|
||||
func profileCountForCorp(cfg *authpkg.ProfilesConfig, corpID string) int {
|
||||
if cfg == nil {
|
||||
return 0
|
||||
}
|
||||
count := 0
|
||||
for _, profile := range cfg.Profiles {
|
||||
if strings.TrimSpace(profile.CorpID) == strings.TrimSpace(corpID) {
|
||||
count++
|
||||
}
|
||||
}
|
||||
return count
|
||||
}
|
||||
|
||||
func profileIsOrgCurrent(profile authpkg.Profile, cfg *authpkg.ProfilesConfig) bool {
|
||||
if cfg == nil {
|
||||
return false
|
||||
}
|
||||
selector := strings.TrimSpace(cfg.OrgCurrentProfiles[strings.TrimSpace(profile.CorpID)])
|
||||
if corpID, userID, exact := authpkg.ParseIdentitySelector(selector); exact {
|
||||
return corpID == strings.TrimSpace(profile.CorpID) && userID == strings.TrimSpace(profile.UserID)
|
||||
}
|
||||
return profileCountForCorp(cfg, profile.CorpID) == 1
|
||||
}
|
||||
|
||||
func clipProfileCell(value string, limit int) string {
|
||||
|
||||
@@ -16,9 +16,11 @@ package app
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
tea "github.com/charmbracelet/bubbletea"
|
||||
@@ -80,8 +82,10 @@ func TestProfileListRootCommandJSONIncludesCorpName(t *testing.T) {
|
||||
if !resp.Success {
|
||||
t.Fatal("success = false, want true")
|
||||
}
|
||||
if resp.PrimaryProfile != "corp_primary" || resp.CurrentProfile != "corp_secondary" || resp.PreviousProfile != "corp_primary" {
|
||||
t.Fatalf("profile pointers = primary %q current %q previous %q, want corp_primary/corp_secondary/corp_primary", resp.PrimaryProfile, resp.CurrentProfile, resp.PreviousProfile)
|
||||
if resp.PrimaryProfile != "" ||
|
||||
resp.CurrentProfile != "corp_secondary:user-corp_secondary" ||
|
||||
resp.PreviousProfile != "corp_primary:user-corp_primary" {
|
||||
t.Fatalf("profile pointers = primary %q current %q previous %q", resp.PrimaryProfile, resp.CurrentProfile, resp.PreviousProfile)
|
||||
}
|
||||
if len(resp.Profiles) != 2 {
|
||||
t.Fatalf("profiles len = %d, want 2", len(resp.Profiles))
|
||||
@@ -96,6 +100,181 @@ func TestProfileListRootCommandJSONIncludesCorpName(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileListRootCommandJSONIncludesAllAccountsInSameCorp(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_same")
|
||||
first.UserID = "user_1"
|
||||
first.UserName = "账号一"
|
||||
second := authLogoutTestToken("corp_same")
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
second.UserName = "账号二"
|
||||
setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"--format", "json", "profile", "list"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("profile list --format json error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
var resp profileListResponse
|
||||
if err := json.Unmarshal(out.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if len(resp.Profiles) != 2 {
|
||||
t.Fatalf("profiles len = %d, want 2: %#v", len(resp.Profiles), resp.Profiles)
|
||||
}
|
||||
got := make(map[string]profileView, len(resp.Profiles))
|
||||
for _, profile := range resp.Profiles {
|
||||
got[profile.Profile] = profile
|
||||
}
|
||||
if _, ok := got["corp_same:user_1"]; !ok {
|
||||
t.Fatalf("profiles missing corp_same:user_1: %#v", resp.Profiles)
|
||||
}
|
||||
current, ok := got["corp_same:user_2"]
|
||||
if !ok {
|
||||
t.Fatalf("profiles missing corp_same:user_2: %#v", resp.Profiles)
|
||||
}
|
||||
if !current.IsOrgCurrent || !current.IsCurrent || current.IsPrimary {
|
||||
t.Fatalf("last login account markers = %#v, want org-current/current and deprecated primary=false", current)
|
||||
}
|
||||
if got["corp_same:user_1"].IsOrgCurrent {
|
||||
t.Fatalf("older account unexpectedly marked org current: %#v", got["corp_same:user_1"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileListUsesRealIdentityTokenState(t *testing.T) {
|
||||
token := authLogoutTestToken("corp_real")
|
||||
token.ExpiresAt = time.Date(2026, 7, 16, 17, 38, 0, 0, time.Local)
|
||||
token.RefreshExpAt = time.Date(2026, 8, 16, 17, 38, 0, 0, time.Local)
|
||||
configDir := setupAuthLogoutProfiles(t, token)
|
||||
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
cfg.Profiles[0].Status = authpkg.ProfileStatusActive
|
||||
cfg.Profiles[0].ExpiresAt = "2026-07-16T22:29:00+08:00"
|
||||
cfg.Profiles[0].RefreshExpAt = "2026-09-16T22:29:00+08:00"
|
||||
if err := authpkg.SaveProfiles(configDir, cfg); err != nil {
|
||||
t.Fatalf("SaveProfiles() error = %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"--format", "json", "profile", "list"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("profile list error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
var resp profileListResponse
|
||||
if err := json.Unmarshal(out.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
if len(resp.Profiles) != 1 {
|
||||
t.Fatalf("profiles len = %d, want 1", len(resp.Profiles))
|
||||
}
|
||||
got := resp.Profiles[0]
|
||||
if got.ExpiresAt != token.ExpiresAt.Format(time.RFC3339) {
|
||||
t.Fatalf("expiresAt = %q, want real token %q", got.ExpiresAt, token.ExpiresAt.Format(time.RFC3339))
|
||||
}
|
||||
if got.RefreshExpAt != token.RefreshExpAt.Format(time.RFC3339) {
|
||||
t.Fatalf("refreshExpAt = %q, want real token %q", got.RefreshExpAt, token.RefreshExpAt.Format(time.RFC3339))
|
||||
}
|
||||
if got.Status != authpkg.ProfileStatusExpired {
|
||||
t.Fatalf("status = %q, want expired", got.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileListDistinguishesMissingAndUnavailableTokenState(t *testing.T) {
|
||||
originalLoad := profileLoadTokenData
|
||||
t.Cleanup(func() { profileLoadTokenData = originalLoad })
|
||||
profile := authpkg.Profile{CorpID: "corp", UserID: "user"}
|
||||
|
||||
profileLoadTokenData = func(string, string) (*authpkg.TokenData, error) {
|
||||
return nil, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
if state := loadProfileTokenState("cfg", profile); state.Status != authpkg.ProfileStatusRevoked {
|
||||
t.Fatalf("missing token status = %q, want revoked", state.Status)
|
||||
}
|
||||
|
||||
profileLoadTokenData = func(string, string) (*authpkg.TokenData, error) {
|
||||
return nil, errors.New("keychain unavailable")
|
||||
}
|
||||
if state := loadProfileTokenState("cfg", profile); state.Status != authpkg.ProfileStatusUnavailable {
|
||||
t.Fatalf("unavailable token status = %q, want unavailable", state.Status)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileListCurrentFlagsUseStoredExactSelectors(t *testing.T) {
|
||||
first := authLogoutTestToken("corp_same")
|
||||
first.UserID = "user_1"
|
||||
first.UserName = "账号一"
|
||||
second := authLogoutTestToken("corp_same")
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
second.UserName = "账号二"
|
||||
configDir := setupAuthLogoutProfiles(t, first, second)
|
||||
|
||||
cfg, err := authpkg.LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
cfg.PrimaryProfile = ""
|
||||
cfg.CurrentProfile = "corp_same:user_1"
|
||||
cfg.OrgCurrentProfiles["corp_same"] = "corp_same:user_1"
|
||||
if err := authpkg.SaveProfiles(configDir, cfg); err != nil {
|
||||
t.Fatalf("SaveProfiles() error = %v", err)
|
||||
}
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"--format", "json", "profile", "list"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("profile list error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
var resp profileListResponse
|
||||
if err := json.Unmarshal(out.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("Unmarshal() error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
got := make(map[string]profileView, len(resp.Profiles))
|
||||
for _, profile := range resp.Profiles {
|
||||
got[profile.Profile] = profile
|
||||
}
|
||||
if !got["corp_same:user_1"].IsCurrent || !got["corp_same:user_1"].IsOrgCurrent {
|
||||
t.Fatalf("first account flags = %#v, want current and org current", got["corp_same:user_1"])
|
||||
}
|
||||
if got["corp_same:user_2"].IsCurrent || got["corp_same:user_2"].IsOrgCurrent {
|
||||
t.Fatalf("second account flags = %#v, want neither current nor org current", got["corp_same:user_2"])
|
||||
}
|
||||
if got["corp_same:user_1"].IsPrimary || got["corp_same:user_2"].IsPrimary {
|
||||
t.Fatalf("deprecated isPrimary should be false without primaryProfile: %#v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileListTableOmitsDeprecatedPrimaryColumn(t *testing.T) {
|
||||
setupAuthLogoutProfiles(t, authLogoutTestToken("corp_table"))
|
||||
|
||||
cmd := NewRootCommand()
|
||||
var out bytes.Buffer
|
||||
cmd.SetOut(&out)
|
||||
cmd.SetErr(&out)
|
||||
cmd.SetArgs([]string{"profile", "list"})
|
||||
if err := cmd.Execute(); err != nil {
|
||||
t.Fatalf("profile list error = %v\noutput:\n%s", err, out.String())
|
||||
}
|
||||
header := strings.SplitN(out.String(), "\n", 2)[0]
|
||||
if strings.Contains(header, "PRI") {
|
||||
t.Fatalf("profile list header still contains deprecated PRI column: %q", header)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileUseRootCommandSwitchesOrganizationAndLegacyMirror(t *testing.T) {
|
||||
configDir := setupAuthLogoutProfiles(t,
|
||||
authLogoutTestToken("corp_primary"),
|
||||
@@ -118,8 +297,9 @@ func TestProfileUseRootCommandSwitchesOrganizationAndLegacyMirror(t *testing.T)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_primary" || cfg.PreviousProfile != "corp_secondary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q, want corp_primary/corp_secondary", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
if cfg.CurrentProfile != "corp_primary:user-corp_primary" ||
|
||||
cfg.PreviousProfile != "corp_secondary:user-corp_secondary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
}
|
||||
legacyToken, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil {
|
||||
@@ -145,8 +325,9 @@ func TestProfileUseRootCommandSwitchesOrganizationAndLegacyMirror(t *testing.T)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_secondary" || cfg.PreviousProfile != "corp_primary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q, want corp_secondary/corp_primary", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
if cfg.CurrentProfile != "corp_secondary:user-corp_secondary" ||
|
||||
cfg.PreviousProfile != "corp_primary:user-corp_primary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
}
|
||||
legacyToken, err = authpkg.LoadTokenData(configDir)
|
||||
if err != nil {
|
||||
@@ -178,8 +359,9 @@ func TestProfileSwitchRootCommandSwitchesPrimaryOrganizationAndLegacyMirror(t *t
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_primary" || cfg.PreviousProfile != "corp_secondary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q, want corp_primary/corp_secondary", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
if cfg.CurrentProfile != "corp_primary:user-corp_primary" ||
|
||||
cfg.PreviousProfile != "corp_secondary:user-corp_secondary" {
|
||||
t.Fatalf("profile pointers = current %q previous %q", cfg.CurrentProfile, cfg.PreviousProfile)
|
||||
}
|
||||
legacyToken, err := authpkg.LoadTokenData(configDir)
|
||||
if err != nil {
|
||||
@@ -209,8 +391,8 @@ func TestProfileSwitchRootCommandSupportsCorpIDFlag(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary", cfg.CurrentProfile)
|
||||
if cfg.CurrentProfile != "corp_primary:user-corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary:user-corp_primary", cfg.CurrentProfile)
|
||||
}
|
||||
|
||||
cmd = NewRootCommand()
|
||||
@@ -226,8 +408,8 @@ func TestProfileSwitchRootCommandSupportsCorpIDFlag(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_secondary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_secondary", cfg.CurrentProfile)
|
||||
if cfg.CurrentProfile != "corp_secondary:user-corp_secondary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_secondary:user-corp_secondary", cfg.CurrentProfile)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -287,8 +469,8 @@ func TestProfileSwitchNoArgsUsesTUISelector(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary", cfg.CurrentProfile)
|
||||
if cfg.CurrentProfile != "corp_primary:user-corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary:user-corp_primary", cfg.CurrentProfile)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -358,7 +540,7 @@ func TestProfileSwitchTUIViewUsesFixedOuterTable(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestProfileSwitchTUISortsLatestLoggedInProfilesFirst(t *testing.T) {
|
||||
func TestProfileSwitchTUIPreservesStoredOrderInsteadOfSortingByTime(t *testing.T) {
|
||||
cfg := &authpkg.ProfilesConfig{
|
||||
PrimaryProfile: "old",
|
||||
CurrentProfile: "old",
|
||||
@@ -370,7 +552,7 @@ func TestProfileSwitchTUISortsLatestLoggedInProfilesFirst(t *testing.T) {
|
||||
}
|
||||
model := newProfileSwitchTUIModel(cfg, "old")
|
||||
gotOrder := []string{model.profiles[0].CorpID, model.profiles[1].CorpID, model.profiles[2].CorpID}
|
||||
wantOrder := []string{"new", "fallback", "old"}
|
||||
wantOrder := []string{"old", "new", "fallback"}
|
||||
if strings.Join(gotOrder, ",") != strings.Join(wantOrder, ",") {
|
||||
t.Fatalf("profile order = %v, want %v", gotOrder, wantOrder)
|
||||
}
|
||||
@@ -507,8 +689,8 @@ func TestProfileUseNoArgsUsesTUISelector(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
if cfg.CurrentProfile != "corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary", cfg.CurrentProfile)
|
||||
if cfg.CurrentProfile != "corp_primary:user-corp_primary" {
|
||||
t.Fatalf("currentProfile = %q, want corp_primary:user-corp_primary", cfg.CurrentProfile)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -549,7 +731,7 @@ func TestWriteProfileListTableIncludesCorpName(t *testing.T) {
|
||||
},
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
writeProfileListTable(&buf, cfg)
|
||||
writeProfileListTable(&buf, "", cfg)
|
||||
out := buf.String()
|
||||
for _, want := range []string{
|
||||
"ORG_NAME",
|
||||
|
||||
@@ -105,7 +105,7 @@ func TestCrossPlatformCoverageProfileRemainingCoverage(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if len(choices) != 3 || choices[0] != "current" || choices[1] != "primary" || choices[2] != "first" {
|
||||
if len(choices) != 3 || choices[0] != "current" || choices[1] != "first" || choices[2] != "first" {
|
||||
t.Fatalf("profile choices = %#v", choices)
|
||||
}
|
||||
|
||||
|
||||
@@ -333,7 +333,11 @@ func (r *recoveryRuntime) CallToolDirect(ctx context.Context, serverID, toolName
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tc := r.transport.WithAuth(resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags)), resolveIdentityHeaders())
|
||||
authToken, err := resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags))
|
||||
if err != nil {
|
||||
return nil, tokenResolutionError(err)
|
||||
}
|
||||
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
result, err := tc.CallTool(ctx, endpoint, toolName, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
+324
-33
@@ -23,6 +23,7 @@ import (
|
||||
"os"
|
||||
"os/signal"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
@@ -69,7 +70,8 @@ var (
|
||||
rootPluginDescriptors = (*plugin.Plugin).ToServerDescriptors
|
||||
rootPluginStdioClients = (*plugin.Plugin).StdioClients
|
||||
rootRegisterPluginHTTPServer = registerPluginHTTPServer
|
||||
rootRegisterStdioManifest = registerStdioServerFromManifest
|
||||
rootPluginStdioDescriptor = stdioServerDescriptorFromManifest
|
||||
rootRegisterResolvedStdioServer = registerResolvedStdioServer
|
||||
rootPluginLoadHooks = (*plugin.Plugin).LoadHooks
|
||||
rootPluginSyncSkills = plugin.SyncSkills
|
||||
rootAuthLoadTokenData = authpkg.LoadTokenData
|
||||
@@ -306,13 +308,28 @@ func NewRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
if len(ctx) > 0 && ctx[0] != nil {
|
||||
rootCtx = ctx[0]
|
||||
}
|
||||
return NewRootCommandWithEngine(rootCtx, nil)
|
||||
return newRootCommandWithEngine(rootCtx, nil, true)
|
||||
}
|
||||
|
||||
// NewSchemaSourceRootCommand constructs the distribution-owned command tree
|
||||
// used by Schema generation and command-surface policy. Installed plugins and
|
||||
// user-defined shortcuts must not change the reviewed embedded Schema.
|
||||
func NewSchemaSourceRootCommand(ctx ...context.Context) *cobra.Command {
|
||||
var rootCtx context.Context
|
||||
if len(ctx) > 0 && ctx[0] != nil {
|
||||
rootCtx = ctx[0]
|
||||
}
|
||||
return newRootCommandWithEngine(rootCtx, nil, false)
|
||||
}
|
||||
|
||||
// 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 {
|
||||
return newRootCommandWithEngine(rootCtx, engine, true)
|
||||
}
|
||||
|
||||
func newRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine, loadRuntimeExtensions bool) *cobra.Command {
|
||||
if rootCtx == nil {
|
||||
rootCtx = context.Background()
|
||||
}
|
||||
@@ -396,16 +413,9 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
}
|
||||
root.AddCommand(utilityCommands...)
|
||||
|
||||
root.AddCommand(newLegacyPublicCommands(runner, patCaller)...)
|
||||
root.AddCommand(newLegacyPublicCommands(runner, patCaller, loadRuntimeExtensions)...)
|
||||
root.AddCommand(newLegacyHiddenCommands(runner)...)
|
||||
|
||||
// --- Plugin loading: runs AFTER legacy commands so plugin endpoints can
|
||||
// be appended on top of the static endpoint registry.
|
||||
pluginCmds := rootLoadPlugins(engine, runner)
|
||||
if len(pluginCmds) > 0 {
|
||||
addPluginCommandsSafe(root, pluginCmds)
|
||||
}
|
||||
|
||||
// PAT authorization commands (open-source core)
|
||||
pat.RegisterCommands(root, patCaller)
|
||||
|
||||
@@ -414,6 +424,15 @@ func NewRootCommandWithEngine(rootCtx context.Context, engine *pipeline.Engine)
|
||||
fn(root, caller)
|
||||
deduplicateCommands(root)
|
||||
}
|
||||
if loadRuntimeExtensions {
|
||||
// Resolve plugins only after the complete distribution command tree is
|
||||
// present, so endpoint and Cobra conflict checks see PAT and edition
|
||||
// commands as well as the open-source base.
|
||||
pluginCmds := rootLoadPlugins(root, engine, runner)
|
||||
if len(pluginCmds) > 0 {
|
||||
addPluginCommandsSafe(root, pluginCmds)
|
||||
}
|
||||
}
|
||||
hideNonDirectRuntimeCommands(root)
|
||||
configureRootHelp(root)
|
||||
// Set custom flag error handler for better UX
|
||||
@@ -631,12 +650,17 @@ var reservedCommands = map[string]bool{
|
||||
"schema": true, "mcp": true, "help": true,
|
||||
}
|
||||
|
||||
var replaceablePluginFallbacks = map[string]bool{
|
||||
"conference": true,
|
||||
}
|
||||
|
||||
// addPluginCommandsSafe registers plugin commands with conflict detection.
|
||||
//
|
||||
// Rules:
|
||||
// - Plugin vs reserved (auth/plugin/cache/...) → reject, warn
|
||||
// - Plugin vs plugin (same name) → reject later one, warn
|
||||
// - Plugin vs Market dynamic command → allow, plugin wins
|
||||
// - Plugin vs hidden compatibility fallback → allow, plugin wins
|
||||
// - Plugin vs visible distribution command → reject, warn
|
||||
func addPluginCommandsSafe(root *cobra.Command, pluginCmds []*cobra.Command) {
|
||||
// Build index of existing commands before plugin registration.
|
||||
existing := make(map[string]bool)
|
||||
@@ -664,17 +688,47 @@ func addPluginCommandsSafe(root *cobra.Command, pluginCmds []*cobra.Command) {
|
||||
}
|
||||
pluginSeen[name] = true
|
||||
|
||||
// Rule 3: plugin vs Market — plugin wins, remove the old one.
|
||||
// An alias must not bypass the same protections applied to primary
|
||||
// plugin command names or shadow another root command.
|
||||
filteredAliases := make([]string, 0, len(cmd.Aliases))
|
||||
for _, rawAlias := range cmd.Aliases {
|
||||
alias := strings.TrimSpace(rawAlias)
|
||||
if alias == "" || alias == name || reservedCommands[alias] ||
|
||||
existing[alias] || pluginSeen[alias] {
|
||||
if alias != "" {
|
||||
slog.Warn("plugin: command alias conflicts with an existing command, skipping",
|
||||
"command", name, "alias", alias)
|
||||
}
|
||||
continue
|
||||
}
|
||||
pluginSeen[alias] = true
|
||||
filteredAliases = append(filteredAliases, alias)
|
||||
}
|
||||
cmd.Aliases = filteredAliases
|
||||
|
||||
// Rule 3: an installed plugin may replace a hidden compatibility
|
||||
// fallback (for example conference), but never a visible distribution
|
||||
// command that participates in the reviewed base interface.
|
||||
if existing[name] {
|
||||
for _, old := range root.Commands() {
|
||||
if old.Name() == name {
|
||||
if !old.Hidden || !replaceablePluginFallbacks[name] ||
|
||||
cmdutil.IsPluginSourced(old) {
|
||||
slog.Warn("plugin: command conflicts with a visible distribution command, skipping",
|
||||
"command", name)
|
||||
cmd = nil
|
||||
break
|
||||
}
|
||||
root.RemoveCommand(old)
|
||||
slog.Debug("plugin: overriding Market command",
|
||||
slog.Debug("plugin: overriding hidden compatibility command",
|
||||
"command", name)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if cmd == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
root.AddCommand(cmd)
|
||||
}
|
||||
@@ -811,7 +865,21 @@ func CloseFileLogger() {
|
||||
// loadPlugins registers versioned plugin manifests, stdio clients, hooks, and
|
||||
// skills. It deliberately does not initialize MCP transports or call
|
||||
// tools/list while constructing the command tree.
|
||||
func loadPlugins(engine *pipeline.Engine, _ executor.Runner) []*cobra.Command {
|
||||
type pluginServerCandidate struct {
|
||||
owner *plugin.Plugin
|
||||
order int
|
||||
descriptor mcptypes.ServerDescriptor
|
||||
stdioClient *plugin.StdioServerClient
|
||||
}
|
||||
|
||||
type pluginIdentityOwner struct {
|
||||
plugin *plugin.Plugin
|
||||
serverKey string
|
||||
rootName string
|
||||
shareable bool
|
||||
}
|
||||
|
||||
func loadPlugins(root *cobra.Command, engine *pipeline.Engine, runner executor.Runner) []*cobra.Command {
|
||||
pluginLoader := plugin.NewLoader(RawVersion())
|
||||
|
||||
// 0a. Inject plugin config values from settings.json as environment
|
||||
@@ -838,25 +906,34 @@ func loadPlugins(engine *pipeline.Engine, _ executor.Runner) []*cobra.Command {
|
||||
|
||||
// 2. Load dev plugins (registered via `dws plugin dev`)
|
||||
devPlugins := rootPluginLoadDev(pluginLoader)
|
||||
sortPluginsForRegistration(userPlugins)
|
||||
sortPluginsForRegistration(devPlugins)
|
||||
|
||||
allPlugins := append(userPlugins, devPlugins...)
|
||||
descriptorsByPlugin := make(map[*plugin.Plugin][]mcptypes.ServerDescriptor, len(allPlugins))
|
||||
|
||||
// 3. Register HTTP descriptors and authentication from the manifest.
|
||||
for _, p := range allPlugins {
|
||||
for _, srv := range rootPluginDescriptors(p) {
|
||||
rootRegisterPluginHTTPServer(srv)
|
||||
// 3. Resolve every descriptor once, then choose identity winners before
|
||||
// mutating endpoint, auth, or stdio-client registries. This keeps the
|
||||
// visible command and its transport owned by the same plugin.
|
||||
candidates := collectPluginServerCandidates(allPlugins, userCtx)
|
||||
accepted := selectPluginServerCandidates(root, candidates)
|
||||
for _, candidate := range accepted {
|
||||
if candidate.stdioClient != nil {
|
||||
rootRegisterResolvedStdioServer(
|
||||
candidate.owner,
|
||||
*candidate.stdioClient,
|
||||
candidate.descriptor,
|
||||
)
|
||||
} else {
|
||||
rootRegisterPluginHTTPServer(candidate.descriptor)
|
||||
}
|
||||
descriptorsByPlugin[candidate.owner] = append(
|
||||
descriptorsByPlugin[candidate.owner],
|
||||
candidate.descriptor,
|
||||
)
|
||||
}
|
||||
|
||||
// 4. Register stdio descriptors and unstarted clients. The subprocess is
|
||||
// started and initialized only when a command is actually executed.
|
||||
for _, p := range allPlugins {
|
||||
for _, sc := range rootPluginStdioClients(p, userCtx) {
|
||||
rootRegisterStdioManifest(p, sc)
|
||||
}
|
||||
}
|
||||
|
||||
// 5. Register plugin hooks into pipeline engine
|
||||
// 4. Register plugin hooks into pipeline engine
|
||||
if engine != nil {
|
||||
for _, p := range allPlugins {
|
||||
hooksCfg, err := rootPluginLoadHooks(p)
|
||||
@@ -874,7 +951,7 @@ func loadPlugins(engine *pipeline.Engine, _ executor.Runner) []*cobra.Command {
|
||||
}
|
||||
}
|
||||
|
||||
// 7. Sync plugin skills to agent directories
|
||||
// 5. Sync plugin skills to agent directories
|
||||
rootPluginSyncSkills(allPlugins)
|
||||
|
||||
if len(allPlugins) > 0 {
|
||||
@@ -884,11 +961,228 @@ func loadPlugins(engine *pipeline.Engine, _ executor.Runner) []*cobra.Command {
|
||||
)
|
||||
}
|
||||
|
||||
return nil
|
||||
var pluginCommands []*cobra.Command
|
||||
for _, p := range allPlugins {
|
||||
// Build each plugin independently. addPluginCommandsSafe deliberately
|
||||
// resolves cross-plugin root conflicts with first-plugin-wins semantics.
|
||||
pluginCommands = append(pluginCommands, buildPluginCommands(descriptorsByPlugin[p], runner, root)...)
|
||||
}
|
||||
return pluginCommands
|
||||
}
|
||||
|
||||
func sortPluginsForRegistration(plugins []*plugin.Plugin) {
|
||||
sort.SliceStable(plugins, func(i, j int) bool {
|
||||
left := strings.TrimSpace(plugins[i].Manifest.Name) + "\x00" + strings.TrimSpace(plugins[i].Root)
|
||||
right := strings.TrimSpace(plugins[j].Manifest.Name) + "\x00" + strings.TrimSpace(plugins[j].Root)
|
||||
return left < right
|
||||
})
|
||||
}
|
||||
|
||||
func collectPluginServerCandidates(
|
||||
plugins []*plugin.Plugin,
|
||||
userCtx *plugin.UserContext,
|
||||
) []pluginServerCandidate {
|
||||
var candidates []pluginServerCandidate
|
||||
for order, owner := range plugins {
|
||||
for _, descriptor := range rootPluginDescriptors(owner) {
|
||||
candidates = append(candidates, pluginServerCandidate{
|
||||
owner: owner,
|
||||
order: order,
|
||||
descriptor: descriptor,
|
||||
})
|
||||
}
|
||||
for _, stdioClient := range rootPluginStdioClients(owner, userCtx) {
|
||||
descriptor, ok := rootPluginStdioDescriptor(owner, stdioClient)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
clientCopy := stdioClient
|
||||
candidates = append(candidates, pluginServerCandidate{
|
||||
owner: owner,
|
||||
order: order,
|
||||
descriptor: descriptor,
|
||||
stdioClient: &clientCopy,
|
||||
})
|
||||
}
|
||||
}
|
||||
sort.SliceStable(candidates, func(i, j int) bool {
|
||||
if candidates[i].order != candidates[j].order {
|
||||
return candidates[i].order < candidates[j].order
|
||||
}
|
||||
left := strings.TrimSpace(candidates[i].descriptor.Key)
|
||||
right := strings.TrimSpace(candidates[j].descriptor.Key)
|
||||
if left != right {
|
||||
return left < right
|
||||
}
|
||||
return candidates[i].stdioClient == nil && candidates[j].stdioClient != nil
|
||||
})
|
||||
return candidates
|
||||
}
|
||||
|
||||
func selectPluginServerCandidates(
|
||||
root *cobra.Command,
|
||||
candidates []pluginServerCandidate,
|
||||
) []pluginServerCandidate {
|
||||
distributionProducts := DirectRuntimeProductIDs()
|
||||
owners := make(map[string]pluginIdentityOwner)
|
||||
for identity := range distributionProducts {
|
||||
if replaceablePluginFallbacks[identity] {
|
||||
continue
|
||||
}
|
||||
owners[identity] = pluginIdentityOwner{serverKey: "distribution"}
|
||||
}
|
||||
|
||||
accepted := make([]pluginServerCandidate, 0, len(candidates))
|
||||
for _, candidate := range candidates {
|
||||
descriptor := candidate.descriptor
|
||||
if descriptor.CLI.Skip {
|
||||
continue
|
||||
}
|
||||
if reason := unsupportedPluginDescriptor(root, descriptor); reason != "" {
|
||||
slog.Warn("plugin: descriptor CLI semantics are unsupported, skipping",
|
||||
"plugin", candidate.owner.Manifest.Name,
|
||||
"server", descriptor.Key,
|
||||
"field", reason)
|
||||
continue
|
||||
}
|
||||
if pluginDescriptorConflictsWithDistribution(root, descriptor, distributionProducts) {
|
||||
continue
|
||||
}
|
||||
claims := pluginDescriptorIdentityClaims(descriptor)
|
||||
conflict := ""
|
||||
for identity, shareable := range claims {
|
||||
existing, exists := owners[identity]
|
||||
if !exists {
|
||||
continue
|
||||
}
|
||||
rootName := pluginDescriptorRootName(descriptor)
|
||||
if shareable && existing.shareable &&
|
||||
existing.plugin == candidate.owner &&
|
||||
existing.rootName == rootName {
|
||||
continue
|
||||
}
|
||||
conflict = identity
|
||||
break
|
||||
}
|
||||
if conflict != "" {
|
||||
slog.Warn("plugin: descriptor identity already owned, skipping",
|
||||
"plugin", candidate.owner.Manifest.Name,
|
||||
"server", descriptor.Key,
|
||||
"identity", conflict)
|
||||
continue
|
||||
}
|
||||
rootName := pluginDescriptorRootName(descriptor)
|
||||
for identity, shareable := range claims {
|
||||
if existing, exists := owners[identity]; exists &&
|
||||
shareable && existing.shareable &&
|
||||
existing.plugin == candidate.owner &&
|
||||
existing.rootName == rootName {
|
||||
continue
|
||||
}
|
||||
owners[identity] = pluginIdentityOwner{
|
||||
plugin: candidate.owner,
|
||||
serverKey: descriptor.Key,
|
||||
rootName: rootName,
|
||||
shareable: shareable,
|
||||
}
|
||||
}
|
||||
accepted = append(accepted, candidate)
|
||||
}
|
||||
return accepted
|
||||
}
|
||||
|
||||
func pluginDescriptorIdentityClaims(descriptor mcptypes.ServerDescriptor) map[string]bool {
|
||||
claims := make(map[string]bool)
|
||||
canonicalID := firstNonEmptyPluginString(descriptor.CLI.ID, descriptor.Key)
|
||||
if canonicalID != "" {
|
||||
claims[canonicalID] = false
|
||||
}
|
||||
for _, identity := range append(
|
||||
[]string{pluginDescriptorRootName(descriptor)},
|
||||
descriptor.CLI.Aliases...,
|
||||
) {
|
||||
identity = strings.TrimSpace(identity)
|
||||
if identity == "" {
|
||||
continue
|
||||
}
|
||||
if _, exists := claims[identity]; !exists {
|
||||
claims[identity] = true
|
||||
}
|
||||
}
|
||||
return claims
|
||||
}
|
||||
|
||||
func pluginDescriptorRootName(descriptor mcptypes.ServerDescriptor) string {
|
||||
return firstNonEmptyPluginString(
|
||||
descriptor.CLI.Command,
|
||||
descriptor.CLI.ID,
|
||||
descriptor.Key,
|
||||
)
|
||||
}
|
||||
|
||||
func pluginDescriptorConflictsWithDistribution(
|
||||
root *cobra.Command,
|
||||
descriptor mcptypes.ServerDescriptor,
|
||||
distributionProducts map[string]bool,
|
||||
) bool {
|
||||
candidates := append(
|
||||
[]string{
|
||||
firstNonEmptyPluginString(descriptor.CLI.ID, descriptor.Key),
|
||||
pluginDescriptorRootName(descriptor),
|
||||
},
|
||||
descriptor.CLI.Aliases...,
|
||||
)
|
||||
for _, candidate := range candidates {
|
||||
candidate = strings.TrimSpace(candidate)
|
||||
if candidate == "" {
|
||||
continue
|
||||
}
|
||||
if !reservedCommands[candidate] && replaceablePluginFallbacks[candidate] {
|
||||
// The distribution ships only a hidden compatibility fallback for
|
||||
// this name; plugins may claim it and the later command merge in
|
||||
// addPluginCommandsSafe still rejects visible non-fallback owners.
|
||||
continue
|
||||
}
|
||||
if reservedCommands[candidate] ||
|
||||
distributionProducts[candidate] ||
|
||||
distributionRootOwns(root, candidate) {
|
||||
slog.Warn("plugin: descriptor conflicts with a distribution command, skipping",
|
||||
"plugin", descriptor.DisplayName,
|
||||
"server", descriptor.Key,
|
||||
"identity", candidate)
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func distributionRootOwns(root *cobra.Command, name string) bool {
|
||||
if root == nil {
|
||||
return false
|
||||
}
|
||||
for _, command := range root.Commands() {
|
||||
if cmdutil.IsPluginSourced(command) {
|
||||
continue
|
||||
}
|
||||
if command.Name() == name {
|
||||
if command.Hidden && replaceablePluginFallbacks[name] {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
for _, alias := range command.Aliases {
|
||||
if strings.TrimSpace(alias) == name {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func registerPluginHTTPServer(srv mcptypes.ServerDescriptor) {
|
||||
AppendDynamicServer(srv)
|
||||
productID := firstNonEmptyPluginString(srv.CLI.ID, srv.Key)
|
||||
ClearPluginAuth(productID)
|
||||
if len(srv.AuthHeaders) > 0 {
|
||||
registerPluginAuthFromHeaders(srv)
|
||||
}
|
||||
@@ -917,10 +1211,7 @@ func registerPluginAuthFromHeaders(srv mcptypes.ServerDescriptor) {
|
||||
host := parsed.Hostname()
|
||||
trustedDomains = []string{host, "*." + host}
|
||||
}
|
||||
productID := strings.TrimSpace(srv.CLI.ID)
|
||||
if productID == "" {
|
||||
productID = srv.Key
|
||||
}
|
||||
productID := firstNonEmptyPluginString(srv.CLI.ID, srv.Key)
|
||||
RegisterPluginAuth(productID, &PluginAuth{
|
||||
Token: authToken,
|
||||
ExtraHeaders: extraHeaders,
|
||||
|
||||
@@ -75,7 +75,7 @@ func TestCrossPlatformCoverageRootConstructionHooksAndVersionCoverage(t *testing
|
||||
version, buildTime, gitCommit = oldVersion, oldBuild, oldCommit
|
||||
})
|
||||
|
||||
rootLoadPlugins = func(*pipeline.Engine, executor.Runner) []*cobra.Command {
|
||||
rootLoadPlugins = func(*cobra.Command, *pipeline.Engine, executor.Runner) []*cobra.Command {
|
||||
return []*cobra.Command{{Use: "plugin-added", Run: func(*cobra.Command, []string) {}}}
|
||||
}
|
||||
preRunCalled := false
|
||||
@@ -236,7 +236,8 @@ func TestCrossPlatformCoverageRootLoadPluginsRemainingCoverage(t *testing.T) {
|
||||
oldDescriptors := rootPluginDescriptors
|
||||
oldStdioClients := rootPluginStdioClients
|
||||
oldHTTP := rootRegisterPluginHTTPServer
|
||||
oldStdio := rootRegisterStdioManifest
|
||||
oldStdioDescriptor := rootPluginStdioDescriptor
|
||||
oldStdioRegister := rootRegisterResolvedStdioServer
|
||||
oldHooks := rootPluginLoadHooks
|
||||
oldSync := rootPluginSyncSkills
|
||||
oldToken := rootAuthLoadTokenData
|
||||
@@ -247,7 +248,8 @@ func TestCrossPlatformCoverageRootLoadPluginsRemainingCoverage(t *testing.T) {
|
||||
rootPluginDescriptors = oldDescriptors
|
||||
rootPluginStdioClients = oldStdioClients
|
||||
rootRegisterPluginHTTPServer = oldHTTP
|
||||
rootRegisterStdioManifest = oldStdio
|
||||
rootPluginStdioDescriptor = oldStdioDescriptor
|
||||
rootRegisterResolvedStdioServer = oldStdioRegister
|
||||
rootPluginLoadHooks = oldHooks
|
||||
rootPluginSyncSkills = oldSync
|
||||
rootAuthLoadTokenData = oldToken
|
||||
@@ -264,9 +266,17 @@ func TestCrossPlatformCoverageRootLoadPluginsRemainingCoverage(t *testing.T) {
|
||||
}
|
||||
rootPluginDescriptors = func(p *plugin.Plugin) []mcptypes.ServerDescriptor {
|
||||
if p == p1 {
|
||||
return []mcptypes.ServerDescriptor{{Key: "http", Endpoint: "https://example.test"}}
|
||||
return []mcptypes.ServerDescriptor{{
|
||||
Key: "http", Endpoint: "https://example.test",
|
||||
CLI: mcptypes.CLIOverlay{
|
||||
ID: "http", Command: "one-http",
|
||||
ToolOverrides: map[string]mcptypes.CLIToolOverride{
|
||||
"ping": {CLIName: "ping"},
|
||||
},
|
||||
},
|
||||
}}
|
||||
}
|
||||
return []mcptypes.ServerDescriptor{{Key: "no-cli", Endpoint: "https://example.test"}}
|
||||
return []mcptypes.ServerDescriptor{{Key: p.Manifest.Name + "-no-cli", Endpoint: "https://example.test"}}
|
||||
}
|
||||
client := transport.NewStdioClient("ignored", nil, nil)
|
||||
rootPluginStdioClients = func(p *plugin.Plugin, uc *plugin.UserContext) []plugin.StdioServerClient {
|
||||
@@ -278,9 +288,23 @@ func TestCrossPlatformCoverageRootLoadPluginsRemainingCoverage(t *testing.T) {
|
||||
httpCount := 0
|
||||
stdioCount := 0
|
||||
rootRegisterPluginHTTPServer = func(mcptypes.ServerDescriptor) { httpCount++ }
|
||||
rootRegisterStdioManifest = func(*plugin.Plugin, plugin.StdioServerClient) mcptypes.ServerDescriptor {
|
||||
rootPluginStdioDescriptor = func(*plugin.Plugin, plugin.StdioServerClient) (mcptypes.ServerDescriptor, bool) {
|
||||
return mcptypes.ServerDescriptor{
|
||||
Key: "local",
|
||||
CLI: mcptypes.CLIOverlay{
|
||||
ID: "local", Command: "one-stdio",
|
||||
ToolOverrides: map[string]mcptypes.CLIToolOverride{
|
||||
"pong": {CLIName: "pong"},
|
||||
},
|
||||
},
|
||||
}, true
|
||||
}
|
||||
rootRegisterResolvedStdioServer = func(
|
||||
*plugin.Plugin,
|
||||
plugin.StdioServerClient,
|
||||
mcptypes.ServerDescriptor,
|
||||
) {
|
||||
stdioCount++
|
||||
return mcptypes.ServerDescriptor{}
|
||||
}
|
||||
rootPluginLoadHooks = func(p *plugin.Plugin) (*plugin.HooksConfig, error) {
|
||||
switch p {
|
||||
@@ -294,7 +318,8 @@ func TestCrossPlatformCoverageRootLoadPluginsRemainingCoverage(t *testing.T) {
|
||||
}
|
||||
synced := false
|
||||
rootPluginSyncSkills = func([]*plugin.Plugin) { synced = true }
|
||||
if got := loadPlugins(pipeline.NewEngine(), runnerCoverageFallback{}); got != nil {
|
||||
got := loadPlugins(nil, pipeline.NewEngine(), runnerCoverageFallback{})
|
||||
if len(got) != 2 || got[0].Name() != "one-http" || got[1].Name() != "one-stdio" {
|
||||
t.Fatalf("loaded plugin commands = %#v", got)
|
||||
}
|
||||
if httpCount != 3 || stdioCount != 1 || !synced {
|
||||
|
||||
+122
-61
@@ -168,6 +168,7 @@ var (
|
||||
runnerPreflightDocDownload = (*runtimeRunner).preflightDocDownload
|
||||
runnerCallTool = (*transport.Client).CallTool
|
||||
runnerStdioEnsureInitialized = (*transport.StdioClient).EnsureInitialized
|
||||
runnerStdioListTools = (*transport.StdioClient).ListTools
|
||||
runnerStdioCallTool = (*transport.StdioClient).CallTool
|
||||
runnerHandlePatAuthCheck func(context.Context, *runtimeRunner, executor.Invocation, *apperrors.PATError, string, io.Writer) (executor.Result, error)
|
||||
runnerRetryWithPatAuthRetry func(context.Context, executor.Runner, executor.Invocation, *PatScopeError, string, io.Writer) (executor.Result, error)
|
||||
@@ -193,13 +194,25 @@ func (r *runtimeRunner) Run(ctx context.Context, invocation executor.Invocation)
|
||||
// invocations within the same process free.
|
||||
logHostOwnedPATDecisionOnce()
|
||||
|
||||
selections, multi, err := runnerResolveMultiProfileSelections(defaultConfigDir(), authpkg.RuntimeProfile())
|
||||
rawProfile := authpkg.RuntimeProfile()
|
||||
selections, multi, err := runnerResolveMultiProfileSelections(defaultConfigDir(), rawProfile)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewValidation(err.Error())
|
||||
}
|
||||
if multi {
|
||||
return r.runMultiProfile(ctx, invocation, selections)
|
||||
}
|
||||
if strings.TrimSpace(rawProfile) != "" {
|
||||
profile, err := authpkg.ResolveProfile(defaultConfigDir(), rawProfile)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewValidation(err.Error())
|
||||
}
|
||||
if profile == nil {
|
||||
return executor.Result{}, apperrors.NewValidation(fmt.Sprintf("profile %q not found", rawProfile))
|
||||
}
|
||||
authpkg.SetRuntimeProfile(authpkg.ProfileSelector(*profile))
|
||||
defer authpkg.SetRuntimeProfile(rawProfile)
|
||||
}
|
||||
|
||||
return r.runSingle(ctx, invocation, true)
|
||||
}
|
||||
@@ -223,7 +236,9 @@ func (r *runtimeRunner) runSingle(ctx context.Context, invocation executor.Invoc
|
||||
// ~70ms on macOS; starting it here lets the load overlap with endpoint
|
||||
// resolution and catalog loading below.
|
||||
if prefetchToken {
|
||||
go runnerGetCachedRuntimeToken(ctx)
|
||||
go func() {
|
||||
_, _ = runnerGetCachedRuntimeToken(ctx)
|
||||
}()
|
||||
}
|
||||
|
||||
if shouldUseDirectRuntime(invocation) {
|
||||
@@ -313,10 +328,11 @@ func resolveMultiProfileSelections(configDir, rawSelector string) ([]multiProfil
|
||||
if profile == nil {
|
||||
return nil, false, fmt.Errorf("profile %q not found", selector)
|
||||
}
|
||||
if seen[profile.CorpID] {
|
||||
identitySelector := authpkg.ProfileSelector(*profile)
|
||||
if seen[identitySelector] {
|
||||
continue
|
||||
}
|
||||
seen[profile.CorpID] = true
|
||||
seen[identitySelector] = true
|
||||
selections = append(selections, multiProfileSelection{
|
||||
Selector: selector,
|
||||
Profile: *profile,
|
||||
@@ -334,13 +350,17 @@ func (r *runtimeRunner) runMultiProfile(ctx context.Context, invocation executor
|
||||
failed := 0
|
||||
|
||||
for _, selection := range selections {
|
||||
authpkg.SetRuntimeProfile(selection.Profile.CorpID)
|
||||
resolvedSelector := authpkg.ProfileSelector(selection.Profile)
|
||||
authpkg.SetRuntimeProfile(resolvedSelector)
|
||||
result, err := r.runSingle(ctx, cloneInvocation(invocation), false)
|
||||
|
||||
entry := map[string]any{
|
||||
"selector": selection.Selector,
|
||||
"profile": resolvedSelector,
|
||||
"corpId": selection.Profile.CorpID,
|
||||
"corpName": selection.Profile.CorpName,
|
||||
"userId": selection.Profile.UserID,
|
||||
"userName": selection.Profile.UserName,
|
||||
"ok": err == nil,
|
||||
}
|
||||
if err != nil {
|
||||
@@ -466,7 +486,7 @@ func (r *runtimeRunner) handleCatalogMiss(ctx context.Context, invocation execut
|
||||
func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string, invocation executor.Invocation) (result executor.Result, retErr error) {
|
||||
// Route stdio:// endpoints to the local StdioClient — no HTTP, no auth.
|
||||
if IsStdioEndpoint(endpoint) {
|
||||
return r.executeStdioInvocation(ctx, invocation)
|
||||
return r.executeStdioInvocationAtEndpoint(ctx, endpoint, invocation)
|
||||
}
|
||||
|
||||
// Constructing the Cobra tree is also used for help, schema, and command
|
||||
@@ -517,8 +537,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
authToken := ""
|
||||
if hasPluginAuth {
|
||||
authToken = pluginAuth.Token
|
||||
} else {
|
||||
authToken = r.resolveAuthToken(ctx)
|
||||
} else if !invocation.DryRun && (r.globalFlags == nil || !r.globalFlags.Mock) {
|
||||
var tokenErr error
|
||||
authToken, tokenErr = r.resolveAuthToken(ctx)
|
||||
if tokenErr != nil {
|
||||
return executor.Result{}, tokenResolutionError(tokenErr)
|
||||
}
|
||||
}
|
||||
|
||||
var timeoutSec int
|
||||
@@ -600,6 +624,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}
|
||||
return runnerHandlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, err, hasPluginAuth); handled {
|
||||
if retryErr != nil {
|
||||
runnerCaptureRuntimeFailure(invocation, err, retryErr)
|
||||
}
|
||||
return result, retryErr
|
||||
}
|
||||
runnerCaptureRuntimeFailure(invocation, err, err)
|
||||
return executor.Result{}, err
|
||||
}
|
||||
@@ -608,9 +638,15 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
callResult, err := runnerCallTool(tc, callCtx, endpoint, invocation.Tool, invocation.Params)
|
||||
RecordTiming(ctx, "mcp_call", time.Since(callStart))
|
||||
if err != nil {
|
||||
if isAuthError(err) {
|
||||
if isRefreshableTransportAuthError(err) {
|
||||
if fn := edition.Get().OnAuthError; fn != nil {
|
||||
if overrideErr := fn(defaultConfigDir(), err); overrideErr != nil {
|
||||
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, overrideErr, hasPluginAuth); handled {
|
||||
if retryErr != nil {
|
||||
runnerCaptureRuntimeFailure(invocation, err, retryErr)
|
||||
}
|
||||
return result, retryErr
|
||||
}
|
||||
runnerCaptureRuntimeFailure(invocation, err, overrideErr)
|
||||
return executor.Result{}, overrideErr
|
||||
}
|
||||
@@ -635,6 +671,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
}
|
||||
return runnerHandlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
|
||||
}
|
||||
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, editionErr, hasPluginAuth); handled {
|
||||
if retryErr != nil {
|
||||
runnerCaptureRuntimeFailure(invocation, editionErr, retryErr)
|
||||
}
|
||||
return result, retryErr
|
||||
}
|
||||
return executor.Result{}, editionErr
|
||||
}
|
||||
}
|
||||
@@ -655,6 +697,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
// patterns (PAT permission, gateway-auth) before generic handling.
|
||||
if classify := edition.Get().ClassifyToolResult; classify != nil {
|
||||
if hookErr := classify(callResult.Content); hookErr != nil {
|
||||
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, hookErr, hasPluginAuth); handled {
|
||||
if retryErr != nil {
|
||||
runnerCaptureRuntimeFailure(invocation, hookErr, retryErr)
|
||||
}
|
||||
return result, retryErr
|
||||
}
|
||||
runnerCaptureRuntimeFailure(invocation, hookErr, hookErr)
|
||||
return executor.Result{}, hookErr
|
||||
}
|
||||
@@ -719,6 +767,14 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
// subprocess instead of the HTTP transport. This is used for plugin stdio
|
||||
// servers whose endpoints use the stdio:// scheme.
|
||||
func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
return r.executeStdioInvocationAtEndpoint(ctx, "", invocation)
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) executeStdioInvocationAtEndpoint(
|
||||
ctx context.Context,
|
||||
endpoint string,
|
||||
invocation executor.Invocation,
|
||||
) (executor.Result, error) {
|
||||
if invocation.DryRun {
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
@@ -731,10 +787,14 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
|
||||
}, nil
|
||||
}
|
||||
|
||||
client, ok := LookupStdioClient(invocation.CanonicalProduct)
|
||||
lookupKey := strings.Trim(strings.TrimPrefix(strings.TrimSpace(endpoint), stdioEndpointScheme), "/")
|
||||
if lookupKey == "" {
|
||||
lookupKey = invocation.CanonicalProduct
|
||||
}
|
||||
client, ok := LookupStdioClient(lookupKey)
|
||||
if !ok {
|
||||
return executor.Result{}, apperrors.NewInternal(
|
||||
fmt.Sprintf("stdio client not found for %q", invocation.CanonicalProduct))
|
||||
fmt.Sprintf("stdio client not found for %q", lookupKey))
|
||||
}
|
||||
|
||||
callCtx := ctx
|
||||
@@ -751,6 +811,27 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
|
||||
)
|
||||
}
|
||||
|
||||
tools, err := runnerStdioListTools(client, callCtx)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
fmt.Sprintf("stdio tools/list failed: %v", err),
|
||||
apperrors.WithOperation("tools/list"),
|
||||
apperrors.WithReason("stdio_tools_list_error"),
|
||||
)
|
||||
}
|
||||
schema, ok := pluginToolInputSchema(tools, invocation.Tool)
|
||||
if !ok {
|
||||
return executor.Result{}, apperrors.NewValidation(
|
||||
fmt.Sprintf("plugin tool %q is not declared by tools/list", invocation.Tool),
|
||||
apperrors.WithReason("plugin_tool_not_found"),
|
||||
)
|
||||
}
|
||||
normalizedParams, err := normalizePluginInputParams(invocation.Params, schema)
|
||||
if err != nil {
|
||||
return executor.Result{}, err
|
||||
}
|
||||
invocation.Params = normalizedParams
|
||||
|
||||
callResult, err := runnerStdioCallTool(client, callCtx, invocation.Tool, invocation.Params)
|
||||
if err != nil {
|
||||
return executor.Result{}, apperrors.NewAPI(
|
||||
@@ -779,67 +860,49 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) (string, error) {
|
||||
explicitToken := ""
|
||||
if r != nil && r.globalFlags != nil {
|
||||
explicitToken = r.globalFlags.Token
|
||||
}
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
}
|
||||
if tp := edition.Get().TokenProvider; tp != nil {
|
||||
token, _ := tp(ctx, func() (string, error) {
|
||||
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
|
||||
})
|
||||
return token
|
||||
}
|
||||
return getCachedRuntimeToken(ctx)
|
||||
return resolveRuntimeAuthToken(ctx, explicitToken)
|
||||
}
|
||||
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) (string, error) {
|
||||
snapshot, err := runtimeTokenManager.Get(ctx, defaultConfigDir(), explicitToken)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// Use cached token to avoid repeated Keychain access (~70ms per call)
|
||||
return getCachedRuntimeToken(ctx)
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
// Cached token state for process lifetime
|
||||
var (
|
||||
cachedRuntimeTokenMu sync.Mutex
|
||||
cachedRuntimeTokens = map[string]string{}
|
||||
)
|
||||
|
||||
// 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 {
|
||||
cacheKey := strings.TrimSpace(authpkg.RuntimeProfile())
|
||||
if cacheKey == "" {
|
||||
cacheKey = "__default__"
|
||||
}
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
if token := cachedRuntimeTokens[cacheKey]; token != "" {
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
return token
|
||||
}
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
|
||||
// getCachedRuntimeToken is kept as the prefetch seam used by runner tests. The
|
||||
// cache itself lives exclusively in TokenManager.
|
||||
func getCachedRuntimeToken(ctx context.Context) (string, error) {
|
||||
loadStart := time.Now()
|
||||
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
|
||||
return resolveRuntimeAuthToken(ctx, "")
|
||||
}
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return ""
|
||||
func tokenResolutionError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if token == "" {
|
||||
return ""
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return err
|
||||
}
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
cachedRuntimeTokens[cacheKey] = token
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
return token
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return apperrors.NewAuth(
|
||||
"未登录,请先执行 dws auth login",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("运行 'dws auth login' 完成登录后重试"),
|
||||
apperrors.WithActions("dws auth login"),
|
||||
apperrors.WithCause(err),
|
||||
)
|
||||
}
|
||||
// Keychain, parse, permission, lock, and refresh failures are real local or
|
||||
// network errors. Preserve their cause instead of disguising them as logout.
|
||||
return fmt.Errorf("resolve access token: %w", err)
|
||||
}
|
||||
|
||||
// generateExecutionID returns a random 16-char hex string used to correlate
|
||||
@@ -854,9 +917,7 @@ func generateExecutionID() string {
|
||||
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
|
||||
// This should be called after login/logout operations.
|
||||
func ResetRuntimeTokenCache() {
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
defer cachedRuntimeTokenMu.Unlock()
|
||||
cachedRuntimeTokens = map[string]string{}
|
||||
runtimeTokenManager.Invalidate()
|
||||
}
|
||||
|
||||
func newRuntimeContentScanner() safety.Scanner {
|
||||
|
||||
@@ -50,9 +50,9 @@ func TestCrossPlatformCoverageRunnerRemainingRoutingCoverage(t *testing.T) {
|
||||
|
||||
inv := executor.Invocation{CanonicalProduct: "product", Tool: "tool"}
|
||||
prefetched := make(chan struct{}, 1)
|
||||
runnerGetCachedRuntimeToken = func(context.Context) string {
|
||||
runnerGetCachedRuntimeToken = func(context.Context) (string, error) {
|
||||
prefetched <- struct{}{}
|
||||
return ""
|
||||
return "", nil
|
||||
}
|
||||
r := &runtimeRunner{
|
||||
loader: cli.CatalogLoaderFrom(cli.Catalog{}, wantErr),
|
||||
@@ -197,7 +197,7 @@ func TestCrossPlatformCoverageRunnerRemainingExecutionCoverage(t *testing.T) {
|
||||
return nil
|
||||
}
|
||||
|
||||
authErr := apperrors.NewAuth("expired")
|
||||
authErr := apperrors.NewAuth("expired", apperrors.WithReason("http_401"))
|
||||
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
return transport.ToolCallResult{}, authErr
|
||||
}
|
||||
@@ -290,10 +290,12 @@ func TestCrossPlatformCoverageRunnerRemainingExecutionCoverage(t *testing.T) {
|
||||
|
||||
func TestCrossPlatformCoverageRunnerRemainingStdioAuthAndHeadersCoverage(t *testing.T) {
|
||||
oldStdioInit := runnerStdioEnsureInitialized
|
||||
oldStdioList := runnerStdioListTools
|
||||
oldStdioCall := runnerStdioCallTool
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() {
|
||||
runnerStdioEnsureInitialized = oldStdioInit
|
||||
runnerStdioListTools = oldStdioList
|
||||
runnerStdioCallTool = oldStdioCall
|
||||
edition.Override(oldEdition)
|
||||
StopAllStdioClients()
|
||||
@@ -309,6 +311,14 @@ func TestCrossPlatformCoverageRunnerRemainingStdioAuthAndHeadersCoverage(t *test
|
||||
t.Fatalf("stdio initialize error = %v", err)
|
||||
}
|
||||
runnerStdioEnsureInitialized = func(*transport.StdioClient, context.Context) error { return nil }
|
||||
runnerStdioListTools = func(*transport.StdioClient, context.Context) (transport.ToolsListResult, error) {
|
||||
return transport.ToolsListResult{
|
||||
Tools: []transport.ToolDescriptor{{
|
||||
Name: "tool",
|
||||
InputSchema: map[string]any{"type": "object"},
|
||||
}},
|
||||
}, nil
|
||||
}
|
||||
runnerStdioCallTool = func(*transport.StdioClient, context.Context, string, map[string]any) (transport.ToolCallResult, error) {
|
||||
return transport.ToolCallResult{}, wantErr
|
||||
}
|
||||
@@ -327,21 +337,31 @@ func TestCrossPlatformCoverageRunnerRemainingStdioAuthAndHeadersCoverage(t *test
|
||||
if got, err := r.executeStdioInvocation(context.Background(), inv); err != nil || !got.Invocation.Implemented {
|
||||
t.Fatalf("stdio success = %#v, %v", got, err)
|
||||
}
|
||||
RegisterStdioClient("plugin/server-key", client)
|
||||
overlayIDInvocation := inv
|
||||
overlayIDInvocation.CanonicalProduct = "overlay-id"
|
||||
if got, err := r.executeInvocation(
|
||||
context.Background(),
|
||||
"stdio://plugin/server-key",
|
||||
overlayIDInvocation,
|
||||
); err != nil || !got.Invocation.Implemented {
|
||||
t.Fatalf("stdio endpoint-key lookup = %#v, %v", got, err)
|
||||
}
|
||||
|
||||
r.globalFlags.Token = " explicit "
|
||||
if got := r.resolveAuthToken(context.Background()); got != "explicit" {
|
||||
t.Fatalf("explicit auth token = %q", got)
|
||||
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "explicit" {
|
||||
t.Fatalf("explicit auth token = %q, %v", got, err)
|
||||
}
|
||||
edition.Override(&edition.Hooks{TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
|
||||
_, _ = fallback()
|
||||
return "provided", nil
|
||||
}})
|
||||
r.globalFlags.Token = ""
|
||||
if got := r.resolveAuthToken(context.Background()); got != "provided" {
|
||||
t.Fatalf("provided auth token = %q", got)
|
||||
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "provided" {
|
||||
t.Fatalf("provided auth token = %q, %v", got, err)
|
||||
}
|
||||
if got := resolveRuntimeAuthToken(context.Background(), " runtime "); got != "runtime" {
|
||||
t.Fatalf("runtime explicit token = %q", got)
|
||||
if got, err := resolveRuntimeAuthToken(context.Background(), " runtime "); err != nil || got != "runtime" {
|
||||
t.Fatalf("runtime explicit token = %q, %v", got, err)
|
||||
}
|
||||
|
||||
t.Setenv(envDWSChannel, "channel")
|
||||
|
||||
@@ -35,6 +35,7 @@ import (
|
||||
var (
|
||||
manualAgentExamplePlaceholderPattern = regexp.MustCompile(`<([^>]+)>`)
|
||||
manualAgentExampleDryRunJSONPattern = regexp.MustCompile(`(?i)"dry_run"\s*:\s*true`)
|
||||
manualAgentExampleDryRunPlanPattern = regexp.MustCompile(`(?i)"preview_kind"\s*:\s*"plan"`)
|
||||
)
|
||||
|
||||
// TestManualAgentExamplesContract is the always-on gate. It validates every
|
||||
@@ -391,6 +392,9 @@ func manualAgentExampleDryRunObserved(capture manualAgentExampleCapture) bool {
|
||||
|
||||
func manualAgentExampleDryRunEvidence(capture manualAgentExampleCapture) (string, bool) {
|
||||
normalized := strings.ToLower(capture.Output)
|
||||
if manualAgentExampleDryRunJSONPattern.MatchString(capture.Output) && manualAgentExampleDryRunPlanPattern.MatchString(capture.Output) {
|
||||
return cli.DryRunPreviewPlan, true
|
||||
}
|
||||
if manualAgentExampleDryRunJSONPattern.MatchString(capture.Output) {
|
||||
return cli.DryRunPreviewRequest, true
|
||||
}
|
||||
@@ -428,6 +432,22 @@ func TestManualAgentExampleDryRunEvidenceAcceptsSharedAndCommandPlans(t *testing
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageManualAgentExampleDryRunEvidenceClassifiesStructuredPlan(t *testing.T) {
|
||||
kind, observed := manualAgentExampleDryRunEvidence(manualAgentExampleCapture{
|
||||
Output: `{"dry_run":true,"executed":false,"preview_kind":"plan"}`,
|
||||
DryRunChecks: 1,
|
||||
})
|
||||
if !observed || kind != cli.DryRunPreviewPlan {
|
||||
t.Fatalf("structured plan classified as kind=%q observed=%v", kind, observed)
|
||||
}
|
||||
if manualAgentExampleDryRunObserved(manualAgentExampleCapture{
|
||||
Output: `{"dry_run":false,"executed":false,"preview_kind":"plan"}`,
|
||||
DryRunChecks: 1,
|
||||
}) {
|
||||
t.Fatal("non-dry-run structured plan was accepted as evidence")
|
||||
}
|
||||
}
|
||||
|
||||
func manualAgentExamplePromptObserved(output string) bool {
|
||||
normalized := strings.ToLower(output)
|
||||
for _, marker := range []string{
|
||||
|
||||
@@ -14,7 +14,7 @@ func TestRuntimeSchemaCompletenessCoversPublicCommandTree(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root := NewRootCommand()
|
||||
root := NewSchemaSourceRootCommand()
|
||||
if err := cli.ValidateEmbeddedRuntimeSchemaCompleteness(root); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -61,15 +61,24 @@ func TestEventRegistryDeliversOneTypedSchemaPath(t *testing.T) {
|
||||
}
|
||||
}
|
||||
for flag, wantType := range map[string]string{
|
||||
"dry-run": "boolean",
|
||||
"duration": "string",
|
||||
"event-types": "array",
|
||||
"max-events": "integer",
|
||||
"dry-run": "boolean",
|
||||
"duration": "string",
|
||||
"event-types": "array",
|
||||
"max-events": "integer",
|
||||
"open-dingtalk-id": "string",
|
||||
} {
|
||||
if got := schemaContractString(consumeParams[flag]["type"]); got != wantType {
|
||||
t.Errorf("event.consume --%s type = %q, want %q", flag, got, wantType)
|
||||
}
|
||||
}
|
||||
if _, exists := consumeParams["odid"]; exists {
|
||||
t.Error("event.consume exposes unsupported --odid alias")
|
||||
}
|
||||
for _, name := range []string{"user", "open-dingtalk-id", "group"} {
|
||||
if _, exists := consumeParams[name]["required_when"]; exists {
|
||||
t.Errorf("event.consume --%s unexpectedly declares required_when", name)
|
||||
}
|
||||
}
|
||||
if _, exists := consumeParams["duration"]["default"]; exists {
|
||||
t.Error("event.consume --duration leaked zero default 0s")
|
||||
}
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"archive/zip"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
@@ -36,25 +37,25 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
skillLoadAccessToken = loadSkillAccessToken
|
||||
skillDownloadToTmp = downloadSkillToTmpDir
|
||||
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
|
||||
skillNewRequest = http.NewRequestWithContext
|
||||
skillLoadTokenData = authpkg.LoadTokenData
|
||||
skillResolveTargetPath = resolveSkillTargetPath
|
||||
skillFetchDownloadInfo = fetchSkillDownloadInfo
|
||||
skillDownloadFile = downloadSkillFile
|
||||
skillExtractZip = extractSkillZip
|
||||
skillUserHomeDir = os.UserHomeDir
|
||||
skillMkdirTemp = os.MkdirTemp
|
||||
skillCreate = os.Create
|
||||
skillCreateTemp = os.CreateTemp
|
||||
skillRemoveAll = os.RemoveAll
|
||||
skillRemove = os.Remove
|
||||
skillMkdirAll = os.MkdirAll
|
||||
skillOpenFile = os.OpenFile
|
||||
skillCopy = io.Copy
|
||||
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
|
||||
skillLoadAccessToken = loadSkillAccessToken
|
||||
skillDownloadToTmp = downloadSkillToTmpDir
|
||||
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
|
||||
skillNewRequest = http.NewRequestWithContext
|
||||
skillResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
skillResolveTargetPath = resolveSkillTargetPath
|
||||
skillFetchDownloadInfo = fetchSkillDownloadInfo
|
||||
skillDownloadFile = downloadSkillFile
|
||||
skillExtractZip = extractSkillZip
|
||||
skillUserHomeDir = os.UserHomeDir
|
||||
skillMkdirTemp = os.MkdirTemp
|
||||
skillCreate = os.Create
|
||||
skillCreateTemp = os.CreateTemp
|
||||
skillRemoveAll = os.RemoveAll
|
||||
skillRemove = os.Remove
|
||||
skillMkdirAll = os.MkdirAll
|
||||
skillOpenFile = os.OpenFile
|
||||
skillCopy = io.Copy
|
||||
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -296,7 +297,7 @@ func newSkillAddHintCommand() *cobra.Command {
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -319,7 +320,7 @@ func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
if source == "" {
|
||||
source, _ = cmd.Flags().GetString("scopes")
|
||||
}
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -388,7 +389,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
|
||||
}
|
||||
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -441,13 +442,16 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSkillAccessToken() (string, error) {
|
||||
func loadSkillAccessToken(ctx context.Context) (string, error) {
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := skillLoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
token, err := skillResolveAccessToken(ctx, configDir, "")
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return "", skillAuthError()
|
||||
}
|
||||
return tokenData.AccessToken, nil
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolve skill access token: %w", err)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func skillAuthError() error {
|
||||
|
||||
@@ -57,11 +57,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
})
|
||||
fail := errors.New("failure")
|
||||
cmd := skillCoverageCommand()
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillGet(cmd, nil); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill get auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
|
||||
if err := runSkillFind(cmd, nil); err == nil {
|
||||
t.Fatal("skill find request failure should propagate")
|
||||
@@ -72,11 +72,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
t.Fatalf("skill get download error = %v", err)
|
||||
}
|
||||
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillFind(cmd, nil); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill find auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillHTTPDo = func(*http.Client, *http.Request) (*http.Response, error) { return nil, fail }
|
||||
if err := runSkillFind(cmd, nil); err == nil {
|
||||
t.Fatal("skill find network failure should propagate")
|
||||
@@ -111,11 +111,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
t.Fatal("invalid skill target should fail")
|
||||
}
|
||||
skillResolveTargetPath = func(string) (string, error) { return "dest", nil }
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill add auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillFetchDownloadInfo = func(context.Context, string, string) (*downloadSkillResponse, error) { return nil, fail }
|
||||
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill info error = %v", err)
|
||||
@@ -152,25 +152,35 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
|
||||
func TestCrossPlatformCoverageSkillCommandLowLevelRemainingCoverage(t *testing.T) {
|
||||
oldHTTP := skillHTTPDo
|
||||
oldNewRequest, oldLoadToken := skillNewRequest, skillLoadTokenData
|
||||
oldNewRequest, oldResolveToken := skillNewRequest, skillResolveAccessToken
|
||||
oldHome := skillUserHomeDir
|
||||
oldMkdirTemp, oldCreate, oldCreateTemp := skillMkdirTemp, skillCreate, skillCreateTemp
|
||||
oldRemoveAll, oldRemove, oldMkdir := skillRemoveAll, skillRemove, skillMkdirAll
|
||||
oldOpen, oldCopy, oldZipOpen := skillOpenFile, skillCopy, skillOpenZipFile
|
||||
t.Cleanup(func() {
|
||||
skillHTTPDo = oldHTTP
|
||||
skillNewRequest, skillLoadTokenData = oldNewRequest, oldLoadToken
|
||||
skillNewRequest, skillResolveAccessToken = oldNewRequest, oldResolveToken
|
||||
skillUserHomeDir = oldHome
|
||||
skillMkdirTemp, skillCreate, skillCreateTemp = oldMkdirTemp, oldCreate, oldCreateTemp
|
||||
skillRemoveAll, skillRemove, skillMkdirAll = oldRemoveAll, oldRemove, oldMkdir
|
||||
skillOpenFile, skillCopy, skillOpenZipFile = oldOpen, oldCopy, oldZipOpen
|
||||
})
|
||||
fail := errors.New("failure")
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
|
||||
if _, err := loadSkillAccessToken(); err == nil {
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
if _, err := loadSkillAccessToken(context.Background()); err == nil {
|
||||
t.Fatal("invalid skill access token succeeded")
|
||||
}
|
||||
skillLoadTokenData = oldLoadToken
|
||||
canceled, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
skillResolveAccessToken = func(ctx context.Context, _, _ string) (string, error) {
|
||||
return "", ctx.Err()
|
||||
}
|
||||
if _, err := loadSkillAccessToken(canceled); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("skill token cancellation = %v", err)
|
||||
}
|
||||
skillResolveAccessToken = oldResolveToken
|
||||
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
|
||||
if _, err := fetchSkillDownloadInfo(context.Background(), "token", "id"); err == nil {
|
||||
t.Fatal("download-info request failure should propagate")
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -397,9 +396,11 @@ func TestSkillInstallRequiresAuth(t *testing.T) {
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
t.Cleanup(CloseFileLogger)
|
||||
originalLoadToken := skillLoadTokenData
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, errors.New("missing") }
|
||||
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
|
||||
originalResolveToken := skillResolveAccessToken
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
|
||||
|
||||
// Ensure the config directory exists but has no token
|
||||
if err := os.MkdirAll(configDir, 0755); err != nil {
|
||||
@@ -679,16 +680,11 @@ func TestSkillSearchUsesSourceQueryAndKeepsScopesCompat(t *testing.T) {
|
||||
configDir := filepath.Join(t.TempDir(), "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
t.Cleanup(CloseFileLogger)
|
||||
originalLoadToken := skillLoadTokenData
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{
|
||||
AccessToken: "test-token",
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
}, nil
|
||||
originalResolveToken := skillResolveAccessToken
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "test-token", nil
|
||||
}
|
||||
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
|
||||
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
|
||||
|
||||
var gotSources []string
|
||||
var gotScopes []string
|
||||
|
||||
@@ -52,6 +52,10 @@ func TestMain(m *testing.M) {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
panic("set " + keychain.StorageDirEnv + ": " + err.Error())
|
||||
}
|
||||
if err := os.Setenv(keychain.TestNamespaceEnv, tmpDir); err != nil {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
panic("set " + keychain.TestNamespaceEnv + ": " + err.Error())
|
||||
}
|
||||
if err := os.Setenv("DWS_CONFIG_DIR", filepath.Join(tmpDir, "config")); err != nil {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
panic("set DWS_CONFIG_DIR: " + err.Error())
|
||||
@@ -79,6 +83,12 @@ func TestMain(m *testing.M) {
|
||||
StopAllStdioClients()
|
||||
CloseAuditSink()
|
||||
CloseFileLogger()
|
||||
if err := keychain.RemoveAuthTokenEntries(keychain.Service); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "internal/app keychain cleanup: %v\n", err)
|
||||
if code == 0 {
|
||||
code = 1
|
||||
}
|
||||
}
|
||||
if err := os.RemoveAll(tmpDir); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "internal/app test cleanup %s: %v\n", tmpDir, err)
|
||||
if code == 0 {
|
||||
|
||||
@@ -17,7 +17,9 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/jsonutil"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
@@ -27,8 +29,9 @@ import (
|
||||
// interface so that private overlays can invoke MCP tools without importing
|
||||
// internal packages.
|
||||
type toolCallerAdapter struct {
|
||||
runner executor.Runner
|
||||
flags *GlobalFlags
|
||||
runner executor.Runner
|
||||
flags *GlobalFlags
|
||||
tokenMu sync.Mutex
|
||||
}
|
||||
|
||||
var toolCallerDryRun = func(ctx context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
@@ -63,6 +66,26 @@ func (a *toolCallerAdapter) CallTool(ctx context.Context, productID, toolName st
|
||||
return convertResult(result), nil
|
||||
}
|
||||
|
||||
// CallToolWithToken invokes a helper with an in-memory token override. It is
|
||||
// used during login before the new token has been persisted to any profile
|
||||
// slot.
|
||||
func (a *toolCallerAdapter) CallToolWithToken(ctx context.Context, token, productID, toolName string, args map[string]any) (*edition.ToolResult, error) {
|
||||
if a == nil || a.flags == nil {
|
||||
return nil, fmt.Errorf("ToolCaller token override is not configured")
|
||||
}
|
||||
a.tokenMu.Lock()
|
||||
defer a.tokenMu.Unlock()
|
||||
previousToken := a.flags.Token
|
||||
previousProfile := authpkg.RuntimeProfile()
|
||||
a.flags.Token = token
|
||||
authpkg.SetRuntimeProfile("")
|
||||
defer func() {
|
||||
a.flags.Token = previousToken
|
||||
authpkg.SetRuntimeProfile(previousProfile)
|
||||
}()
|
||||
return a.CallTool(ctx, productID, toolName, args)
|
||||
}
|
||||
|
||||
func (a *toolCallerAdapter) Format() string {
|
||||
if a != nil && a.flags != nil {
|
||||
return a.flags.Format
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
// 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"
|
||||
"testing"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
|
||||
)
|
||||
|
||||
func TestToolCallerTokenOverrideClearsUnpersistedRuntimeProfile(t *testing.T) {
|
||||
authpkg.SetRuntimeProfile("corp_not_persisted")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
|
||||
flags := &GlobalFlags{}
|
||||
runner := runtimeProfileCaptureRunner{flags: flags}
|
||||
caller := &toolCallerAdapter{runner: runner, flags: flags}
|
||||
result, err := caller.CallToolWithToken(
|
||||
context.Background(),
|
||||
"temporary-access-token",
|
||||
"contact",
|
||||
"get_current_user_profile",
|
||||
nil,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("CallToolWithToken() error = %v", err)
|
||||
}
|
||||
if got := result.Content[0].Text; got != `{"profile":"","token":"temporary-access-token"}` {
|
||||
t.Fatalf("CallToolWithToken() result = %s", got)
|
||||
}
|
||||
if authpkg.RuntimeProfile() != "corp_not_persisted" {
|
||||
t.Fatalf("runtime profile = %q, want restored selector", authpkg.RuntimeProfile())
|
||||
}
|
||||
if flags.Token != "" {
|
||||
t.Fatalf("token override leaked after call: %q", flags.Token)
|
||||
}
|
||||
}
|
||||
|
||||
type runtimeProfileCaptureRunner struct {
|
||||
flags *GlobalFlags
|
||||
}
|
||||
|
||||
func (r runtimeProfileCaptureRunner) Run(_ context.Context, invocation executor.Invocation) (executor.Result, error) {
|
||||
return executor.Result{
|
||||
Invocation: invocation,
|
||||
Response: map[string]any{
|
||||
"content": []any{map[string]any{
|
||||
"type": "text",
|
||||
"text": `{"profile":"` + authpkg.RuntimeProfile() + `","token":"` + r.flags.Token + `"}`,
|
||||
}},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -15,6 +15,7 @@ package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/shortcut/usage"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
@@ -38,6 +39,21 @@ func (r recordingToolCaller) CallTool(ctx context.Context, product, tool string,
|
||||
return res, err
|
||||
}
|
||||
|
||||
func (r recordingToolCaller) CallToolWithToken(
|
||||
ctx context.Context,
|
||||
token, product, tool string,
|
||||
args map[string]any,
|
||||
) (*edition.ToolResult, error) {
|
||||
inner, ok := r.inner.(tokenOverrideToolCaller)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("ToolCaller token override is not configured")
|
||||
}
|
||||
recordedArgs := cloneToolArgs(args)
|
||||
res, err := inner.CallToolWithToken(ctx, token, product, tool, args)
|
||||
usage.Append(product, tool, recordedArgs, err == nil, r.inner.DryRun())
|
||||
return res, err
|
||||
}
|
||||
|
||||
func (r recordingToolCaller) Format() string { return r.inner.Format() }
|
||||
func (r recordingToolCaller) DryRun() bool { return r.inner.DryRun() }
|
||||
func (r recordingToolCaller) Fields() string { return r.inner.Fields() }
|
||||
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
|
||||
type crossPlatformCoverageCaller struct {
|
||||
args map[string]any
|
||||
token string
|
||||
dryRun bool
|
||||
}
|
||||
|
||||
@@ -33,6 +34,16 @@ func (c *crossPlatformCoverageCaller) CallTool(_ context.Context, _, _ string, a
|
||||
return &edition.ToolResult{}, nil
|
||||
}
|
||||
|
||||
func (c *crossPlatformCoverageCaller) CallToolWithToken(
|
||||
_ context.Context,
|
||||
token, _, _ string,
|
||||
args map[string]any,
|
||||
) (*edition.ToolResult, error) {
|
||||
c.token = token
|
||||
c.args = args
|
||||
return &edition.ToolResult{}, nil
|
||||
}
|
||||
|
||||
func (*crossPlatformCoverageCaller) Format() string { return "json" }
|
||||
func (c *crossPlatformCoverageCaller) DryRun() bool { return c.dryRun }
|
||||
func (*crossPlatformCoverageCaller) Fields() string { return "id,name" }
|
||||
@@ -88,6 +99,23 @@ func TestCrossPlatformCoverageRecordingToolCaller(t *testing.T) {
|
||||
t.Fatal("sensitive text must not be recorded")
|
||||
}
|
||||
|
||||
tokenCaller, ok := caller.(tokenOverrideToolCaller)
|
||||
if !ok {
|
||||
t.Fatal("recording caller dropped token override support")
|
||||
}
|
||||
if _, err := tokenCaller.CallToolWithToken(
|
||||
context.Background(),
|
||||
"temporary-token",
|
||||
"contact",
|
||||
"get_current_user_profile",
|
||||
nil,
|
||||
); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if inner.token != "temporary-token" || len(inner.args) != 0 {
|
||||
t.Fatalf("token override forwarding = token %q args %#v", inner.token, inner.args)
|
||||
}
|
||||
|
||||
inner.dryRun = true
|
||||
if _, err := caller.CallTool(context.Background(), "chat", "send_message", args); err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -96,7 +124,7 @@ func TestCrossPlatformCoverageRecordingToolCaller(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(records) != 1 {
|
||||
if len(records) != 2 {
|
||||
t.Fatalf("dry-run call must not be recorded: %#v", records)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotPreservesLoadFailure(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
want := errors.New("keychain permission denied")
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
_, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
if errors.Is(err, ErrTokenDataNotFound) {
|
||||
t.Fatalf("load failure was misclassified as missing credentials: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderLoginPreservesLoadFailure(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
want := errors.New("keychain permission denied")
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
_, err := NewOAuthProvider(t.TempDir(), nil).Login(context.Background(), false)
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotReturnsExpiryMetadata(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
oauthLoadToken = func(string) (*TokenData, error) {
|
||||
return &TokenData{AccessToken: "token", ExpiresAt: expiresAt}, nil
|
||||
}
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
snapshot, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if snapshot.AccessToken != "token" || !snapshot.ExpiresAt.Equal(expiresAt) {
|
||||
t.Fatalf("snapshot = %#v", snapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenMarkerRevisionChangesOnEveryPublication(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, present, err := ReadTokenMarkerRevision(configDir)
|
||||
if err != nil || !present || first == "" {
|
||||
t.Fatalf("first marker = %q, %v, %v", first, present, err)
|
||||
}
|
||||
if err := WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, present, err := ReadTokenMarkerRevision(configDir)
|
||||
if err != nil || !present || second == "" || second == first {
|
||||
t.Fatalf("second marker = %q, %v, %v; first=%q", second, present, err, first)
|
||||
}
|
||||
}
|
||||
@@ -306,6 +306,43 @@ func TestRevokeTokenRemote(t *testing.T) {
|
||||
// Can't easily test since LogoutURL is a const. Just test that it doesn't panic with real URL.
|
||||
}
|
||||
|
||||
func TestRevokeTokenRemoteForDataUsesExactTokenMetadata(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
var got struct {
|
||||
ClientID string `json:"clientId"`
|
||||
AccessToken string `json:"accessToken"`
|
||||
}
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != MCPRevokeTokenPath {
|
||||
t.Errorf("revoke path = %q, want %q", r.URL.Path, MCPRevokeTokenPath)
|
||||
}
|
||||
if err := json.NewDecoder(r.Body).Decode(&got); err != nil {
|
||||
t.Errorf("decode revoke body: %v", err)
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
if err := os.WriteFile(filepath.Join(configDir, "mcp_url"), []byte(srv.URL), 0o600); err != nil {
|
||||
t.Fatalf("write mcp_url: %v", err)
|
||||
}
|
||||
SetClientID("wrong-global-client")
|
||||
t.Cleanup(func() { SetClientID("") })
|
||||
|
||||
data := &TokenData{
|
||||
AccessToken: "exact-account-token",
|
||||
ClientID: "exact-account-client",
|
||||
Source: "mcp",
|
||||
}
|
||||
if err := RevokeTokenRemoteForData(t.Context(), data); err != nil {
|
||||
t.Fatalf("RevokeTokenRemoteForData() error = %v", err)
|
||||
}
|
||||
if got.ClientID != data.ClientID || got.AccessToken != data.AccessToken {
|
||||
t.Fatalf("revoke body = %#v, want exact token metadata", got)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── oauth_helpers.go ──────────────────────────────────────────────────
|
||||
|
||||
type tokenResponse struct {
|
||||
|
||||
@@ -495,6 +495,10 @@ func TestDeviceFlow_LoginOnce_CLIAuthEnabled_Success(t *testing.T) {
|
||||
logger: newDeviceFlowTestLogger(),
|
||||
Output: io.Discard,
|
||||
httpClient: srv.Client(),
|
||||
IdentityEnricher: func(_ context.Context, data *TokenData) error {
|
||||
data.UserID = "user123"
|
||||
return nil
|
||||
},
|
||||
}
|
||||
|
||||
token, err := provider.loginOnce(context.Background(), 1)
|
||||
|
||||
@@ -396,7 +396,8 @@ func TestCrossPlatformCoverageDeviceFlowLoginRetry(t *testing.T) {
|
||||
return
|
||||
}
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"accessToken": "device-access", "refreshToken": "refresh", "expiresIn": 7200, "corpId": "corp-device",
|
||||
"accessToken": "device-access", "refreshToken": "refresh", "expiresIn": 7200,
|
||||
"corpId": "corp-device", "userId": "user-device",
|
||||
})
|
||||
case strings.HasSuffix(r.URL.Path, CLIAuthEnabledPath):
|
||||
_ = json.NewEncoder(w).Encode(CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: true}})
|
||||
@@ -493,7 +494,7 @@ func TestCrossPlatformCoverageProfilesLifecycleEdges(t *testing.T) {
|
||||
}
|
||||
normalizeProfilesConfig(cfg)
|
||||
if len(cfg.Profiles) != 1 || cfg.Profiles[0].Name != "Acme" ||
|
||||
cfg.PrimaryProfile != "corp-a" || cfg.CurrentProfile != "corp-a" || cfg.PreviousProfile != "" {
|
||||
cfg.PrimaryProfile != "" || cfg.CurrentProfile != "" || cfg.PreviousProfile != "" {
|
||||
t.Fatalf("normalized profiles = %#v", cfg)
|
||||
}
|
||||
|
||||
@@ -515,6 +516,16 @@ func TestCrossPlatformCoverageProfilesLifecycleEdges(t *testing.T) {
|
||||
if err := UpsertProfileFromTokenWithCurrent(dir, second, true); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := SaveTokenDataKeychainForIdentity(first.CorpID, first.UserID, first); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := SaveTokenDataKeychainForCorpID(second.CorpID, second); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
_ = DeleteTokenDataKeychainForIdentity(first.CorpID, first.UserID)
|
||||
_ = DeleteTokenDataKeychainForCorpID(second.CorpID)
|
||||
})
|
||||
profilesForAmbiguity, err := LoadProfiles(dir)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -575,7 +586,7 @@ func TestCrossPlatformCoverageProfilesLifecycleEdges(t *testing.T) {
|
||||
{CorpID: "two", Name: "Acme-12345678"},
|
||||
{CorpID: "three", Name: "Acme-12345678-2"},
|
||||
}}
|
||||
if got := chooseProfileName(nameCfg, &TokenData{CorpID: "corp-12345678", CorpName: "Acme"}); got != "Acme-12345678-3" {
|
||||
if got := chooseProfileName(nameCfg, &TokenData{CorpID: "corp-12345678", CorpName: "Acme"}); got != "Acme-2" {
|
||||
t.Fatalf("collision profile name = %q", got)
|
||||
}
|
||||
if chooseProfileName(&ProfilesConfig{}, &TokenData{}) != "profile" {
|
||||
@@ -585,12 +596,7 @@ func TestCrossPlatformCoverageProfilesLifecycleEdges(t *testing.T) {
|
||||
!shouldRefreshProfileName(&Profile{}, first) {
|
||||
t.Fatal("profile refresh-name decisions failed")
|
||||
}
|
||||
if findProfile(nil, "x") != nil || findProfile(cfg, "") != nil ||
|
||||
profileIndexByCorpID(nil, "x") != -1 || firstProfileCorpID(nil) != "" {
|
||||
t.Fatal("nil profile helper failed")
|
||||
}
|
||||
if shortCorpID("short") != "short" || shortCorpID("corp-12345678") != "12345678" ||
|
||||
timeOrRFC3339(time.Time{}) != "" {
|
||||
if shortCorpID("short") != "short" || shortCorpID("corp-12345678") != "12345678" {
|
||||
t.Fatal("profile formatting helpers failed")
|
||||
}
|
||||
}
|
||||
@@ -1449,31 +1455,43 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
oldMarshalIndent := tokenJSONMarshalIndent
|
||||
oldMarshal := tokenJSONMarshal
|
||||
oldMkdir := tokenMkdirAll
|
||||
oldRead := tokenReadFile
|
||||
oldWrite := tokenWriteFile
|
||||
oldRename := tokenRename
|
||||
oldRemove := tokenRemove
|
||||
oldGlob := tokenGlob
|
||||
oldSaveCorp := tokenSaveKeychainForCorpID
|
||||
oldSaveIdentity := tokenSaveKeychainForIdentity
|
||||
oldSaveLegacy := tokenSaveKeychain
|
||||
oldLoadCorp := tokenLoadKeychainForCorpID
|
||||
oldLoadIdentity := tokenLoadKeychainIdentity
|
||||
oldLoadLegacy := tokenLoadKeychain
|
||||
oldExists := tokenKeychainExists
|
||||
oldDeleteCorp := tokenDeleteKeychainForCorpID
|
||||
oldDeleteIdentity := tokenDeleteKeychainIdentity
|
||||
oldDeleteLegacy := tokenDeleteKeychain
|
||||
oldRemoveAuthEntries := tokenRemoveAuthTokenEntries
|
||||
oldLoadSecure := tokenLoadSecure
|
||||
oldDeleteSecure := tokenDeleteSecure
|
||||
oldEnsureProfiles := profilesEnsureMigration
|
||||
oldResolve := tokenResolveProfile
|
||||
oldResolveDeletion := tokenResolveDeletion
|
||||
oldResolveSelection := tokenResolveSelection
|
||||
oldUpsert := tokenUpsertProfile
|
||||
oldRemoveProfile := tokenRemoveProfile
|
||||
oldSync := tokenSyncLegacyMirror
|
||||
oldSyncOrganization := tokenSyncOrganizationMirror
|
||||
oldLoadProfiles := tokenLoadProfiles
|
||||
oldSaveProfiles := tokenSaveProfiles
|
||||
oldWriteMarker := tokenWriteMarker
|
||||
oldWriteManualMarker := tokenWriteManualMarker
|
||||
oldDeleteMarker := tokenDeleteMarker
|
||||
oldParseURL := tokenParseURL
|
||||
oldNewRequest := tokenNewRequest
|
||||
oldDefaultDir := tokenDefaultConfigDir
|
||||
oldLoadData := tokenLoadData
|
||||
oldRevokeURL := tokenRevokeURL
|
||||
oldMCPBaseURL := tokenMCPBaseURL
|
||||
oldLogoutURL := tokenLogoutURL
|
||||
oldLogoutContinue := tokenLogoutContinueURL
|
||||
oldLogoutClient := tokenLogoutHTTPClient
|
||||
@@ -1483,31 +1501,43 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
tokenJSONMarshalIndent = oldMarshalIndent
|
||||
tokenJSONMarshal = oldMarshal
|
||||
tokenMkdirAll = oldMkdir
|
||||
tokenReadFile = oldRead
|
||||
tokenWriteFile = oldWrite
|
||||
tokenRename = oldRename
|
||||
tokenRemove = oldRemove
|
||||
tokenGlob = oldGlob
|
||||
tokenSaveKeychainForCorpID = oldSaveCorp
|
||||
tokenSaveKeychainForIdentity = oldSaveIdentity
|
||||
tokenSaveKeychain = oldSaveLegacy
|
||||
tokenLoadKeychainForCorpID = oldLoadCorp
|
||||
tokenLoadKeychainIdentity = oldLoadIdentity
|
||||
tokenLoadKeychain = oldLoadLegacy
|
||||
tokenKeychainExists = oldExists
|
||||
tokenDeleteKeychainForCorpID = oldDeleteCorp
|
||||
tokenDeleteKeychainIdentity = oldDeleteIdentity
|
||||
tokenDeleteKeychain = oldDeleteLegacy
|
||||
tokenRemoveAuthTokenEntries = oldRemoveAuthEntries
|
||||
tokenLoadSecure = oldLoadSecure
|
||||
tokenDeleteSecure = oldDeleteSecure
|
||||
profilesEnsureMigration = oldEnsureProfiles
|
||||
tokenResolveProfile = oldResolve
|
||||
tokenResolveDeletion = oldResolveDeletion
|
||||
tokenResolveSelection = oldResolveSelection
|
||||
tokenUpsertProfile = oldUpsert
|
||||
tokenRemoveProfile = oldRemoveProfile
|
||||
tokenSyncLegacyMirror = oldSync
|
||||
tokenSyncOrganizationMirror = oldSyncOrganization
|
||||
tokenLoadProfiles = oldLoadProfiles
|
||||
tokenSaveProfiles = oldSaveProfiles
|
||||
tokenWriteMarker = oldWriteMarker
|
||||
tokenWriteManualMarker = oldWriteManualMarker
|
||||
tokenDeleteMarker = oldDeleteMarker
|
||||
tokenParseURL = oldParseURL
|
||||
tokenNewRequest = oldNewRequest
|
||||
tokenDefaultConfigDir = oldDefaultDir
|
||||
tokenLoadData = oldLoadData
|
||||
tokenRevokeURL = oldRevokeURL
|
||||
tokenMCPBaseURL = oldMCPBaseURL
|
||||
tokenLogoutURL = oldLogoutURL
|
||||
tokenLogoutContinueURL = oldLogoutContinue
|
||||
tokenLogoutHTTPClient = oldLogoutClient
|
||||
@@ -1572,10 +1602,12 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Fatalf("corp marker error = %v", err)
|
||||
}
|
||||
SetRuntimeProfile("")
|
||||
tokenWriteManualMarker = func(string) error { return fail }
|
||||
if err := saveTokenDataLocked(dir, &TokenData{}); !errors.Is(err, fail) {
|
||||
t.Fatalf("legacy marker error = %v", err)
|
||||
}
|
||||
tokenWriteMarker = oldWriteMarker
|
||||
tokenWriteManualMarker = oldWriteManualMarker
|
||||
|
||||
edition.Override(&edition.Hooks{Name: "coverage", SaveToken: func(string, []byte) error { return nil }})
|
||||
if err := saveTokenDataLocked(dir, data); err != nil {
|
||||
@@ -1630,6 +1662,9 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Fatalf("failed migration load = %#v %v", got, err)
|
||||
}
|
||||
tokenSaveKeychain = func(*TokenData) error { return nil }
|
||||
tokenLoadKeychain = func() (*TokenData, error) { return nil, ErrTokenDataNotFound }
|
||||
tokenReadFile = func(string) ([]byte, error) { return nil, os.ErrNotExist }
|
||||
tokenWriteManualMarker = func(string) error { return nil }
|
||||
deleted := false
|
||||
tokenDeleteSecure = func(string) error { deleted = true; return nil }
|
||||
if _, err := LoadTokenDataForProfile(dir, ""); err != nil || !deleted {
|
||||
@@ -1653,25 +1688,70 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Run("delete branches", func(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
fail := errors.New("fail")
|
||||
selected := &Profile{CorpID: "corp"}
|
||||
tokenResolveProfile = func(string, string) (*Profile, error) { return nil, fail }
|
||||
selected := &Profile{CorpID: "corp", UserID: "user"}
|
||||
profilesEnsureMigration = func(string) error { return fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete migration error = %v", err)
|
||||
}
|
||||
profilesEnsureMigration = func(string) error { return nil }
|
||||
tokenLoadProfiles = func(string) (*ProfilesConfig, error) { return nil, fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete profiles load error = %v", err)
|
||||
}
|
||||
cfg := &ProfilesConfig{
|
||||
CurrentProfile: ProfileSelector(*selected),
|
||||
OrgCurrentProfiles: map[string]string{"corp": ProfileSelector(*selected)},
|
||||
Profiles: []Profile{*selected},
|
||||
}
|
||||
tokenLoadProfiles = func(string) (*ProfilesConfig, error) { return cfg, nil }
|
||||
tokenResolveDeletion = func(*ProfilesConfig, string) (*Profile, bool, error) { return nil, false, fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete resolution error = %v", err)
|
||||
}
|
||||
tokenResolveProfile = func(string, string) (*Profile, error) { return selected, nil }
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return fail }
|
||||
tokenResolveDeletion = func(*ProfilesConfig, string) (*Profile, bool, error) { return selected, true, nil }
|
||||
tokenLoadKeychainIdentity = func(string, string) (*TokenData, error) {
|
||||
return &TokenData{CorpID: "corp", UserID: "user"}, nil
|
||||
}
|
||||
tokenLoadKeychainForCorpID = func(string) (*TokenData, error) {
|
||||
return &TokenData{CorpID: "corp", UserID: "user"}, nil
|
||||
}
|
||||
tokenLoadKeychain = func() (*TokenData, error) {
|
||||
return &TokenData{CorpID: "corp", UserID: "user"}, nil
|
||||
}
|
||||
tokenReadFile = func(string) ([]byte, error) { return nil, os.ErrNotExist }
|
||||
tokenSaveProfiles = func(string, *ProfilesConfig) error { return nil }
|
||||
tokenSaveKeychainForIdentity = func(string, string, *TokenData) error { return nil }
|
||||
tokenSaveKeychainForCorpID = func(string, *TokenData) error { return nil }
|
||||
tokenSaveKeychain = func(*TokenData) error { return nil }
|
||||
tokenWriteMarker = func(string) error { return nil }
|
||||
tokenDeleteMarker = func(string) error { return nil }
|
||||
tokenRemoveProfile = func(string, string) (*Profile, error) { return selected, nil }
|
||||
tokenSyncLegacyMirror = func(string) error { return nil }
|
||||
tokenSyncOrganizationMirror = func(Profile) error { return nil }
|
||||
tokenDeleteSecure = func(string) error { return nil }
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return nil }
|
||||
tokenDeleteKeychainIdentity = func(string, string) error { return fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete identity keychain error = %v", err)
|
||||
}
|
||||
tokenDeleteKeychainIdentity = func(string, string) error { return nil }
|
||||
cfg.OrgCurrentProfiles = map[string]string{"corp": ProfileSelector(*selected)}
|
||||
tokenRemoveProfile = func(string, string) (*Profile, error) {
|
||||
cfg.OrgCurrentProfiles = nil
|
||||
return selected, nil
|
||||
}
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete corp keychain error = %v", err)
|
||||
}
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return nil }
|
||||
cfg.OrgCurrentProfiles = map[string]string{"corp": ProfileSelector(*selected)}
|
||||
tokenRemoveProfile = func(string, string) (*Profile, error) { return nil, fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("remove profile error = %v", err)
|
||||
}
|
||||
tokenRemoveProfile = func(string, string) (*Profile, error) { return selected, nil }
|
||||
cfg.OrgCurrentProfiles = nil
|
||||
tokenSyncLegacyMirror = func(string) error { return fail }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
t.Fatalf("sync mirror error = %v", err)
|
||||
@@ -1682,7 +1762,8 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Fatalf("secure cleanup error = %v", err)
|
||||
}
|
||||
|
||||
tokenResolveProfile = func(string, string) (*Profile, error) { return nil, nil }
|
||||
cfg.CurrentProfile = ""
|
||||
cfg.Profiles = nil
|
||||
tokenDeleteKeychain = func() error { return fail }
|
||||
tokenDeleteMarker = func(string) error { return nil }
|
||||
if err := deleteTokenDataForProfileLocked(dir, ""); !errors.Is(err, fail) {
|
||||
@@ -1710,22 +1791,17 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Run("delete all", func(t *testing.T) {
|
||||
fail := errors.New("fail")
|
||||
base := func() string {
|
||||
tokenLoadProfiles = func(string) (*ProfilesConfig, error) { return &ProfilesConfig{}, nil }
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return nil }
|
||||
tokenRemoveAuthTokenEntries = func(string) error { return nil }
|
||||
tokenRemove = func(string) error { return os.ErrNotExist }
|
||||
tokenGlob = func(string) ([]string, error) { return nil, nil }
|
||||
tokenDeleteKeychain = func() error { return nil }
|
||||
tokenDeleteSecure = func(string) error { return nil }
|
||||
tokenDeleteMarker = func(string) error { return nil }
|
||||
return t.TempDir()
|
||||
}
|
||||
dir := base()
|
||||
tokenLoadProfiles = func(string) (*ProfilesConfig, error) {
|
||||
return &ProfilesConfig{Profiles: []Profile{{CorpID: "corp"}}}, nil
|
||||
}
|
||||
tokenDeleteKeychainForCorpID = func(string) error { return fail }
|
||||
tokenRemoveAuthTokenEntries = func(string) error { return fail }
|
||||
if err := DeleteAllTokenData(dir); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete-all corp error = %v", err)
|
||||
t.Fatalf("delete-all auth namespace error = %v", err)
|
||||
}
|
||||
dir = base()
|
||||
tokenRemove = func(string) error { return fail }
|
||||
@@ -1745,11 +1821,6 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
t.Fatalf("delete-all quarantine error = %v", err)
|
||||
}
|
||||
dir = base()
|
||||
tokenDeleteKeychain = func() error { return fail }
|
||||
if err := DeleteAllTokenData(dir); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete-all legacy keychain error = %v", err)
|
||||
}
|
||||
dir = base()
|
||||
tokenDeleteSecure = func(string) error { return fail }
|
||||
if err := DeleteAllTokenData(dir); !errors.Is(err, fail) {
|
||||
t.Fatalf("delete-all secure error = %v", err)
|
||||
@@ -1765,6 +1836,9 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
resetClientIDFromMCP()
|
||||
SetClientID("client")
|
||||
fail := errors.New("fail")
|
||||
tokenLoadData = func(string) (*TokenData, error) {
|
||||
return &TokenData{AccessToken: "access", ClientID: "client", Source: "direct"}, nil
|
||||
}
|
||||
tokenParseURL = func(string) (*url.URL, error) { return nil, fail }
|
||||
if err := RevokeTokenRemote(context.Background()); !errors.Is(err, fail) {
|
||||
t.Fatalf("logout parse error = %v", err)
|
||||
@@ -1797,16 +1871,14 @@ func TestCrossPlatformCoverageTokenStorageAndRevocationCoverageEdges(t *testing.
|
||||
}
|
||||
|
||||
SetClientIDFromMCP("mcp-client")
|
||||
tokenRevokeURL = func() string { return "" }
|
||||
if err := RevokeTokenRemote(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tokenRevokeURL = func() string { return "https://revoke.test" }
|
||||
tokenLoadData = func(string) (*TokenData, error) { return nil, fail }
|
||||
if err := RevokeTokenRemote(context.Background()); err != nil {
|
||||
t.Fatal("missing token revoke should be a no-op")
|
||||
}
|
||||
tokenLoadData = func(string) (*TokenData, error) { return &TokenData{AccessToken: "access"}, nil }
|
||||
tokenMCPBaseURL = func() string { return "https://revoke.test" }
|
||||
tokenLoadData = func(string) (*TokenData, error) {
|
||||
return &TokenData{AccessToken: "access", ClientID: "mcp-client", Source: "mcp"}, nil
|
||||
}
|
||||
tokenJSONMarshal = func(any) ([]byte, error) { return nil, fail }
|
||||
if err := RevokeTokenRemote(context.Background()); !errors.Is(err, fail) {
|
||||
t.Fatalf("revoke marshal error = %v", err)
|
||||
@@ -2207,8 +2279,8 @@ func TestCrossPlatformCoverageProfilesCoverageEdges(t *testing.T) {
|
||||
}
|
||||
primaryOnly := &ProfilesConfig{PrimaryProfile: "primary", Profiles: []Profile{{Name: "primary", CorpID: "primary"}}}
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return primaryOnly, nil }
|
||||
if got, err := ResolveProfile(dir, ""); err != nil || got == nil || got.CorpID != "primary" {
|
||||
t.Fatalf("primary profile resolution = %#v %v", got, err)
|
||||
if got, err := ResolveProfile(dir, ""); err != nil || got != nil {
|
||||
t.Fatalf("primary-only profile resolution = %#v %v", got, err)
|
||||
}
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return &ProfilesConfig{}, nil }
|
||||
if got, err := ResolveProfile(dir, ""); err != nil || got != nil {
|
||||
@@ -2222,8 +2294,8 @@ func TestCrossPlatformCoverageProfilesCoverageEdges(t *testing.T) {
|
||||
}
|
||||
cfgNoCurrent := &ProfilesConfig{PrimaryProfile: "primary", Profiles: []Profile{{CorpID: "primary"}}}
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return cfgNoCurrent, nil }
|
||||
if got, err := resolveProfileForLoad(dir, ""); err != nil || got == nil || got.CorpID != "primary" {
|
||||
t.Fatalf("profile-for-load primary fallback = %#v %v", got, err)
|
||||
if got, err := resolveProfileForLoad(dir, ""); err != nil || got != nil {
|
||||
t.Fatalf("profile-for-load primary-only result = %#v %v", got, err)
|
||||
}
|
||||
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return cfg, nil }
|
||||
@@ -2242,6 +2314,10 @@ func TestCrossPlatformCoverageProfilesCoverageEdges(t *testing.T) {
|
||||
Profiles: []Profile{{Name: "primary", CorpID: "primary"}, {Name: "current", CorpID: "current"}, {Name: "previous", CorpID: "previous"}},
|
||||
}
|
||||
profilesSave = func(string, *ProfilesConfig) error { return nil }
|
||||
profilesLoadCorp = func(corpID string) (*TokenData, error) {
|
||||
return &TokenData{CorpID: corpID, AccessToken: "x"}, nil
|
||||
}
|
||||
profilesSaveCorp = func(string, *TokenData) error { return nil }
|
||||
profilesSyncLegacyMirror = func(string) error { return fail }
|
||||
if _, err := setCurrentProfileLocked(dir, "current"); !errors.Is(err, fail) {
|
||||
t.Fatalf("set-current mirror error = %v", err)
|
||||
@@ -2269,9 +2345,13 @@ func TestCrossPlatformCoverageProfilesCoverageEdges(t *testing.T) {
|
||||
deletedLegacy, deletedMarker := false, false
|
||||
profilesDeleteLegacy = func() error { deletedLegacy = true; return fail }
|
||||
profilesDeleteMarker = func(string) error { deletedMarker = true; return fail }
|
||||
if err := syncLegacyTokenMirrorLocked(dir); err != nil || !deletedLegacy || !deletedMarker {
|
||||
if err := syncLegacyTokenMirrorLocked(dir); !errors.Is(err, fail) || !deletedLegacy || deletedMarker {
|
||||
t.Fatalf("empty mirror cleanup = %v %v %v", err, deletedLegacy, deletedMarker)
|
||||
}
|
||||
profilesDeleteLegacy = func() error { return nil }
|
||||
if err := syncLegacyTokenMirrorLocked(dir); !errors.Is(err, fail) || !deletedMarker {
|
||||
t.Fatalf("empty mirror marker cleanup = %v %v", err, deletedMarker)
|
||||
}
|
||||
|
||||
normalizeProfilesConfig(nil)
|
||||
}
|
||||
@@ -2911,3 +2991,175 @@ func TestCrossPlatformCoverageDeviceFlowHighLevelCoverageEdges(t *testing.T) {
|
||||
}
|
||||
dfPrintBox(io.Discard, []string{"short"})
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageMultiAccountSelectorAndIdentityLoadEdges(t *testing.T) {
|
||||
cfg := &ProfilesConfig{
|
||||
Profiles: []Profile{
|
||||
{Name: "local-a", CorpID: "corp-a", CorpName: "Shared", UserID: "u1", UserName: "Alice"},
|
||||
{Name: "duplicate", CorpID: "corp-a", CorpName: "Shared", UserID: "u2", UserName: "Alice"},
|
||||
{Name: "duplicate", CorpID: "corp-b", CorpName: "Shared", UserID: "u3", UserName: "Bob"},
|
||||
{Name: "solo", CorpID: "corp-c", CorpName: "Solo", UserID: "u4", UserName: "Carol"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
selector string
|
||||
wantErr bool
|
||||
}{
|
||||
{"", true},
|
||||
{"missing:u1", true},
|
||||
{"Shared:u1", true},
|
||||
{"corp-a:u1", false},
|
||||
{"corp-a:Alice", true},
|
||||
{"corp-a:missing", true},
|
||||
{"corp-a", true},
|
||||
{"corp-c", false},
|
||||
{"Solo", false},
|
||||
{"local-a", false},
|
||||
{"duplicate", true},
|
||||
{"missing", true},
|
||||
} {
|
||||
_, _, err := resolveProfileSelection("", cfg, tc.selector)
|
||||
if (err != nil) != tc.wantErr {
|
||||
t.Fatalf("resolveProfileSelection(%q) error = %v", tc.selector, err)
|
||||
}
|
||||
}
|
||||
if _, _, err := resolveProfileSelection("", nil, "x"); err == nil {
|
||||
t.Fatal("nil profile config selection succeeded")
|
||||
}
|
||||
cfg.OrgCurrentProfiles = map[string]string{"corp-a": "corp-a:u2"}
|
||||
if got, _, err := resolveProfileSelection("", cfg, "corp-a"); err != nil || got.UserID != "u2" {
|
||||
t.Fatalf("organization current selection = %#v %v", got, err)
|
||||
}
|
||||
|
||||
for _, tc := range []struct {
|
||||
selector string
|
||||
wantErr bool
|
||||
exact bool
|
||||
}{
|
||||
{"", true, false},
|
||||
{"corp-a:u1", false, true},
|
||||
{"corp-a", false, false},
|
||||
{"Solo", false, false},
|
||||
{"local-a", false, true},
|
||||
{"duplicate", true, true},
|
||||
{"missing", true, false},
|
||||
} {
|
||||
_, exact, err := resolveProfileDeletionSelection(cfg, tc.selector)
|
||||
if (err != nil) != tc.wantErr || (!tc.wantErr && exact != tc.exact) {
|
||||
t.Fatalf("resolveProfileDeletionSelection(%q) = exact %v, %v", tc.selector, exact, err)
|
||||
}
|
||||
}
|
||||
if _, _, err := resolveProfileDeletionSelection(nil, "x"); err == nil {
|
||||
t.Fatal("nil deletion config succeeded")
|
||||
}
|
||||
|
||||
for _, selector := range []string{"", "corp-a", "missing", "Solo", "Shared"} {
|
||||
_, _ = resolveOrganizationCorpID(cfg, selector)
|
||||
}
|
||||
if got, err := resolveOrganizationCorpID(nil, "x"); err != nil || got != "" {
|
||||
t.Fatalf("nil organization resolution = %q %v", got, err)
|
||||
}
|
||||
if _, _, err := resolveOrganizationDefault(cfg, "missing", "missing", nil); err == nil {
|
||||
t.Fatal("empty organization default succeeded")
|
||||
}
|
||||
if got, _, err := resolveOrganizationDefault(cfg, "corp-c", "Solo", profilesForCorpID(cfg, "corp-c")); err != nil || got.UserID != "u4" {
|
||||
t.Fatalf("single account organization default = %#v %v", got, err)
|
||||
}
|
||||
if _, _, err := resolveOrganizationDefault(&ProfilesConfig{}, "corp-a", "corp-a", profilesForCorpID(cfg, "corp-a")); err == nil {
|
||||
t.Fatal("ambiguous organization default succeeded")
|
||||
}
|
||||
if got := profileSelectorCandidates([]*Profile{nil, &cfg.Profiles[1], &cfg.Profiles[0]}); len(got) != 2 || got[0] != "corp-a:u1" {
|
||||
t.Fatalf("selector candidates = %#v", got)
|
||||
}
|
||||
for selector, want := range map[string]bool{
|
||||
"": false,
|
||||
"corp-a:u1": true,
|
||||
"Shared:Bob": true,
|
||||
"corp-c": true,
|
||||
"local-a": true,
|
||||
"definitely-missing": false,
|
||||
} {
|
||||
if got := profileSelectorReferenceExists(cfg, selector); got != want {
|
||||
t.Fatalf("profileSelectorReferenceExists(%q) = %v", selector, got)
|
||||
}
|
||||
}
|
||||
if profileSelectorReferenceExists(nil, "corp-a") {
|
||||
t.Fatal("nil profile selector reference exists")
|
||||
}
|
||||
|
||||
oldAcquire := profilesAcquireDualLock
|
||||
oldEnsure := profilesEnsureMigration
|
||||
oldLoad := profilesLoad
|
||||
t.Cleanup(func() {
|
||||
profilesAcquireDualLock = oldAcquire
|
||||
profilesEnsureMigration = oldEnsure
|
||||
profilesLoad = oldLoad
|
||||
})
|
||||
profilesAcquireDualLock = func(context.Context, string) (*DualLock, error) { return &DualLock{}, nil }
|
||||
profilesEnsureMigration = func(string) error { return nil }
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return cfg, nil }
|
||||
if got, exact, err := ResolveProfileDeletionScope("cfg", "corp-a:u1"); err != nil || !exact || got.UserID != "u1" {
|
||||
t.Fatalf("deletion scope = %#v %v %v", got, exact, err)
|
||||
}
|
||||
profilesEnsureMigration = func(string) error { return errors.New("migration") }
|
||||
if _, _, err := ResolveProfileDeletionScope("cfg", "corp-a"); err == nil {
|
||||
t.Fatal("deletion scope migration failure succeeded")
|
||||
}
|
||||
profilesEnsureMigration = func(string) error { return nil }
|
||||
profilesLoad = func(string) (*ProfilesConfig, error) { return nil, errors.New("load") }
|
||||
if _, _, err := ResolveProfileDeletionScope("cfg", "corp-a"); err == nil {
|
||||
t.Fatal("deletion scope load failure succeeded")
|
||||
}
|
||||
|
||||
oldLoadIdentity := profilesLoadIdentity
|
||||
oldLoadCorp := profilesLoadCorp
|
||||
oldSaveIdentity := profilesSaveIdentity
|
||||
t.Cleanup(func() {
|
||||
profilesLoadIdentity = oldLoadIdentity
|
||||
profilesLoadCorp = oldLoadCorp
|
||||
profilesSaveIdentity = oldSaveIdentity
|
||||
})
|
||||
profile := Profile{CorpID: "corp-a", UserID: "u1"}
|
||||
profilesLoadIdentity = func(string, string) (*TokenData, error) { return &TokenData{AccessToken: "identity"}, nil }
|
||||
if got, err := loadTokenForProfileIdentity(profile); err != nil || got.AccessToken != "identity" {
|
||||
t.Fatalf("identity token load = %#v %v", got, err)
|
||||
}
|
||||
fail := errors.New("fail")
|
||||
profilesLoadIdentity = func(string, string) (*TokenData, error) { return nil, fail }
|
||||
if _, err := loadTokenForProfileIdentity(profile); !errors.Is(err, fail) {
|
||||
t.Fatalf("identity load failure = %v", err)
|
||||
}
|
||||
profilesLoadIdentity = func(string, string) (*TokenData, error) { return nil, ErrTokenDataNotFound }
|
||||
profilesLoadCorp = func(string) (*TokenData, error) { return nil, ErrTokenDataNotFound }
|
||||
if _, err := loadTokenForProfileIdentity(profile); !errors.Is(err, ErrTokenDataNotFound) {
|
||||
t.Fatalf("missing organization mirror = %v", err)
|
||||
}
|
||||
profilesLoadCorp = func(string) (*TokenData, error) { return nil, fail }
|
||||
if _, err := loadTokenForProfileIdentity(profile); !errors.Is(err, fail) {
|
||||
t.Fatalf("organization mirror failure = %v", err)
|
||||
}
|
||||
for _, data := range []*TokenData{
|
||||
{CorpID: "corp-a"},
|
||||
{CorpID: "corp-a", UserID: "other"},
|
||||
} {
|
||||
profilesLoadCorp = func(string) (*TokenData, error) { return data, nil }
|
||||
if _, err := loadTokenForProfileIdentity(profile); err == nil {
|
||||
t.Fatalf("invalid organization mirror %#v succeeded", data)
|
||||
}
|
||||
}
|
||||
profilesLoadCorp = func(string) (*TokenData, error) {
|
||||
return &TokenData{CorpID: "corp-a", UserID: "u1", AccessToken: "mirror"}, nil
|
||||
}
|
||||
profilesSaveIdentity = func(string, string, *TokenData) error { return fail }
|
||||
if _, err := loadTokenForProfileIdentity(profile); !errors.Is(err, fail) {
|
||||
t.Fatalf("identity repair save failure = %v", err)
|
||||
}
|
||||
profilesSaveIdentity = func(string, string, *TokenData) error { return nil }
|
||||
if got, err := loadTokenForProfileIdentity(profile); err != nil || got.AccessToken != "mirror" {
|
||||
t.Fatalf("identity repair = %#v %v", got, err)
|
||||
}
|
||||
if _, err := loadTokenForProfileIdentity(Profile{CorpID: "corp-a"}); err != nil {
|
||||
t.Fatalf("organization-only token load = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -75,15 +75,16 @@ var (
|
||||
)
|
||||
|
||||
type DeviceFlowProvider struct {
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
terminalBaseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
NoBrowser bool
|
||||
configDir string
|
||||
clientID string
|
||||
scope string
|
||||
baseURL string
|
||||
terminalBaseURL string
|
||||
logger *slog.Logger
|
||||
Output io.Writer
|
||||
httpClient *http.Client
|
||||
NoBrowser bool
|
||||
IdentityEnricher func(context.Context, *TokenData) error
|
||||
}
|
||||
|
||||
func NewDeviceFlowProvider(configDir string, logger *slog.Logger) *DeviceFlowProvider {
|
||||
@@ -267,9 +268,10 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
dfPrintStep(p.output(), 3, i18n.T("使用授权码换取 Access Token..."), 0)
|
||||
|
||||
oauthProvider := &OAuthProvider{
|
||||
configDir: p.configDir,
|
||||
clientID: p.clientID,
|
||||
logger: p.logger,
|
||||
configDir: p.configDir,
|
||||
clientID: p.clientID,
|
||||
logger: p.logger,
|
||||
IdentityEnricher: p.IdentityEnricher,
|
||||
}
|
||||
tokenData, err := deviceExchangeCode(oauthProvider, ctx, tokenResult.AuthCode)
|
||||
if err != nil {
|
||||
@@ -352,6 +354,9 @@ func (p *DeviceFlowProvider) loginOnce(ctx context.Context, attempt int) (*Token
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := oauthProvider.prepareLoginToken(ctx, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
if err := deviceSaveToken(p.configDir, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
@@ -717,6 +722,6 @@ func isInvalidGrantError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
msg := strings.ToLower(err.Error() + " " + httpStatusResponseBody(err))
|
||||
return strings.Contains(msg, "invalid_grant") || (strings.Contains(msg, "code") && strings.Contains(msg, "expired"))
|
||||
}
|
||||
|
||||
@@ -13,7 +13,50 @@
|
||||
|
||||
package auth
|
||||
|
||||
import "time"
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
const rejectedTokenRefreshFailureCooldown = 3 * time.Second
|
||||
|
||||
type rejectedTokenRefreshKey struct {
|
||||
configDir string
|
||||
profile string
|
||||
tokenDigest [sha256.Size]byte
|
||||
}
|
||||
|
||||
type rejectedTokenRefreshCall struct {
|
||||
done chan struct{}
|
||||
participants int
|
||||
token string
|
||||
err error
|
||||
}
|
||||
|
||||
type rejectedTokenRefreshFailure struct {
|
||||
at time.Time
|
||||
err error
|
||||
}
|
||||
|
||||
var rejectedTokenRefreshCoordinator = struct {
|
||||
sync.Mutex
|
||||
inFlight map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall
|
||||
failures map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure
|
||||
now func() time.Time
|
||||
}{
|
||||
inFlight: make(map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall),
|
||||
failures: make(map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure),
|
||||
now: time.Now,
|
||||
}
|
||||
|
||||
// MarkAccessTokenStale loads the persisted TokenData, sets ExpiresAt to a past
|
||||
// instant (preserving access_token and refresh_token), and writes it back. The
|
||||
@@ -39,3 +82,192 @@ func MarkAccessTokenStale(configDir string) error {
|
||||
data.ExpiresAt = time.Now().Add(-1 * time.Minute)
|
||||
return SaveTokenData(configDir, data)
|
||||
}
|
||||
|
||||
// ForceRefreshRejectedToken refreshes rejectedAccessToken only while it is
|
||||
// still the credential stored for the active profile. The compare and refresh
|
||||
// run under the same process + file lock used by ordinary expiry refresh, so a
|
||||
// late rejection cannot invalidate or refresh over a token another caller has
|
||||
// already rotated.
|
||||
//
|
||||
// When the stored token no longer matches, the newer token is returned without
|
||||
// calling the refresh endpoint. Refresh failures leave the stored credential in
|
||||
// place; login/logout remain the only owners of credential deletion.
|
||||
func (p *OAuthProvider) ForceRefreshRejectedToken(ctx context.Context, rejectedAccessToken string) (string, error) {
|
||||
if p == nil || strings.TrimSpace(p.configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
}
|
||||
rejectedAccessToken = strings.TrimSpace(rejectedAccessToken)
|
||||
if rejectedAccessToken == "" {
|
||||
return "", fmt.Errorf("rejected access token is empty")
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
|
||||
profile := strings.TrimSpace(RuntimeProfile())
|
||||
key := newRejectedTokenRefreshKey(p.configDir, profile, rejectedAccessToken)
|
||||
call, leader := beginRejectedTokenRefresh(key)
|
||||
if !leader {
|
||||
select {
|
||||
case <-call.done:
|
||||
return call.token, call.err
|
||||
case <-ctx.Done():
|
||||
return "", ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
token, err, recordFailure := p.forceRefreshRejectedTokenOnce(ctx, profile, rejectedAccessToken, key)
|
||||
finishRejectedTokenRefresh(key, call, token, err, recordFailure)
|
||||
return token, err
|
||||
}
|
||||
|
||||
func (p *OAuthProvider) forceRefreshRejectedTokenOnce(
|
||||
ctx context.Context,
|
||||
profile string,
|
||||
rejectedAccessToken string,
|
||||
key rejectedTokenRefreshKey,
|
||||
) (string, error, bool) {
|
||||
lock, err := oauthAcquireLock(ctx, p.configDir)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("acquiring dual lock: %w", err), false
|
||||
}
|
||||
defer lock.Release()
|
||||
|
||||
data, err := loadOAuthTokenUnderHeldLock(p.configDir, profile)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("reload rejected token: %w", err), false
|
||||
}
|
||||
current := strings.TrimSpace(data.AccessToken)
|
||||
if current == "" {
|
||||
clearRejectedTokenRefreshFailure(key)
|
||||
return "", fmt.Errorf("stored access token is empty"), false
|
||||
}
|
||||
if current != rejectedAccessToken {
|
||||
clearRejectedTokenRefreshFailure(key)
|
||||
return current, nil, false
|
||||
}
|
||||
if cachedErr := recentRejectedTokenRefreshFailure(key); cachedErr != nil {
|
||||
return "", cachedErr, false
|
||||
}
|
||||
if !data.IsRefreshTokenValid() {
|
||||
return "", fmt.Errorf("refresh_token 已过期"), true
|
||||
}
|
||||
if err := preflightTokenRefreshPersistence(p.configDir, data); err != nil {
|
||||
return "", fmt.Errorf("本地登录态无法安全更新: %w", err), true
|
||||
}
|
||||
|
||||
refreshed, err := oauthRefreshToken(p, ctx, data)
|
||||
if err != nil {
|
||||
return "", err, true
|
||||
}
|
||||
if refreshed == nil || strings.TrimSpace(refreshed.AccessToken) == "" {
|
||||
return "", fmt.Errorf("force refresh returned empty access token"), true
|
||||
}
|
||||
return strings.TrimSpace(refreshed.AccessToken), nil, false
|
||||
}
|
||||
|
||||
func newRejectedTokenRefreshKey(configDir, profile, rejectedAccessToken string) rejectedTokenRefreshKey {
|
||||
canonicalDir := filepath.Clean(configDir)
|
||||
if absolute, err := filepath.Abs(configDir); err == nil {
|
||||
canonicalDir = filepath.Clean(absolute)
|
||||
}
|
||||
return rejectedTokenRefreshKey{
|
||||
configDir: canonicalDir,
|
||||
profile: strings.TrimSpace(profile),
|
||||
tokenDigest: sha256.Sum256([]byte(strings.TrimSpace(rejectedAccessToken))),
|
||||
}
|
||||
}
|
||||
|
||||
func beginRejectedTokenRefresh(key rejectedTokenRefreshKey) (*rejectedTokenRefreshCall, bool) {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
defer rejectedTokenRefreshCoordinator.Unlock()
|
||||
|
||||
if call := rejectedTokenRefreshCoordinator.inFlight[key]; call != nil {
|
||||
call.participants++
|
||||
return call, false
|
||||
}
|
||||
call := &rejectedTokenRefreshCall{done: make(chan struct{}), participants: 1}
|
||||
rejectedTokenRefreshCoordinator.inFlight[key] = call
|
||||
return call, true
|
||||
}
|
||||
|
||||
func finishRejectedTokenRefresh(
|
||||
key rejectedTokenRefreshKey,
|
||||
call *rejectedTokenRefreshCall,
|
||||
token string,
|
||||
err error,
|
||||
recordFailure bool,
|
||||
) {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
defer rejectedTokenRefreshCoordinator.Unlock()
|
||||
|
||||
call.token = token
|
||||
call.err = err
|
||||
delete(rejectedTokenRefreshCoordinator.inFlight, key)
|
||||
if err == nil {
|
||||
delete(rejectedTokenRefreshCoordinator.failures, key)
|
||||
} else if recordFailure && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
|
||||
now := time.Now()
|
||||
if rejectedTokenRefreshCoordinator.now != nil {
|
||||
now = rejectedTokenRefreshCoordinator.now()
|
||||
}
|
||||
rejectedTokenRefreshCoordinator.failures[key] = rejectedTokenRefreshFailure{at: now, err: err}
|
||||
}
|
||||
close(call.done)
|
||||
}
|
||||
|
||||
func recentRejectedTokenRefreshFailure(key rejectedTokenRefreshKey) error {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
defer rejectedTokenRefreshCoordinator.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if rejectedTokenRefreshCoordinator.now != nil {
|
||||
now = rejectedTokenRefreshCoordinator.now()
|
||||
}
|
||||
for failureKey, failure := range rejectedTokenRefreshCoordinator.failures {
|
||||
age := now.Sub(failure.at)
|
||||
if age < 0 || age >= rejectedTokenRefreshFailureCooldown {
|
||||
delete(rejectedTokenRefreshCoordinator.failures, failureKey)
|
||||
}
|
||||
}
|
||||
if failure, ok := rejectedTokenRefreshCoordinator.failures[key]; ok {
|
||||
return failure.err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func clearRejectedTokenRefreshFailure(key rejectedTokenRefreshKey) {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
delete(rejectedTokenRefreshCoordinator.failures, key)
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
}
|
||||
|
||||
// loadOAuthTokenUnderHeldLock mirrors LoadTokenDataForProfile without taking a
|
||||
// second, non-reentrant auth lock. Opaque edition storage hooks (for example
|
||||
// Wukong's encrypted .data file) are read inside the caller's dual lock so the
|
||||
// compare-and-refresh decision covers both Core and embedded storage.
|
||||
func loadOAuthTokenUnderHeldLock(configDir, profile string) (*TokenData, error) {
|
||||
hooks := edition.Get()
|
||||
if hooks.LoadToken == nil {
|
||||
data, err := oauthLoadTokenLocked(configDir, profile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if data == nil {
|
||||
return nil, fmt.Errorf("stored token data is empty")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return nil, fmt.Errorf("profile selection is not supported by the current auth backend")
|
||||
}
|
||||
blob, err := hooks.LoadToken(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var data TokenData
|
||||
if err := json.Unmarshal(blob, &data); err != nil {
|
||||
return nil, fmt.Errorf("parsing token data from hook: %w", err)
|
||||
}
|
||||
return &data, nil
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -52,6 +53,14 @@ func TokenAccountForCorpID(corpID string) string {
|
||||
return keychain.AccountToken + ":" + strings.TrimSpace(corpID)
|
||||
}
|
||||
|
||||
// TokenAccountForIdentity returns the stable keychain account used for one
|
||||
// DingTalk identity. The hash avoids collisions caused by delimiter escaping
|
||||
// or keychain/file-name restrictions.
|
||||
func TokenAccountForIdentity(corpID, userID string) string {
|
||||
identity := strings.TrimSpace(corpID) + "\x00" + strings.TrimSpace(userID)
|
||||
return fmt.Sprintf("%s:id:%x", keychain.AccountToken, sha256.Sum256([]byte(identity)))
|
||||
}
|
||||
|
||||
// SaveTokenDataKeychainForCorpID saves TokenData to a corp-scoped keychain slot.
|
||||
func SaveTokenDataKeychainForCorpID(corpID string, data *TokenData) error {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
@@ -61,6 +70,16 @@ func SaveTokenDataKeychainForCorpID(corpID string, data *TokenData) error {
|
||||
return saveTokenDataKeychainAccount(TokenAccountForCorpID(corpID), data)
|
||||
}
|
||||
|
||||
// SaveTokenDataKeychainForIdentity saves TokenData to an identity-scoped slot.
|
||||
func SaveTokenDataKeychainForIdentity(corpID, userID string, data *TokenData) error {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
userID = strings.TrimSpace(userID)
|
||||
if corpID == "" || userID == "" {
|
||||
return fmt.Errorf("corpId and userId are required for identity token storage")
|
||||
}
|
||||
return saveTokenDataKeychainAccount(TokenAccountForIdentity(corpID, userID), data)
|
||||
}
|
||||
|
||||
func saveTokenDataKeychainAccount(account string, data *TokenData) error {
|
||||
jsonData, err := authKeychainMarshal(data, "", " ")
|
||||
if err != nil {
|
||||
@@ -93,6 +112,16 @@ func LoadTokenDataKeychainForCorpID(corpID string) (*TokenData, error) {
|
||||
return loadTokenDataKeychainAccount(TokenAccountForCorpID(corpID))
|
||||
}
|
||||
|
||||
// LoadTokenDataKeychainForIdentity loads TokenData from an identity-scoped slot.
|
||||
func LoadTokenDataKeychainForIdentity(corpID, userID string) (*TokenData, error) {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
userID = strings.TrimSpace(userID)
|
||||
if corpID == "" || userID == "" {
|
||||
return nil, fmt.Errorf("corpId and userId are required for identity token storage")
|
||||
}
|
||||
return loadTokenDataKeychainAccount(TokenAccountForIdentity(corpID, userID))
|
||||
}
|
||||
|
||||
func loadTokenDataKeychainAccount(account string) (*TokenData, error) {
|
||||
jsonStr, err := authKeychainGet(keychain.Service, account)
|
||||
if err != nil {
|
||||
@@ -127,6 +156,9 @@ func preflightTokenPersistence(configDir string) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("load token profiles: %w", err)
|
||||
}
|
||||
if err := ensureProfilesWritable(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, profile := range cfg.Profiles {
|
||||
// LoadProfiles normalizes away blank and duplicate corp IDs.
|
||||
corpID := profile.CorpID
|
||||
@@ -137,6 +169,19 @@ func preflightTokenPersistence(configDir string) error {
|
||||
)
|
||||
}
|
||||
}
|
||||
for _, profile := range cfg.Profiles {
|
||||
corpID := strings.TrimSpace(profile.CorpID)
|
||||
userID := strings.TrimSpace(profile.UserID)
|
||||
if corpID == "" || userID == "" {
|
||||
continue
|
||||
}
|
||||
if _, err := LoadTokenDataKeychainForIdentity(corpID, userID); err != nil && !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return fmt.Errorf(
|
||||
"identity token slot %q is unreadable; remove only this account with `dws auth logout --profile %q`, or use `dws auth reset` only when discarding all local profiles: %w",
|
||||
TokenAccountForIdentity(corpID, userID), ProfileSelector(profile), err,
|
||||
)
|
||||
}
|
||||
}
|
||||
if err := authValidateEntries(keychain.Service); err != nil {
|
||||
return fmt.Errorf(
|
||||
"auth token ciphertext inventory is unreadable; on macOS first try `env -u DWS_DISABLE_KEYCHAIN dws auth migrate-keychain --to file-dek --dry-run`; if the ciphertext is damaged, use `dws auth reset` only when discarding all local profiles: %w",
|
||||
@@ -149,10 +194,17 @@ func preflightTokenPersistence(configDir string) error {
|
||||
// preflightTokenRefreshPersistence checks only the slots a refresh can write.
|
||||
// An unrelated broken profile must not prevent the current profile from using
|
||||
// its still-valid credentials.
|
||||
func preflightTokenRefreshPersistence(data *TokenData) error {
|
||||
func preflightTokenRefreshPersistence(configDir string, data *TokenData) error {
|
||||
if h := edition.Get(); h.SaveToken != nil {
|
||||
return nil
|
||||
}
|
||||
cfg, err := LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureProfilesWritable(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := LoadTokenDataKeychain(); err != nil && !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return fmt.Errorf("legacy token slot %q is unreadable: %w", keychain.AccountToken, err)
|
||||
@@ -161,8 +213,22 @@ func preflightTokenRefreshPersistence(data *TokenData) error {
|
||||
return nil
|
||||
}
|
||||
corpID := strings.TrimSpace(data.CorpID)
|
||||
if _, err := LoadTokenDataKeychainForCorpID(corpID); err != nil && !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return fmt.Errorf("profile token slot %q is unreadable: %w", TokenAccountForCorpID(corpID), err)
|
||||
userID := strings.TrimSpace(data.UserID)
|
||||
if userID != "" {
|
||||
if _, err := LoadTokenDataKeychainForIdentity(corpID, userID); err != nil && !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return fmt.Errorf("identity token slot %q is unreadable: %w", TokenAccountForIdentity(corpID, userID), err)
|
||||
}
|
||||
}
|
||||
checkOrganizationMirror := true
|
||||
if _, _, exact := ParseIdentitySelector(RuntimeProfile()); exact {
|
||||
checkOrganizationMirror =
|
||||
exactProfileSelectorForCorp(cfg, corpID, cfg.OrgCurrentProfiles[corpID]) ==
|
||||
profileSelector(corpID, userID)
|
||||
}
|
||||
if checkOrganizationMirror {
|
||||
if _, err := LoadTokenDataKeychainForCorpID(corpID); err != nil && !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return fmt.Errorf("profile token slot %q is unreadable: %w", TokenAccountForCorpID(corpID), err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -181,6 +247,16 @@ func DeleteTokenDataKeychainForCorpID(corpID string) error {
|
||||
return authKeychainRemove(keychain.Service, TokenAccountForCorpID(corpID))
|
||||
}
|
||||
|
||||
// DeleteTokenDataKeychainForIdentity removes one identity-scoped token.
|
||||
func DeleteTokenDataKeychainForIdentity(corpID, userID string) error {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
userID = strings.TrimSpace(userID)
|
||||
if corpID == "" || userID == "" {
|
||||
return fmt.Errorf("corpId and userId are required for identity token storage")
|
||||
}
|
||||
return authKeychainRemove(keychain.Service, TokenAccountForIdentity(corpID, userID))
|
||||
}
|
||||
|
||||
// TokenDataExistsKeychain checks if token data exists in keychain.
|
||||
func TokenDataExistsKeychain() bool {
|
||||
return authKeychainExists(keychain.Service, keychain.AccountToken)
|
||||
@@ -195,6 +271,16 @@ func TokenDataExistsKeychainForCorpID(corpID string) bool {
|
||||
return authKeychainExists(keychain.Service, TokenAccountForCorpID(corpID))
|
||||
}
|
||||
|
||||
// TokenDataExistsKeychainForIdentity checks if an identity-scoped token exists.
|
||||
func TokenDataExistsKeychainForIdentity(corpID, userID string) bool {
|
||||
corpID = strings.TrimSpace(corpID)
|
||||
userID = strings.TrimSpace(userID)
|
||||
if corpID == "" || userID == "" {
|
||||
return false
|
||||
}
|
||||
return authKeychainExists(keychain.Service, TokenAccountForIdentity(corpID, userID))
|
||||
}
|
||||
|
||||
// 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.
|
||||
|
||||
@@ -47,8 +47,10 @@ func (m *Manager) GetToken() (string, string, error) {
|
||||
}
|
||||
return token, "file", nil
|
||||
}
|
||||
|
||||
return "", "", fmt.Errorf("%s", i18n.T("未找到认证信息,请运行 dws auth login"))
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return "", "", fmt.Errorf("load legacy token: %w", err)
|
||||
}
|
||||
return "", "", fmt.Errorf("%s: %w", i18n.T("未找到认证信息,请运行 dws auth login"), ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
func (m *Manager) GetMCPURL() (string, error) {
|
||||
|
||||
@@ -237,12 +237,21 @@ func (p *OAuthProvider) postJSON(ctx context.Context, endpoint string, body any)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading response: %w", err)
|
||||
}
|
||||
data, readErr := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("HTTP %d: %s", resp.StatusCode, truncateBody(data, 200))
|
||||
// Preserve structured HTTP status semantics even when the response
|
||||
// body is truncated. The body is diagnostic-only here, so read it
|
||||
// best-effort and classify retryability from the status code.
|
||||
if readErr != nil {
|
||||
data = nil
|
||||
}
|
||||
return nil, &HTTPStatusError{
|
||||
StatusCode: resp.StatusCode,
|
||||
responseBody: truncateBody(data, 200),
|
||||
}
|
||||
}
|
||||
if readErr != nil {
|
||||
return nil, fmt.Errorf("reading response: %w", readErr)
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
@@ -293,6 +302,8 @@ func (p *OAuthProvider) parseMCPTokenResponse(body []byte) (*TokenData, error) {
|
||||
CorpName string `json:"corpName"`
|
||||
CorpNameSnake string `json:"corp_name"`
|
||||
OrgName string `json:"orgName"`
|
||||
UserID string `json:"userId"`
|
||||
UserName string `json:"userName"`
|
||||
// Error fields (when request fails)
|
||||
ErrorCode string `json:"errorCode,omitempty"`
|
||||
ErrorMsg string `json:"errorMsg,omitempty"`
|
||||
@@ -320,6 +331,8 @@ func (p *OAuthProvider) parseMCPTokenResponse(body []byte) (*TokenData, error) {
|
||||
RefreshExpAt: now.Add(config.DefaultRefreshTokenLifetime),
|
||||
CorpID: resp.CorpID,
|
||||
CorpName: firstNonEmpty(resp.CorpName, resp.CorpNameSnake, resp.OrgName),
|
||||
UserID: resp.UserID,
|
||||
UserName: resp.UserName,
|
||||
}
|
||||
if resp.PersistentCode != "" {
|
||||
data.PersistentCode = resp.PersistentCode
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type postJSONRoundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f postJSONRoundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
type oauthBrokenBody struct{}
|
||||
|
||||
func (oauthBrokenBody) Read([]byte) (int, error) { return 0, io.ErrUnexpectedEOF }
|
||||
func (oauthBrokenBody) Close() error { return nil }
|
||||
|
||||
func TestCrossPlatformCoveragePostJSONTruncatedErrorBodyKeepsHTTPStatus(t *testing.T) {
|
||||
for _, status := range []int{http.StatusTooManyRequests, http.StatusServiceUnavailable} {
|
||||
t.Run(http.StatusText(status), func(t *testing.T) {
|
||||
provider := &OAuthProvider{httpClient: &http.Client{Transport: postJSONRoundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: status, Body: oauthBrokenBody{}, Header: make(http.Header)}, nil
|
||||
})}}
|
||||
|
||||
_, err := provider.postJSON(context.Background(), "https://oauth.test/token", map[string]string{"grantType": "refresh_token"})
|
||||
var statusErr *HTTPStatusError
|
||||
if !errors.As(err, &statusErr) || statusErr.StatusCode != status {
|
||||
t.Fatalf("postJSON() error = %v, want HTTPStatusError %d", err, status)
|
||||
}
|
||||
if errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
t.Fatalf("HTTP status error should not expose diagnostic body read failure: %v", err)
|
||||
}
|
||||
if got := ClassifyRefreshFailure(err); got != RefreshFailureTransient {
|
||||
t.Fatalf("ClassifyRefreshFailure() = %s, want transient", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePostJSONOKTruncatedBodyIsTransient(t *testing.T) {
|
||||
provider := &OAuthProvider{httpClient: &http.Client{Transport: postJSONRoundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: http.StatusOK, Body: oauthBrokenBody{}, Header: make(http.Header)}, nil
|
||||
})}}
|
||||
|
||||
_, err := provider.postJSON(context.Background(), "https://oauth.test/token", map[string]string{"grantType": "refresh_token"})
|
||||
if !errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
t.Fatalf("postJSON() error = %v, want io.ErrUnexpectedEOF", err)
|
||||
}
|
||||
if got := ClassifyRefreshFailure(err); got != RefreshFailureTransient {
|
||||
t.Fatalf("ClassifyRefreshFailure() = %s, want transient", got)
|
||||
}
|
||||
}
|
||||
+169
-16
@@ -18,6 +18,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"html"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
@@ -28,6 +29,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/i18n"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
)
|
||||
|
||||
// oauthHTTPClient is a dedicated HTTP client for OAuth operations with
|
||||
@@ -44,6 +46,7 @@ var (
|
||||
oauthPollInterval = 5 * time.Second
|
||||
oauthSuccessPause = 2 * time.Second
|
||||
oauthLoadToken = LoadTokenData
|
||||
oauthLoadTokenLocked = loadTokenDataForProfileLocked
|
||||
oauthAcquireLock = AcquireDualLock
|
||||
oauthMarkProfile = MarkProfileStatus
|
||||
oauthFetchClientID = FetchClientIDFromMCP
|
||||
@@ -73,6 +76,9 @@ type OAuthProvider struct {
|
||||
httpClient *http.Client
|
||||
NoBrowser bool
|
||||
TargetCorpID string
|
||||
// IdentityEnricher resolves userId/userName/corpName while the freshly
|
||||
// exchanged access token is still only in memory.
|
||||
IdentityEnricher func(context.Context, *TokenData) error
|
||||
}
|
||||
|
||||
// NewOAuthProvider creates a new OAuth provider.
|
||||
@@ -110,6 +116,12 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
// Smart degradation: try silent refresh before opening browser.
|
||||
if !force {
|
||||
data, err := oauthLoadToken(p.configDir)
|
||||
if err != nil && !errors.Is(err, ErrTokenDataNotFound) && !os.IsNotExist(err) {
|
||||
if preflightErr := preflightTokenPersistence(p.configDir); preflightErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("本地登录态无法安全更新"), preflightErr)
|
||||
}
|
||||
return nil, fmt.Errorf("load existing access token: %w", err)
|
||||
}
|
||||
if err == nil {
|
||||
// Case 1: access_token still valid — no action needed.
|
||||
if data.IsAccessTokenValid() {
|
||||
@@ -175,6 +187,14 @@ 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)
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.flow.start",
|
||||
"client_id", strings.TrimSpace(p.clientID),
|
||||
"target_corp_id", strings.TrimSpace(p.TargetCorpID),
|
||||
"callback_port", port,
|
||||
"force", force,
|
||||
"no_browser", p.NoBrowser,
|
||||
)
|
||||
|
||||
// Channel to pass callback result (token data or error with CLI auth status)
|
||||
type callbackResult struct {
|
||||
@@ -205,6 +225,11 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
if code == "" {
|
||||
code = r.URL.Query().Get("code")
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.callback.received",
|
||||
"callback_port", port,
|
||||
"has_authorization_code", code != "",
|
||||
)
|
||||
|
||||
// Check state and handle page refresh or concurrent requests
|
||||
callbackTokenMu.Lock()
|
||||
@@ -264,6 +289,11 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
// Exchange code for token
|
||||
tokenData, exchangeErr := oauthExchange(p, ctx, code)
|
||||
if exchangeErr != nil {
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.token_exchange.error",
|
||||
"callback_port", port,
|
||||
"error", exchangeErr,
|
||||
)
|
||||
// Clear in-progress state on error
|
||||
callbackTokenMu.Lock()
|
||||
if callbackCodeInProgress == code {
|
||||
@@ -272,13 +302,23 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
callbackTokenMu.Unlock()
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
_, _ = fmt.Fprintf(w, "<html><body><h1>授权失败</h1><p>%s</p></body></html>", exchangeErr.Error())
|
||||
_, _ = fmt.Fprintf(w, "<html><body><h1>授权失败</h1><p>%s</p></body></html>", html.EscapeString(oauthExchangeDisplayError(exchangeErr)))
|
||||
select {
|
||||
case resultCh <- callbackResult{err: exchangeErr}:
|
||||
default:
|
||||
}
|
||||
return
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.token_exchange.done",
|
||||
"callback_port", port,
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"user_id", strings.TrimSpace(tokenData.UserID),
|
||||
"user_name", strings.TrimSpace(tokenData.UserName),
|
||||
"source", strings.TrimSpace(tokenData.Source),
|
||||
"access_expires_at", tokenData.ExpiresAt,
|
||||
"refresh_expires_at", tokenData.RefreshExpAt,
|
||||
)
|
||||
|
||||
// Mark as processed immediately after successful exchange
|
||||
callbackTokenMu.Lock()
|
||||
@@ -302,6 +342,13 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
denialReason = classifyDenialReason(authStatus, os.Getenv("DWS_CHANNEL"))
|
||||
}
|
||||
cliAuthEnabled := denialReason == ""
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.organization_access.checked",
|
||||
"callback_port", port,
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"enabled", cliAuthEnabled,
|
||||
"denial_reason", denialReason,
|
||||
)
|
||||
|
||||
// Server-provided errorMsg (nil-safe), surfaced both on the page and to
|
||||
// the terminal so portal can update copy without releasing the CLI.
|
||||
@@ -555,9 +602,30 @@ continueLogin:
|
||||
|
||||
// Save token data with associated client ID for refresh
|
||||
tokenData.ClientID = p.clientID
|
||||
if err := oauthSaveToken(p.configDir, tokenData); err != nil {
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.persistence.start",
|
||||
"callback_port", port,
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"user_id", strings.TrimSpace(tokenData.UserID),
|
||||
"user_name", strings.TrimSpace(tokenData.UserName),
|
||||
)
|
||||
if err := p.persistLoginToken(ctx, tokenData); err != nil {
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.persistence.error",
|
||||
"callback_port", port,
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"user_id", strings.TrimSpace(tokenData.UserID),
|
||||
"error", err,
|
||||
)
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.persistence.done",
|
||||
"callback_port", port,
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"user_id", strings.TrimSpace(tokenData.UserID),
|
||||
"user_name", strings.TrimSpace(tokenData.UserName),
|
||||
)
|
||||
|
||||
// Persist app credentials (with secret) if using custom client credentials.
|
||||
// MUST run BEFORE os.Setenv below to avoid env-matching short circuit.
|
||||
@@ -575,35 +643,63 @@ continueLogin:
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
// GetAccessToken returns a valid access token, auto-refreshing if needed.
|
||||
// Uses a file lock with double-check pattern to prevent concurrent refresh
|
||||
// from multiple CLI processes.
|
||||
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
|
||||
func oauthExchangeDisplayError(err error) string {
|
||||
var statusErr *HTTPStatusError
|
||||
if errors.As(err, &statusErr) && statusErr != nil {
|
||||
return fmt.Sprintf("HTTP %d: token exchange failed", statusErr.StatusCode)
|
||||
}
|
||||
return err.Error()
|
||||
}
|
||||
|
||||
// GetTokenSnapshot returns a valid token together with its expiry metadata.
|
||||
// Storage and refresh failures retain their original cause; only a confirmed
|
||||
// missing credential is reported as ErrTokenDataNotFound.
|
||||
func (p *OAuthProvider) GetTokenSnapshot(ctx context.Context) (*TokenData, error) {
|
||||
data, err := oauthLoadToken(p.configDir)
|
||||
if err != nil {
|
||||
return "", errors.New(i18n.T("未登录,请运行 dws auth login"))
|
||||
if errors.Is(err, ErrTokenDataNotFound) || os.IsNotExist(err) {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("未登录,请运行 dws auth login"), ErrTokenDataNotFound)
|
||||
}
|
||||
return nil, fmt.Errorf("load access token: %w", err)
|
||||
}
|
||||
|
||||
// Fast path: access_token still valid — no lock needed.
|
||||
if data.IsAccessTokenValid() {
|
||||
return data.AccessToken, nil
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// Slow path: token expired — try locked refresh.
|
||||
if data.IsRefreshTokenValid() {
|
||||
refreshed, rErr := p.lockedRefresh(ctx)
|
||||
if rErr == nil {
|
||||
return refreshed.AccessToken, nil
|
||||
return refreshed, nil
|
||||
}
|
||||
// A network, timeout, rate-limit or 5xx failure does not invalidate the
|
||||
// refresh credential. Keep the profile active so a long-running source
|
||||
// can retry after backoff. Terminal and unknown failures remain fatal.
|
||||
if ClassifyRefreshFailure(rErr) != RefreshFailureTransient {
|
||||
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
|
||||
}
|
||||
_ = oauthMarkProfile(p.configDir, data.CorpID, ProfileStatusExpired)
|
||||
if p.logger != nil {
|
||||
p.logger.Warn(i18n.T("refresh_token 刷新失败"), "error", rErr)
|
||||
}
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
|
||||
} else {
|
||||
_ = oauthMarkProfile(p.configDir, data.CorpID, ProfileStatusExpired)
|
||||
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
|
||||
}
|
||||
|
||||
return "", errors.New(i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"))
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"), ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
// GetAccessToken returns a valid access token, auto-refreshing if needed.
|
||||
// Uses a file lock with double-check pattern to prevent concurrent refresh
|
||||
// from multiple CLI processes.
|
||||
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
|
||||
data, err := p.GetTokenSnapshot(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strings.TrimSpace(data.AccessToken), nil
|
||||
}
|
||||
|
||||
// lockedRefresh attempts to refresh the token while holding dual-layer locks.
|
||||
@@ -635,7 +731,7 @@ func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
|
||||
|
||||
// Double-check: re-load from disk — another goroutine/process may have refreshed
|
||||
// while we were waiting for the lock.
|
||||
data, err := oauthLoadToken(p.configDir)
|
||||
data, err := loadOAuthTokenUnderHeldLock(p.configDir, RuntimeProfile())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -654,7 +750,7 @@ func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
|
||||
if !data.IsRefreshTokenValid() {
|
||||
return nil, fmt.Errorf("refresh_token 已过期")
|
||||
}
|
||||
if err := preflightTokenRefreshPersistence(data); err != nil {
|
||||
if err := preflightTokenRefreshPersistence(p.configDir, data); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("本地登录态无法安全更新"), err)
|
||||
}
|
||||
|
||||
@@ -672,14 +768,71 @@ func (p *OAuthProvider) ExchangeAuthCode(ctx context.Context, authCode, uid stri
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("换取 token 失败"), err)
|
||||
}
|
||||
if uid != "" {
|
||||
tokenData.UserID = uid
|
||||
tokenData.UserID = strings.TrimSpace(uid)
|
||||
if err := p.persistKnownLoginToken(tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
return tokenData, nil
|
||||
}
|
||||
if err := oauthSaveToken(p.configDir, tokenData); err != nil {
|
||||
if err := p.persistLoginToken(ctx, tokenData); err != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("保存 token 失败"), err)
|
||||
}
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
func (p *OAuthProvider) persistLoginToken(ctx context.Context, tokenData *TokenData) error {
|
||||
corpID, userID, userName := "", "", ""
|
||||
if tokenData != nil {
|
||||
corpID = strings.TrimSpace(tokenData.CorpID)
|
||||
userID = strings.TrimSpace(tokenData.UserID)
|
||||
userName = strings.TrimSpace(tokenData.UserName)
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.identity.before_enrich",
|
||||
"corp_id", corpID,
|
||||
"user_id", userID,
|
||||
"user_name", userName,
|
||||
)
|
||||
if err := p.prepareLoginToken(ctx, tokenData); err != nil {
|
||||
return err
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.login.oauth.identity.after_enrich",
|
||||
"corp_id", strings.TrimSpace(tokenData.CorpID),
|
||||
"user_id", strings.TrimSpace(tokenData.UserID),
|
||||
"user_name", strings.TrimSpace(tokenData.UserName),
|
||||
)
|
||||
if err := oauthSaveToken(p.configDir, tokenData); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *OAuthProvider) prepareLoginToken(ctx context.Context, tokenData *TokenData) error {
|
||||
if tokenData == nil {
|
||||
return fmt.Errorf("token data is empty")
|
||||
}
|
||||
if p != nil && p.IdentityEnricher != nil {
|
||||
if err := p.IdentityEnricher(ctx, tokenData); err != nil {
|
||||
return fmt.Errorf("resolve login identity: %w", err)
|
||||
}
|
||||
}
|
||||
if strings.TrimSpace(tokenData.CorpID) != "" && strings.TrimSpace(tokenData.UserID) == "" {
|
||||
return fmt.Errorf("resolve login identity: userId is required for corpId %q", tokenData.CorpID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (p *OAuthProvider) persistKnownLoginToken(tokenData *TokenData) error {
|
||||
if tokenData == nil {
|
||||
return fmt.Errorf("token data is empty")
|
||||
}
|
||||
if strings.TrimSpace(tokenData.CorpID) != "" && strings.TrimSpace(tokenData.UserID) == "" {
|
||||
return fmt.Errorf("resolve login identity: userId is required for corpId %q", tokenData.CorpID)
|
||||
}
|
||||
return oauthSaveToken(p.configDir, tokenData)
|
||||
}
|
||||
|
||||
// Logout clears all stored credentials.
|
||||
func (p *OAuthProvider) Logout() error {
|
||||
return DeleteTokenData(p.configDir)
|
||||
|
||||
@@ -18,6 +18,8 @@ import (
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
)
|
||||
|
||||
type oauthLoginFixture struct {
|
||||
@@ -33,8 +35,79 @@ type oauthLoginFixture struct {
|
||||
exchangeError atomic.Bool
|
||||
}
|
||||
|
||||
type oauthLoginResult struct {
|
||||
token *TokenData
|
||||
err error
|
||||
}
|
||||
|
||||
type oauthHTTPResult struct {
|
||||
status int
|
||||
body string
|
||||
err error
|
||||
}
|
||||
|
||||
const oauthTestWaitTimeout = 5 * time.Second
|
||||
|
||||
func isolateOAuthPersistence(t *testing.T) {
|
||||
t.Helper()
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
cleanupKeychain(t)
|
||||
}
|
||||
|
||||
func startOAuthLogin(t *testing.T, parent context.Context, f *oauthLoginFixture) <-chan oauthLoginResult {
|
||||
t.Helper()
|
||||
ctx, cancel := context.WithCancel(parent)
|
||||
done := make(chan oauthLoginResult, 1)
|
||||
finished := make(chan struct{})
|
||||
go func() {
|
||||
defer close(finished)
|
||||
token, err := f.provider.Login(ctx, true)
|
||||
done <- oauthLoginResult{token: token, err: err}
|
||||
}()
|
||||
t.Cleanup(func() {
|
||||
cancel()
|
||||
select {
|
||||
case <-finished:
|
||||
case <-time.After(oauthTestWaitTimeout):
|
||||
t.Errorf("OAuth Login goroutine did not stop after cancellation")
|
||||
}
|
||||
})
|
||||
return done
|
||||
}
|
||||
|
||||
func awaitOAuthLogin(t *testing.T, done <-chan oauthLoginResult) oauthLoginResult {
|
||||
t.Helper()
|
||||
select {
|
||||
case result := <-done:
|
||||
return result
|
||||
case <-time.After(oauthTestWaitTimeout):
|
||||
t.Fatal("timed out waiting for OAuth Login")
|
||||
return oauthLoginResult{}
|
||||
}
|
||||
}
|
||||
|
||||
func waitOAuthSignal(t *testing.T, signal <-chan struct{}, done <-chan oauthLoginResult, name string) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-signal:
|
||||
case result := <-done:
|
||||
t.Fatalf("OAuth Login returned before %s: token=%#v err=%v", name, result.token, result.err)
|
||||
case <-time.After(oauthTestWaitTimeout):
|
||||
t.Fatalf("timed out waiting for OAuth %s", name)
|
||||
}
|
||||
}
|
||||
|
||||
func closeOAuthRelease(ch chan struct{}) {
|
||||
select {
|
||||
case <-ch:
|
||||
default:
|
||||
close(ch)
|
||||
}
|
||||
}
|
||||
|
||||
func newOAuthLoginFixture(t *testing.T, status func(int32) CLIAuthStatus) *oauthLoginFixture {
|
||||
t.Helper()
|
||||
isolateOAuthPersistence(t)
|
||||
SetClientID("")
|
||||
SetClientSecret("")
|
||||
resetClientIDFromMCP()
|
||||
@@ -91,6 +164,10 @@ func newOAuthLoginFixture(t *testing.T, status func(int32) CLIAuthStatus) *oauth
|
||||
}
|
||||
}))
|
||||
t.Cleanup(f.server.Close)
|
||||
t.Cleanup(func() {
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
})
|
||||
f.configDir = setupMCPConfigDir(t, f.server.URL)
|
||||
|
||||
oldClient := oauthHTTPClient
|
||||
@@ -118,17 +195,25 @@ func newOAuthLoginFixture(t *testing.T, status func(int32) CLIAuthStatus) *oauth
|
||||
|
||||
func httpGetBody(t *testing.T, rawURL string) (int, string) {
|
||||
t.Helper()
|
||||
result := getHTTPBody(rawURL)
|
||||
if result.err != nil {
|
||||
t.Fatalf("GET %s: %v", rawURL, result.err)
|
||||
}
|
||||
return result.status, result.body
|
||||
}
|
||||
|
||||
func getHTTPBody(rawURL string) oauthHTTPResult {
|
||||
client := &http.Client{Timeout: 2 * time.Second}
|
||||
resp, err := client.Get(rawURL)
|
||||
if err != nil {
|
||||
t.Fatalf("GET %s: %v", rawURL, err)
|
||||
return oauthHTTPResult{err: err}
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
data, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
return oauthHTTPResult{err: err}
|
||||
}
|
||||
return resp.StatusCode, string(data)
|
||||
return oauthHTTPResult{status: resp.StatusCode, body: string(data)}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthLoginCallbackAndAPIs(t *testing.T) {
|
||||
@@ -136,17 +221,7 @@ func TestCrossPlatformCoverageOAuthLoginCallbackAndAPIs(t *testing.T) {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: true}}
|
||||
})
|
||||
|
||||
loginDone := make(chan struct {
|
||||
token *TokenData
|
||||
err error
|
||||
}, 1)
|
||||
go func() {
|
||||
token, err := f.provider.Login(context.Background(), true)
|
||||
loginDone <- struct {
|
||||
token *TokenData
|
||||
err error
|
||||
}{token, err}
|
||||
}()
|
||||
loginDone := startOAuthLogin(t, context.Background(), f)
|
||||
|
||||
for _, path := range []string{"/api/superAdmin", "/api/sendApply?adminStaffId=admin-1", "/api/cliAuthEnabled"} {
|
||||
_, body := httpGetBody(t, f.callbackBase+path)
|
||||
@@ -166,18 +241,17 @@ func TestCrossPlatformCoverageOAuthLoginCallbackAndAPIs(t *testing.T) {
|
||||
t.Fatalf("success page = %q", body)
|
||||
}
|
||||
|
||||
callbackDone := make(chan string, 1)
|
||||
callbackDone := make(chan oauthHTTPResult, 1)
|
||||
go func() {
|
||||
_, callbackBody := httpGetBody(t, f.callbackBase+CallbackPath+"?code=good")
|
||||
callbackDone <- callbackBody
|
||||
callbackDone <- getHTTPBody(f.callbackBase + CallbackPath + "?code=good")
|
||||
}()
|
||||
<-f.exchangeEntered
|
||||
waitOAuthSignal(t, f.exchangeEntered, loginDone, "token exchange")
|
||||
_, body = httpGetBody(t, f.callbackBase+CallbackPath+"?authCode=good")
|
||||
if !strings.Contains(body, "正在处理授权") {
|
||||
t.Fatalf("concurrent callback = %q", body)
|
||||
}
|
||||
close(f.exchangeRelease)
|
||||
<-f.statusEntered
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
waitOAuthSignal(t, f.statusEntered, loginDone, "CLI auth status check")
|
||||
|
||||
_, body = httpGetBody(t, f.callbackBase+CallbackPath+"?code=good")
|
||||
if !strings.Contains(body, "<html") {
|
||||
@@ -204,11 +278,16 @@ func TestCrossPlatformCoverageOAuthLoginCallbackAndAPIs(t *testing.T) {
|
||||
t.Fatalf("auth enabled API = %q", body)
|
||||
}
|
||||
|
||||
close(f.statusRelease)
|
||||
if callbackBody := <-callbackDone; !strings.Contains(callbackBody, "<html") {
|
||||
t.Fatalf("callback body = %q", callbackBody)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
select {
|
||||
case callback := <-callbackDone:
|
||||
if callback.err != nil || !strings.Contains(callback.body, "<html") {
|
||||
t.Fatalf("callback body = %q, %v", callback.body, callback.err)
|
||||
}
|
||||
case <-time.After(oauthTestWaitTimeout):
|
||||
t.Fatal("timed out waiting for OAuth callback")
|
||||
}
|
||||
result := <-loginDone
|
||||
result := awaitOAuthLogin(t, loginDone)
|
||||
if result.err != nil || result.token == nil || result.token.AccessToken != "access" {
|
||||
t.Fatalf("Login = %#v, %v", result.token, result.err)
|
||||
}
|
||||
@@ -228,23 +307,20 @@ func TestCrossPlatformCoverageOAuthLoginMissingCallbackCode(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: true}}
|
||||
})
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
status, body := httpGetBody(t, f.callbackBase+CallbackPath)
|
||||
if status != http.StatusBadRequest || strings.TrimSpace(body) == "" {
|
||||
t.Fatalf("missing callback = %d %q", status, body)
|
||||
}
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("missing callback code did not fail login")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthLoginEarlyAndListenerEdges(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
var buf bytes.Buffer
|
||||
p := &OAuthProvider{configDir: t.TempDir(), Output: &buf}
|
||||
if p.output() != &buf || (*OAuthProvider)(nil).output() != io.Discard {
|
||||
@@ -311,6 +387,7 @@ func TestCrossPlatformCoverageOAuthLoginTimeoutAndServerError(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderOtherMethods(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
dir := t.TempDir()
|
||||
_ = DeleteTokenDataKeychain()
|
||||
t.Cleanup(func() { _ = DeleteTokenDataKeychain() })
|
||||
@@ -363,6 +440,7 @@ func (f roundTripFunc) RoundTrip(r *http.Request) (*http.Response, error) {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthPersistConfigEdges(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
p := &OAuthProvider{configDir: t.TempDir(), logger: slog.Default()}
|
||||
SetClientID("")
|
||||
SetClientSecret("")
|
||||
@@ -425,16 +503,12 @@ func TestCrossPlatformCoverageOAuthLoginDenialAndPolling(t *testing.T) {
|
||||
for _, tt := range terminal {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return tt.status })
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
f.provider.NoBrowser = false
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
_, _ = httpGetBody(t, f.callbackBase+CallbackPath+"?code=denied")
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("denied login succeeded")
|
||||
}
|
||||
})
|
||||
@@ -445,15 +519,11 @@ func TestCrossPlatformCoverageOAuthLoginDenialAndPolling(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{ChannelScope: "specified", AllowedChannels: []string{"allowed"}}}
|
||||
})
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
_, _ = httpGetBody(t, f.callbackBase+CallbackPath+"?code=denied")
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("channel-denied login succeeded")
|
||||
}
|
||||
})
|
||||
@@ -463,15 +533,11 @@ func TestCrossPlatformCoverageOAuthLoginDenialAndPolling(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{}}
|
||||
})
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
_, _ = httpGetBody(t, f.callbackBase+CallbackPath+"?code=pending")
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("approval timeout login succeeded")
|
||||
}
|
||||
oauthApprovalTimeout = oldApproval
|
||||
@@ -481,16 +547,12 @@ func TestCrossPlatformCoverageOAuthLoginDenialAndPolling(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(call int32) CLIAuthStatus {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: call > 1}}
|
||||
})
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
_, _ = httpGetBody(t, f.callbackBase+CallbackPath+"?code=pending")
|
||||
if err := <-done; err != nil {
|
||||
t.Fatalf("poll-enabled login failed: %v", err)
|
||||
if result := awaitOAuthLogin(t, done); result.err != nil {
|
||||
t.Fatalf("poll-enabled login failed: %v", result.err)
|
||||
}
|
||||
})
|
||||
|
||||
@@ -498,17 +560,13 @@ func TestCrossPlatformCoverageOAuthLoginDenialAndPolling(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{}}
|
||||
})
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(ctx, true)
|
||||
done <- err
|
||||
}()
|
||||
done := startOAuthLogin(t, ctx, f)
|
||||
_, _ = httpGetBody(t, f.callbackBase+CallbackPath+"?code=pending")
|
||||
cancel()
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("canceled pending login succeeded")
|
||||
}
|
||||
})
|
||||
@@ -519,23 +577,20 @@ func TestCrossPlatformCoverageOAuthLoginExchangeFailure(t *testing.T) {
|
||||
return CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: true}}
|
||||
})
|
||||
f.exchangeError.Store(true)
|
||||
close(f.exchangeRelease)
|
||||
close(f.statusRelease)
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := f.provider.Login(context.Background(), true)
|
||||
done <- err
|
||||
}()
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
closeOAuthRelease(f.statusRelease)
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
_, body := httpGetBody(t, f.callbackBase+CallbackPath+"?code=bad")
|
||||
if !strings.Contains(body, "failed") {
|
||||
t.Fatalf("exchange failure page = %q", body)
|
||||
}
|
||||
if err := <-done; err == nil {
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil {
|
||||
t.Fatal("exchange failure login succeeded")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthRefreshAndParsingEdges(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
t.Setenv("DWS_CLIENT_ID", "")
|
||||
t.Setenv("DWS_CLIENT_SECRET", "")
|
||||
dir := t.TempDir()
|
||||
@@ -640,7 +695,9 @@ func TestCrossPlatformCoverageOAuthRefreshAndParsingEdges(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderHighLevelEdges(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
oldLoad := oauthLoadToken
|
||||
oldLoadLocked := oauthLoadTokenLocked
|
||||
oldAcquire := oauthAcquireLock
|
||||
oldMark := oauthMarkProfile
|
||||
oldFetch := oauthFetchClientID
|
||||
@@ -651,6 +708,7 @@ func TestCrossPlatformCoverageOAuthProviderHighLevelEdges(t *testing.T) {
|
||||
oldSaveLocked := oauthSaveTokenLocked
|
||||
t.Cleanup(func() {
|
||||
oauthLoadToken = oldLoad
|
||||
oauthLoadTokenLocked = oldLoadLocked
|
||||
oauthAcquireLock = oldAcquire
|
||||
oauthMarkProfile = oldMark
|
||||
oauthFetchClientID = oldFetch
|
||||
@@ -665,6 +723,9 @@ func TestCrossPlatformCoverageOAuthProviderHighLevelEdges(t *testing.T) {
|
||||
})
|
||||
fail := errors.New("fail")
|
||||
p := &OAuthProvider{configDir: t.TempDir(), logger: slog.Default(), Output: io.Discard}
|
||||
oauthLoadTokenLocked = func(configDir, _ string) (*TokenData, error) {
|
||||
return oauthLoadToken(configDir)
|
||||
}
|
||||
|
||||
valid := &TokenData{AccessToken: "valid", ExpiresAt: time.Now().Add(time.Hour)}
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return valid, nil }
|
||||
@@ -795,29 +856,24 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
oauthSuccessPause = 0
|
||||
oauthSleep = func(time.Duration) {}
|
||||
|
||||
type loginResult struct {
|
||||
token *TokenData
|
||||
err error
|
||||
}
|
||||
startLogin := func(ctx context.Context, f *oauthLoginFixture) <-chan loginResult {
|
||||
finishExchange := func(t *testing.T, f *oauthLoginFixture, done <-chan oauthLoginResult, code string) string {
|
||||
t.Helper()
|
||||
done := make(chan loginResult, 1)
|
||||
bodyCh := make(chan oauthHTTPResult, 1)
|
||||
go func() {
|
||||
token, err := f.provider.Login(ctx, true)
|
||||
done <- loginResult{token: token, err: err}
|
||||
bodyCh <- getHTTPBody(f.callbackBase + CallbackPath + "?code=" + url.QueryEscape(code))
|
||||
}()
|
||||
return done
|
||||
}
|
||||
finishExchange := func(t *testing.T, f *oauthLoginFixture, code string) string {
|
||||
t.Helper()
|
||||
bodyCh := make(chan string, 1)
|
||||
go func() {
|
||||
_, body := httpGetBody(t, f.callbackBase+CallbackPath+"?code="+url.QueryEscape(code))
|
||||
bodyCh <- body
|
||||
}()
|
||||
<-f.exchangeEntered
|
||||
close(f.exchangeRelease)
|
||||
return <-bodyCh
|
||||
waitOAuthSignal(t, f.exchangeEntered, done, "token exchange")
|
||||
closeOAuthRelease(f.exchangeRelease)
|
||||
select {
|
||||
case result := <-bodyCh:
|
||||
if result.err != nil {
|
||||
t.Fatalf("OAuth callback failed: %v", result.err)
|
||||
}
|
||||
return result.body
|
||||
case <-time.After(oauthTestWaitTimeout):
|
||||
t.Fatal("timed out waiting for OAuth callback")
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
t.Run("switch organization and cached disabled pages", func(t *testing.T) {
|
||||
@@ -830,8 +886,8 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
return &CLIAuthStatus{Success: true, Result: &CLIAuthResult{CLIAuthEnabled: true}}, nil
|
||||
}
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return CLIAuthStatus{} })
|
||||
done := startLogin(context.Background(), f)
|
||||
if body := finishExchange(t, f, "first"); !strings.Contains(body, "<html") {
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
if body := finishExchange(t, f, done, "first"); !strings.Contains(body, "<html") {
|
||||
t.Fatalf("disabled callback body = %q", body)
|
||||
}
|
||||
if _, body := httpGetBody(t, f.callbackBase+CallbackPath+"?code=first"); !strings.Contains(body, "<html") {
|
||||
@@ -843,7 +899,7 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
if _, body := httpGetBody(t, f.callbackBase+CallbackPath+"?code=second"); !strings.Contains(body, "<html") {
|
||||
t.Fatalf("switched callback = %q", body)
|
||||
}
|
||||
result := <-done
|
||||
result := awaitOAuthLogin(t, done)
|
||||
if result.err != nil || result.token == nil || result.token.AccessToken != "access" {
|
||||
t.Fatalf("switched login = %#v %v", result.token, result.err)
|
||||
}
|
||||
@@ -857,15 +913,15 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
oauthSendApply = func(context.Context, string, string) (*SendApplyResponse, error) { return nil, fail }
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return CLIAuthStatus{} })
|
||||
done := startLogin(ctx, f)
|
||||
finishExchange(t, f, "errors")
|
||||
done := startOAuthLogin(t, ctx, f)
|
||||
finishExchange(t, f, done, "errors")
|
||||
for _, path := range []string{"/api/superAdmin", "/api/sendApply?adminStaffId=admin", "/api/cliAuthEnabled"} {
|
||||
if _, body := httpGetBody(t, f.callbackBase+path); !strings.Contains(body, "hook failure") {
|
||||
t.Fatalf("API error %s = %q", path, body)
|
||||
}
|
||||
}
|
||||
cancel()
|
||||
if result := <-done; !errors.Is(result.err, context.Canceled) {
|
||||
if result := awaitOAuthLogin(t, done); !errors.Is(result.err, context.Canceled) {
|
||||
t.Fatalf("canceled error login = %v", result.err)
|
||||
}
|
||||
})
|
||||
@@ -882,14 +938,14 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return CLIAuthStatus{} })
|
||||
var output bytes.Buffer
|
||||
f.provider.Output = &output
|
||||
done := startLogin(ctx, f)
|
||||
finishExchange(t, f, "apply")
|
||||
done := startOAuthLogin(t, ctx, f)
|
||||
finishExchange(t, f, done, "apply")
|
||||
if _, body := httpGetBody(t, f.callbackBase+"/api/sendApply?adminStaffId=admin"); !strings.Contains(body, "true") {
|
||||
t.Fatalf("apply response = %q", body)
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
cancel()
|
||||
if result := <-done; !errors.Is(result.err, context.Canceled) {
|
||||
if result := awaitOAuthLogin(t, done); !errors.Is(result.err, context.Canceled) {
|
||||
t.Fatalf("canceled apply login = %v", result.err)
|
||||
}
|
||||
if !strings.Contains(output.String(), "Waiting for admin approval") && !strings.Contains(output.String(), "等待管理员审批中") {
|
||||
@@ -903,9 +959,9 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
return &CLIAuthStatus{Success: false, ErrorCode: "ENTERPRISE_NOT_AUTHORIZED"}, nil
|
||||
}
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return CLIAuthStatus{} })
|
||||
done := startLogin(context.Background(), f)
|
||||
finishExchange(t, f, "enterprise")
|
||||
if result := <-done; result.err == nil || !strings.Contains(result.err.Error(), "企业安全认证") {
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
finishExchange(t, f, done, "enterprise")
|
||||
if result := awaitOAuthLogin(t, done); result.err == nil || !strings.Contains(result.err.Error(), "企业安全认证") {
|
||||
t.Fatalf("enterprise denial = %v", result.err)
|
||||
}
|
||||
})
|
||||
@@ -928,15 +984,16 @@ func TestCrossPlatformCoverageOAuthCallbackRemainingEdges(t *testing.T) {
|
||||
fail := errors.New("save failure")
|
||||
oauthSaveToken = func(string, *TokenData) error { return fail }
|
||||
f := newOAuthLoginFixture(t, func(int32) CLIAuthStatus { return CLIAuthStatus{} })
|
||||
done := startLogin(context.Background(), f)
|
||||
finishExchange(t, f, "save")
|
||||
if result := <-done; !errors.Is(result.err, fail) {
|
||||
done := startOAuthLogin(t, context.Background(), f)
|
||||
finishExchange(t, f, done, "save")
|
||||
if result := awaitOAuthLogin(t, done); !errors.Is(result.err, fail) {
|
||||
t.Fatalf("save failure login = %v", result.err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthHelperRemainingEdges(t *testing.T) {
|
||||
isolateOAuthPersistence(t)
|
||||
oldClient := oauthHTTPClient
|
||||
oldRequest := oauthNewRequest
|
||||
oldRetry := oauthRetryAfter
|
||||
|
||||
@@ -288,7 +288,7 @@ func TestPortableAuthBundleRoundTripPreservesProfiles(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() after import error = %v", err)
|
||||
}
|
||||
if cfg.PrimaryProfile != "corp_a" || cfg.CurrentProfile != "corp_b" || cfg.PreviousProfile != "corp_a" {
|
||||
if cfg.PrimaryProfile != "" || cfg.CurrentProfile != "corp_b" || cfg.PreviousProfile != "corp_a" {
|
||||
t.Fatalf("profiles after import = %#v", cfg)
|
||||
}
|
||||
if len(cfg.Profiles) != 2 {
|
||||
@@ -311,6 +311,65 @@ func TestPortableAuthBundleRoundTripPreservesProfiles(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestPortableAuthBundleRoundTripPreservesSameCorpAccounts(t *testing.T) {
|
||||
requirePortableFileBackend(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
SetRuntimeProfile("")
|
||||
t.Cleanup(func() { SetRuntimeProfile("") })
|
||||
|
||||
sourceKeychain := filepath.Join(t.TempDir(), "source-keychain")
|
||||
t.Setenv(keychain.StorageDirEnv, sourceKeychain)
|
||||
sourceConfig := filepath.Join(t.TempDir(), ".dws")
|
||||
|
||||
first := &TokenData{
|
||||
AccessToken: "access-first",
|
||||
RefreshToken: "refresh-first",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(30 * 24 * time.Hour),
|
||||
CorpID: "corp_same",
|
||||
CorpName: "Same Org",
|
||||
UserID: "user_1",
|
||||
UserName: "账号一",
|
||||
}
|
||||
second := *first
|
||||
second.AccessToken = "access-second"
|
||||
second.RefreshToken = "refresh-second"
|
||||
second.UserID = "user_2"
|
||||
second.UserName = "账号二"
|
||||
if err := SaveTokenData(sourceConfig, first); err != nil {
|
||||
t.Fatalf("SaveTokenData(first) error = %v", err)
|
||||
}
|
||||
if err := SaveTokenData(sourceConfig, &second); err != nil {
|
||||
t.Fatalf("SaveTokenData(second) error = %v", err)
|
||||
}
|
||||
|
||||
var bundle bytes.Buffer
|
||||
if err := ExportPortableAuthBundle(sourceConfig, &bundle); err != nil {
|
||||
t.Fatalf("ExportPortableAuthBundle() error = %v", err)
|
||||
}
|
||||
|
||||
targetKeychain := filepath.Join(t.TempDir(), "target-keychain")
|
||||
t.Setenv(keychain.StorageDirEnv, targetKeychain)
|
||||
targetConfig := filepath.Join(t.TempDir(), ".dws")
|
||||
if _, err := ImportPortableAuthBundle(targetConfig, bytes.NewReader(bundle.Bytes())); err != nil {
|
||||
t.Fatalf("ImportPortableAuthBundle() error = %v", err)
|
||||
}
|
||||
|
||||
for selector, wantToken := range map[string]string{
|
||||
"corp_same:user_1": "access-first",
|
||||
"corp_same:user_2": "access-second",
|
||||
"corp_same": "access-second",
|
||||
} {
|
||||
loaded, err := LoadTokenDataForProfile(targetConfig, selector)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenDataForProfile(%s) error = %v", selector, err)
|
||||
}
|
||||
if loaded.AccessToken != wantToken {
|
||||
t.Fatalf("profile %s token = %q, want %q", selector, loaded.AccessToken, wantToken)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func requirePortableFileBackend(t *testing.T) {
|
||||
t.Helper()
|
||||
if runtime.GOOS == "windows" {
|
||||
|
||||
+1011
-208
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,89 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// RefreshFailureClass separates refresh failures that may recover after a
|
||||
// delay from failures that require new credentials or local intervention.
|
||||
type RefreshFailureClass string
|
||||
|
||||
const (
|
||||
RefreshFailureUnknown RefreshFailureClass = "unknown"
|
||||
RefreshFailureTransient RefreshFailureClass = "transient"
|
||||
RefreshFailureTerminal RefreshFailureClass = "terminal"
|
||||
)
|
||||
|
||||
// HTTPStatusError preserves an OAuth endpoint status for structured retry
|
||||
// decisions without copying an untrusted response body into logs.
|
||||
type HTTPStatusError struct {
|
||||
StatusCode int
|
||||
responseBody string
|
||||
}
|
||||
|
||||
func (e *HTTPStatusError) Error() string {
|
||||
if e == nil {
|
||||
return "OAuth endpoint request failed"
|
||||
}
|
||||
return fmt.Sprintf("HTTP %d", e.StatusCode)
|
||||
}
|
||||
|
||||
func httpStatusResponseBody(err error) string {
|
||||
var statusErr *HTTPStatusError
|
||||
if !errors.As(err, &statusErr) || statusErr == nil {
|
||||
return ""
|
||||
}
|
||||
return statusErr.responseBody
|
||||
}
|
||||
|
||||
// ClassifyRefreshFailure uses only structured transport and HTTP signals.
|
||||
// Unknown errors, including parse, keychain and persistence failures, remain
|
||||
// fatal so a long-running source cannot retry an error that needs user action.
|
||||
func ClassifyRefreshFailure(err error) RefreshFailureClass {
|
||||
if err == nil {
|
||||
return RefreshFailureUnknown
|
||||
}
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return RefreshFailureTransient
|
||||
}
|
||||
if errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
return RefreshFailureTransient
|
||||
}
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) {
|
||||
return RefreshFailureTransient
|
||||
}
|
||||
var statusErr *HTTPStatusError
|
||||
if !errors.As(err, &statusErr) || statusErr == nil {
|
||||
return RefreshFailureUnknown
|
||||
}
|
||||
if statusErr.StatusCode == http.StatusRequestTimeout ||
|
||||
statusErr.StatusCode == http.StatusTooManyRequests ||
|
||||
statusErr.StatusCode >= http.StatusInternalServerError {
|
||||
return RefreshFailureTransient
|
||||
}
|
||||
if statusErr.StatusCode == http.StatusBadRequest ||
|
||||
statusErr.StatusCode == http.StatusUnauthorized ||
|
||||
statusErr.StatusCode == http.StatusForbidden {
|
||||
return RefreshFailureTerminal
|
||||
}
|
||||
return RefreshFailureUnknown
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageClassifyRefreshFailureUsesStructuredSignals(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
err error
|
||||
want RefreshFailureClass
|
||||
}{
|
||||
{name: "deadline", err: context.DeadlineExceeded, want: RefreshFailureTransient},
|
||||
{name: "network", err: &url.Error{Op: "Post", URL: "https://oauth.test", Err: context.DeadlineExceeded}, want: RefreshFailureTransient},
|
||||
{name: "request timeout", err: &HTTPStatusError{StatusCode: http.StatusRequestTimeout}, want: RefreshFailureTransient},
|
||||
{name: "rate limited", err: &HTTPStatusError{StatusCode: http.StatusTooManyRequests}, want: RefreshFailureTransient},
|
||||
{name: "server unavailable", err: &HTTPStatusError{StatusCode: http.StatusServiceUnavailable}, want: RefreshFailureTransient},
|
||||
{name: "refresh rejected", err: &HTTPStatusError{StatusCode: http.StatusUnauthorized}, want: RefreshFailureTerminal},
|
||||
{name: "invalid grant", err: &HTTPStatusError{StatusCode: http.StatusBadRequest}, want: RefreshFailureTerminal},
|
||||
{name: "forbidden", err: &HTTPStatusError{StatusCode: http.StatusForbidden}, want: RefreshFailureTerminal},
|
||||
{name: "local persistence", err: errors.New("save refreshed token failed"), want: RefreshFailureUnknown},
|
||||
{name: "nil error", err: nil, want: RefreshFailureUnknown},
|
||||
{name: "dns failure", err: &net.DNSError{Err: "no such host", Name: "oauth.test"}, want: RefreshFailureTransient},
|
||||
{name: "redirect status", err: &HTTPStatusError{StatusCode: http.StatusFound}, want: RefreshFailureUnknown},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := ClassifyRefreshFailure(tt.err); got != tt.want {
|
||||
t.Fatalf("ClassifyRefreshFailure() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageHTTPStatusErrorRetainsStatusThroughWrapping(t *testing.T) {
|
||||
want := &HTTPStatusError{StatusCode: http.StatusTooManyRequests}
|
||||
err := errors.Join(errors.New("refresh failed"), want)
|
||||
if got := ClassifyRefreshFailure(err); got != RefreshFailureTransient {
|
||||
t.Fatalf("ClassifyRefreshFailure() = %q, want transient", got)
|
||||
}
|
||||
var statusErr *HTTPStatusError
|
||||
if !errors.As(err, &statusErr) || statusErr.StatusCode != http.StatusTooManyRequests {
|
||||
t.Fatalf("HTTP status error not retained: %v", err)
|
||||
}
|
||||
if got, want := statusErr.Error(), "HTTP 429"; got != want {
|
||||
t.Fatalf("HTTP status error = %q, want %q", got, want)
|
||||
}
|
||||
var nilStatus *HTTPStatusError
|
||||
if got, want := nilStatus.Error(), "OAuth endpoint request failed"; got != want {
|
||||
t.Fatalf("nil HTTP status error = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthExchangeDisplayErrorFallsBackToPlainError(t *testing.T) {
|
||||
if got, want := oauthExchangeDisplayError(&HTTPStatusError{StatusCode: http.StatusBadGateway}), "HTTP 502: token exchange failed"; got != want {
|
||||
t.Fatalf("status display error = %q, want %q", got, want)
|
||||
}
|
||||
if got, want := oauthExchangeDisplayError(errors.New("exchange failed")), "exchange failed"; got != want {
|
||||
t.Fatalf("plain display error = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePostJSONClassifiesStatusWithoutLoggingResponseBody(t *testing.T) {
|
||||
const secretBody = `{"refreshToken":"must-not-reach-logs"}`
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
_, _ = w.Write([]byte(secretBody))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
provider := &OAuthProvider{httpClient: server.Client()}
|
||||
_, err := provider.postJSON(context.Background(), server.URL, map[string]string{"grantType": "refresh_token"})
|
||||
if got := ClassifyRefreshFailure(err); got != RefreshFailureTransient {
|
||||
t.Fatalf("ClassifyRefreshFailure() = %q, want transient: %v", got, err)
|
||||
}
|
||||
if strings.Contains(err.Error(), "must-not-reach-logs") {
|
||||
t.Fatalf("postJSON error leaked response body: %v", err)
|
||||
}
|
||||
if got := httpStatusResponseBody(err); !strings.Contains(got, "must-not-reach-logs") {
|
||||
t.Fatalf("postJSON did not retain bounded response details for internal classification: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageGetTokenSnapshotOnlyExpiresProfileForNonTransientRefreshFailures(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
oldLoadLocked := oauthLoadTokenLocked
|
||||
oldAcquire := oauthAcquireLock
|
||||
oldRefresh := oauthRefreshToken
|
||||
oldMark := oauthMarkProfile
|
||||
oldEdition := edition.Get()
|
||||
t.Cleanup(func() {
|
||||
oauthLoadToken = oldLoad
|
||||
oauthLoadTokenLocked = oldLoadLocked
|
||||
oauthAcquireLock = oldAcquire
|
||||
oauthRefreshToken = oldRefresh
|
||||
oauthMarkProfile = oldMark
|
||||
edition.Override(oldEdition)
|
||||
})
|
||||
edition.Override(&edition.Hooks{})
|
||||
|
||||
expired := &TokenData{
|
||||
AccessToken: "expired-access",
|
||||
ExpiresAt: time.Now().Add(-time.Hour),
|
||||
RefreshToken: "refresh",
|
||||
RefreshExpAt: time.Now().Add(time.Hour),
|
||||
CorpID: "corp",
|
||||
UserID: "user",
|
||||
}
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return expired, nil }
|
||||
oauthLoadTokenLocked = func(string, string) (*TokenData, error) { return expired, nil }
|
||||
oauthAcquireLock = func(context.Context, string) (*DualLock, error) { return &DualLock{}, nil }
|
||||
|
||||
markCalls := 0
|
||||
oauthMarkProfile = func(_, _, status string) error {
|
||||
if status != ProfileStatusExpired {
|
||||
t.Fatalf("profile status = %q, want %q", status, ProfileStatusExpired)
|
||||
}
|
||||
markCalls++
|
||||
return nil
|
||||
}
|
||||
provider := NewOAuthProvider(t.TempDir(), nil)
|
||||
|
||||
oauthRefreshToken = func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
|
||||
return nil, &HTTPStatusError{StatusCode: http.StatusServiceUnavailable}
|
||||
}
|
||||
if _, err := provider.GetTokenSnapshot(context.Background()); ClassifyRefreshFailure(err) != RefreshFailureTransient {
|
||||
t.Fatalf("transient refresh error = %v", err)
|
||||
}
|
||||
if markCalls != 0 {
|
||||
t.Fatalf("transient refresh marked profile expired %d times", markCalls)
|
||||
}
|
||||
|
||||
oauthRefreshToken = func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
|
||||
return nil, &HTTPStatusError{StatusCode: http.StatusUnauthorized}
|
||||
}
|
||||
if _, err := provider.GetTokenSnapshot(context.Background()); ClassifyRefreshFailure(err) != RefreshFailureTerminal {
|
||||
t.Fatalf("terminal refresh error = %v", err)
|
||||
}
|
||||
if markCalls != 1 {
|
||||
t.Fatalf("terminal refresh marked profile expired %d times, want 1", markCalls)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,510 @@
|
||||
// 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"
|
||||
"errors"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
type rejectedTokenHookStore struct {
|
||||
mu sync.Mutex
|
||||
data TokenData
|
||||
deletes int
|
||||
}
|
||||
|
||||
func (s *rejectedTokenHookStore) load(string) ([]byte, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return json.Marshal(s.data)
|
||||
}
|
||||
|
||||
func (s *rejectedTokenHookStore) save(_ string, blob []byte) error {
|
||||
var data TokenData
|
||||
if err := json.Unmarshal(blob, &data); err != nil {
|
||||
return err
|
||||
}
|
||||
s.mu.Lock()
|
||||
s.data = data
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *rejectedTokenHookStore) delete(string) error {
|
||||
s.mu.Lock()
|
||||
s.data = TokenData{}
|
||||
s.deletes++
|
||||
s.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *rejectedTokenHookStore) snapshot() (TokenData, int) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
return s.data, s.deletes
|
||||
}
|
||||
|
||||
func installRejectedTokenHookStore(t *testing.T, data TokenData) *rejectedTokenHookStore {
|
||||
t.Helper()
|
||||
store := &rejectedTokenHookStore{data: data}
|
||||
previousHooks := edition.Get()
|
||||
edition.Override(&edition.Hooks{
|
||||
LoadToken: store.load,
|
||||
SaveToken: store.save,
|
||||
DeleteToken: store.delete,
|
||||
})
|
||||
t.Cleanup(func() { edition.Override(previousHooks) })
|
||||
return store
|
||||
}
|
||||
|
||||
func installOAuthRefreshStub(t *testing.T, fn func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error)) {
|
||||
t.Helper()
|
||||
resetRejectedTokenRefreshCoordinator(t)
|
||||
previous := oauthRefreshToken
|
||||
oauthRefreshToken = fn
|
||||
t.Cleanup(func() { oauthRefreshToken = previous })
|
||||
}
|
||||
|
||||
func resetRejectedTokenRefreshCoordinator(t *testing.T) {
|
||||
t.Helper()
|
||||
reset := func() {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
rejectedTokenRefreshCoordinator.inFlight = make(map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall)
|
||||
rejectedTokenRefreshCoordinator.failures = make(map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure)
|
||||
rejectedTokenRefreshCoordinator.now = time.Now
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
}
|
||||
reset()
|
||||
t.Cleanup(reset)
|
||||
}
|
||||
|
||||
func waitForRejectedTokenRefreshParticipants(t *testing.T, key rejectedTokenRefreshKey, want int) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for {
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
call := rejectedTokenRefreshCoordinator.inFlight[key]
|
||||
got := 0
|
||||
if call != nil {
|
||||
got = call.participants
|
||||
}
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
if got >= want {
|
||||
return
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("refresh participants = %d, want %d", got, want)
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
}
|
||||
|
||||
func installProfilesAcquireProbe(t *testing.T) <-chan struct{} {
|
||||
t.Helper()
|
||||
previous := profilesAcquireDualLock
|
||||
attempted := make(chan struct{}, 1)
|
||||
profilesAcquireDualLock = func(ctx context.Context, configDir string) (*DualLock, error) {
|
||||
attempted <- struct{}{}
|
||||
return previous(ctx, configDir)
|
||||
}
|
||||
t.Cleanup(func() { profilesAcquireDualLock = previous })
|
||||
return attempted
|
||||
}
|
||||
|
||||
func waitForProfilesAcquire(t *testing.T, attempted <-chan struct{}) {
|
||||
t.Helper()
|
||||
select {
|
||||
case <-attempted:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("public opaque token mutation did not enter the Core dual lock")
|
||||
}
|
||||
}
|
||||
|
||||
func validRejectedTokenData(accessToken string) TokenData {
|
||||
return TokenData{
|
||||
AccessToken: accessToken,
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
Source: "mcp",
|
||||
ClientID: "client-id",
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageForceRefreshRejectedTokenConcurrentCallersExchangeOnce(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
var refreshCalls atomic.Int32
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
|
||||
if refreshCalls.Add(1) == 1 {
|
||||
close(started)
|
||||
}
|
||||
<-release
|
||||
updated := *data
|
||||
updated.AccessToken = "new-access"
|
||||
updated.ExpiresAt = time.Now().Add(time.Hour)
|
||||
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &updated, nil
|
||||
})
|
||||
|
||||
provider := NewOAuthProvider(t.TempDir(), nil)
|
||||
const workers = 8
|
||||
results := make(chan string, workers)
|
||||
errs := make(chan error, workers)
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
results <- token
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
<-started
|
||||
close(release)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
close(errs)
|
||||
|
||||
for err := range errs {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
for token := range results {
|
||||
if token != "new-access" {
|
||||
t.Fatalf("token = %q, want new-access", token)
|
||||
}
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("refresh calls = %d, want 1", got)
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "new-access" || deletes != 0 {
|
||||
t.Fatalf("stored token = %q, deletes = %d", stored.AccessToken, deletes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageForceRefreshRejectedTokenFailurePreservesCredential(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
refreshErr := errors.New("temporary refresh failure")
|
||||
installOAuthRefreshStub(t, func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
|
||||
return nil, refreshErr
|
||||
})
|
||||
|
||||
_, err := NewOAuthProvider(t.TempDir(), nil).ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
if !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("error = %v, want refresh cause", err)
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "old-access" || stored.RefreshToken != "refresh-token" || deletes != 0 {
|
||||
t.Fatalf("credential changed after transient failure: %#v, deletes=%d", stored, deletes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageForceRefreshRejectedTokenFailureIsSingleflightAndCooledDown(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
previousProfile := RuntimeProfile()
|
||||
SetRuntimeProfile("")
|
||||
t.Cleanup(func() { SetRuntimeProfile(previousProfile) })
|
||||
|
||||
refreshErr := errors.New("temporary refresh failure")
|
||||
var refreshCalls atomic.Int32
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
baseNow := time.Now()
|
||||
var nowNanos atomic.Int64
|
||||
nowNanos.Store(baseNow.UnixNano())
|
||||
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
|
||||
call := refreshCalls.Add(1)
|
||||
if call == 1 {
|
||||
close(started)
|
||||
<-release
|
||||
return nil, refreshErr
|
||||
}
|
||||
updated := *data
|
||||
updated.AccessToken = "recovered-access"
|
||||
updated.ExpiresAt = time.Now().Add(time.Hour)
|
||||
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &updated, nil
|
||||
})
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
rejectedTokenRefreshCoordinator.now = func() time.Time {
|
||||
return time.Unix(0, nowNanos.Load())
|
||||
}
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
|
||||
configDir := t.TempDir()
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
const workers = 8
|
||||
start := make(chan struct{})
|
||||
ready := make(chan struct{}, workers)
|
||||
errs := make(chan error, workers)
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
ready <- struct{}{}
|
||||
<-start
|
||||
_, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
for range workers {
|
||||
<-ready
|
||||
}
|
||||
close(start)
|
||||
<-started
|
||||
key := newRejectedTokenRefreshKey(configDir, "", "old-access")
|
||||
waitForRejectedTokenRefreshParticipants(t, key, workers)
|
||||
close(release)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
|
||||
for err := range errs {
|
||||
if !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("shared refresh error = %v, want %v", err, refreshErr)
|
||||
}
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("refresh calls after concurrent failure = %d, want 1", got)
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "old-access" || stored.RefreshToken != "refresh-token" || deletes != 0 {
|
||||
t.Fatalf("credential changed after shared failure: %#v, deletes=%d", stored, deletes)
|
||||
}
|
||||
|
||||
if _, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access"); !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("cooldown error = %v, want %v", err, refreshErr)
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("refresh calls inside cooldown = %d, want 1", got)
|
||||
}
|
||||
|
||||
nowNanos.Store(baseNow.Add(rejectedTokenRefreshFailureCooldown + time.Nanosecond).UnixNano())
|
||||
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
if err != nil || token != "recovered-access" {
|
||||
t.Fatalf("refresh after cooldown = %q, %v", token, err)
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 2 {
|
||||
t.Fatalf("refresh calls after cooldown = %d, want 2", got)
|
||||
}
|
||||
stored, deletes = store.snapshot()
|
||||
if stored.AccessToken != "recovered-access" || deletes != 0 {
|
||||
t.Fatalf("stored token after cooldown recovery = %q, deletes=%d", stored.AccessToken, deletes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageForceRefreshRejectedTokenChangedDuringCooldownUsesNewToken(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
previousProfile := RuntimeProfile()
|
||||
SetRuntimeProfile("")
|
||||
t.Cleanup(func() { SetRuntimeProfile(previousProfile) })
|
||||
|
||||
refreshErr := errors.New("temporary refresh failure")
|
||||
var refreshCalls atomic.Int32
|
||||
baseNow := time.Now()
|
||||
installOAuthRefreshStub(t, func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
|
||||
refreshCalls.Add(1)
|
||||
return nil, refreshErr
|
||||
})
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
rejectedTokenRefreshCoordinator.now = func() time.Time { return baseNow }
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
|
||||
configDir := t.TempDir()
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
if _, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access"); !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("initial refresh error = %v, want %v", err, refreshErr)
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("initial refresh calls = %d, want 1", got)
|
||||
}
|
||||
|
||||
store.mu.Lock()
|
||||
store.data = validRejectedTokenData("externally-refreshed")
|
||||
store.mu.Unlock()
|
||||
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
if err != nil || token != "externally-refreshed" {
|
||||
t.Fatalf("refresh after external publication = %q, %v", token, err)
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("external publication triggered another exchange: calls=%d", got)
|
||||
}
|
||||
key := newRejectedTokenRefreshKey(configDir, "", "old-access")
|
||||
rejectedTokenRefreshCoordinator.Lock()
|
||||
_, failurePresent := rejectedTokenRefreshCoordinator.failures[key]
|
||||
rejectedTokenRefreshCoordinator.Unlock()
|
||||
if failurePresent {
|
||||
t.Fatal("old-token failure cache was not cleared after external publication")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOpaquePublisherWaitsForRejectedTokenRefresh(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
var releaseOnce sync.Once
|
||||
t.Cleanup(func() { releaseOnce.Do(func() { close(release) }) })
|
||||
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
|
||||
close(started)
|
||||
<-release
|
||||
updated := *data
|
||||
updated.AccessToken = "refreshed-from-old"
|
||||
updated.ExpiresAt = time.Now().Add(time.Hour)
|
||||
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &updated, nil
|
||||
})
|
||||
|
||||
configDir := t.TempDir()
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
refreshResult := make(chan struct {
|
||||
token string
|
||||
err error
|
||||
}, 1)
|
||||
go func() {
|
||||
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
refreshResult <- struct {
|
||||
token string
|
||||
err error
|
||||
}{token: token, err: err}
|
||||
}()
|
||||
<-started
|
||||
|
||||
acquireAttempted := installProfilesAcquireProbe(t)
|
||||
publishResult := make(chan error, 1)
|
||||
go func() {
|
||||
publishResult <- SaveTokenData(configDir, ptrTokenData(validRejectedTokenData("login-published")))
|
||||
}()
|
||||
waitForProfilesAcquire(t, acquireAttempted)
|
||||
releaseOnce.Do(func() { close(release) })
|
||||
|
||||
refresh := <-refreshResult
|
||||
if refresh.err != nil || refresh.token != "refreshed-from-old" {
|
||||
t.Fatalf("refresh result = %q, %v", refresh.token, refresh.err)
|
||||
}
|
||||
if err := <-publishResult; err != nil {
|
||||
t.Fatalf("publish token: %v", err)
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "login-published" || deletes != 0 {
|
||||
t.Fatalf("older refresh overwrote login publication: token=%q deletes=%d", stored.AccessToken, deletes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOpaqueLogoutWaitsForRejectedTokenRefresh(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
logout func(string) error
|
||||
}{
|
||||
{name: "current profile", logout: func(configDir string) error {
|
||||
return DeleteTokenDataForProfile(configDir, "")
|
||||
}},
|
||||
{name: "all profiles", logout: DeleteAllTokenData},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
|
||||
started := make(chan struct{})
|
||||
release := make(chan struct{})
|
||||
var releaseOnce sync.Once
|
||||
t.Cleanup(func() { releaseOnce.Do(func() { close(release) }) })
|
||||
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
|
||||
close(started)
|
||||
<-release
|
||||
updated := *data
|
||||
updated.AccessToken = "refreshed-before-logout"
|
||||
updated.ExpiresAt = time.Now().Add(time.Hour)
|
||||
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &updated, nil
|
||||
})
|
||||
|
||||
configDir := t.TempDir()
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
refreshResult := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
|
||||
refreshResult <- err
|
||||
}()
|
||||
<-started
|
||||
|
||||
acquireAttempted := installProfilesAcquireProbe(t)
|
||||
logoutResult := make(chan error, 1)
|
||||
go func() { logoutResult <- tc.logout(configDir) }()
|
||||
waitForProfilesAcquire(t, acquireAttempted)
|
||||
releaseOnce.Do(func() { close(release) })
|
||||
|
||||
if err := <-refreshResult; err != nil {
|
||||
t.Fatalf("refresh: %v", err)
|
||||
}
|
||||
if err := <-logoutResult; err != nil {
|
||||
t.Fatalf("logout: %v", err)
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "" || stored.RefreshToken != "" || deletes != 1 {
|
||||
t.Fatalf("refresh resurrected logged-out credential: %#v deletes=%d", stored, deletes)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func ptrTokenData(data TokenData) *TokenData {
|
||||
return &data
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthLockedRefreshReadsOpaqueEditionStore(t *testing.T) {
|
||||
data := validRejectedTokenData("expired-access")
|
||||
data.ExpiresAt = time.Now().Add(-time.Hour)
|
||||
store := installRejectedTokenHookStore(t, data)
|
||||
var refreshCalls atomic.Int32
|
||||
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, current *TokenData) (*TokenData, error) {
|
||||
refreshCalls.Add(1)
|
||||
updated := *current
|
||||
updated.AccessToken = "proactively-refreshed"
|
||||
updated.ExpiresAt = time.Now().Add(time.Hour)
|
||||
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &updated, nil
|
||||
})
|
||||
|
||||
token, err := NewOAuthProvider(t.TempDir(), nil).GetAccessToken(context.Background())
|
||||
if err != nil || token != "proactively-refreshed" {
|
||||
t.Fatalf("GetAccessToken() = %q, %v", token, err)
|
||||
}
|
||||
if refreshCalls.Load() != 1 {
|
||||
t.Fatalf("refresh calls = %d, want 1", refreshCalls.Load())
|
||||
}
|
||||
stored, deletes := store.snapshot()
|
||||
if stored.AccessToken != "proactively-refreshed" || deletes != 0 {
|
||||
t.Fatalf("stored token = %q, deletes = %d", stored.AccessToken, deletes)
|
||||
}
|
||||
}
|
||||
@@ -46,7 +46,7 @@ func TestCrossPlatformCoverageTokenPersistencePreflightRemainingEdges(t *testing
|
||||
if err := preflightTokenPersistence(t.TempDir()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := preflightTokenRefreshPersistence(&TokenData{CorpID: "corp"}); err != nil {
|
||||
if err := preflightTokenRefreshPersistence(t.TempDir(), &TokenData{CorpID: "corp"}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -79,7 +79,7 @@ func TestCrossPlatformCoverageTokenPersistencePreflightRemainingEdges(t *testing
|
||||
if err := preflightTokenPersistence(dir); !errors.Is(err, fail) {
|
||||
t.Fatalf("profile slot error = %v", err)
|
||||
}
|
||||
if err := preflightTokenRefreshPersistence(&TokenData{CorpID: "corp"}); !errors.Is(err, fail) {
|
||||
if err := preflightTokenRefreshPersistence(t.TempDir(), &TokenData{CorpID: "corp"}); !errors.Is(err, fail) {
|
||||
t.Fatalf("refresh profile slot error = %v", err)
|
||||
}
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
@@ -34,7 +35,17 @@ func TestMain(m *testing.M) {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
panic("set " + keychain.StorageDirEnv + ": " + err.Error())
|
||||
}
|
||||
if err := os.Setenv(keychain.TestNamespaceEnv, tmpDir); err != nil {
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
panic("set " + keychain.TestNamespaceEnv + ": " + err.Error())
|
||||
}
|
||||
code := m.Run()
|
||||
if err := keychain.RemoveAuthTokenEntries(keychain.Service); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "internal/auth keychain cleanup: %v\n", err)
|
||||
if code == 0 {
|
||||
code = 1
|
||||
}
|
||||
}
|
||||
_ = os.RemoveAll(tmpDir)
|
||||
os.Exit(code)
|
||||
}
|
||||
|
||||
+722
-80
@@ -28,6 +28,8 @@ import (
|
||||
|
||||
"github.com/google/uuid"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/keychain"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
@@ -35,35 +37,49 @@ var (
|
||||
tokenJSONMarshalIndent = json.MarshalIndent
|
||||
tokenJSONMarshal = json.Marshal
|
||||
tokenMkdirAll = os.MkdirAll
|
||||
tokenReadFile = os.ReadFile
|
||||
tokenWriteFile = os.WriteFile
|
||||
tokenRename = os.Rename
|
||||
tokenRemove = os.Remove
|
||||
tokenGlob = filepath.Glob
|
||||
tokenSaveKeychainForCorpID = SaveTokenDataKeychainForCorpID
|
||||
tokenSaveKeychainForIdentity = SaveTokenDataKeychainForIdentity
|
||||
tokenSaveKeychain = SaveTokenDataKeychain
|
||||
tokenLoadKeychainForCorpID = LoadTokenDataKeychainForCorpID
|
||||
tokenLoadKeychainIdentity = LoadTokenDataKeychainForIdentity
|
||||
tokenLoadKeychain = LoadTokenDataKeychain
|
||||
tokenKeychainExists = TokenDataExistsKeychain
|
||||
tokenDeleteKeychainForCorpID = DeleteTokenDataKeychainForCorpID
|
||||
tokenDeleteKeychainIdentity = DeleteTokenDataKeychainForIdentity
|
||||
tokenDeleteKeychain = DeleteTokenDataKeychain
|
||||
tokenRemoveAuthTokenEntries = keychain.RemoveAuthTokenEntries
|
||||
tokenLoadSecure = LoadSecureTokenData
|
||||
tokenDeleteSecure = DeleteSecureData
|
||||
tokenResolveProfile = resolveProfileForLoad
|
||||
tokenUpsertProfile = upsertProfileFromTokenWithCurrentLocked
|
||||
tokenRemoveProfile = removeProfileLocked
|
||||
tokenSyncLegacyMirror = syncLegacyTokenMirrorLocked
|
||||
tokenLoadProfiles = LoadProfiles
|
||||
tokenWriteMarker = WriteTokenMarker
|
||||
tokenDeleteMarker = DeleteTokenMarker
|
||||
tokenParseURL = url.Parse
|
||||
tokenNewRequest = http.NewRequestWithContext
|
||||
tokenDefaultConfigDir = getDefaultConfigDir
|
||||
tokenLoadData = LoadTokenData
|
||||
tokenRevokeURL = GetRevokeTokenURL
|
||||
tokenLogoutURL = LogoutURL
|
||||
tokenLogoutContinueURL = LogoutContinueURL
|
||||
tokenLogoutHTTPClient = &http.Client{Timeout: 10 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}
|
||||
tokenRevokeHTTPClient = &http.Client{Timeout: 10 * time.Second}
|
||||
tokenResolveProfile = func(configDir, selector string) (*Profile, error) {
|
||||
profile, _, err := resolveProfileForLoadLocked(configDir, selector)
|
||||
return profile, err
|
||||
}
|
||||
tokenResolveDeletion = resolveProfileDeletionSelection
|
||||
tokenResolveSelection = resolveProfileSelection
|
||||
tokenUpsertProfile = upsertProfileFromTokenWithCurrentLocked
|
||||
tokenRemoveProfile = removeProfileLocked
|
||||
tokenSyncLegacyMirror = syncLegacyTokenMirrorLocked
|
||||
tokenSyncOrganizationMirror = syncOrganizationTokenMirrorForProfile
|
||||
tokenLoadProfiles = LoadProfiles
|
||||
tokenSaveProfiles = SaveProfiles
|
||||
tokenWriteMarker = WriteTokenMarker
|
||||
tokenWriteManualMarker = WriteManualTokenMarker
|
||||
tokenDeleteMarker = DeleteTokenMarker
|
||||
tokenParseURL = url.Parse
|
||||
tokenNewRequest = http.NewRequestWithContext
|
||||
tokenDefaultConfigDir = getDefaultConfigDir
|
||||
tokenLoadData = LoadTokenData
|
||||
tokenRevokeURL = GetRevokeTokenURL
|
||||
tokenMCPBaseURL = GetMCPBaseURL
|
||||
tokenLogoutURL = LogoutURL
|
||||
tokenLogoutContinueURL = LogoutContinueURL
|
||||
tokenLogoutHTTPClient = &http.Client{Timeout: 10 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}
|
||||
tokenRevokeHTTPClient = &http.Client{Timeout: 10 * time.Second}
|
||||
)
|
||||
|
||||
// TokenData holds the OAuth token set persisted to disk.
|
||||
@@ -109,14 +125,34 @@ const tokenJSONFile = "token.json"
|
||||
// TokenMarker is a lightweight file the host application reads to detect
|
||||
// whether the CLI has a valid token without accessing the keychain.
|
||||
type TokenMarker struct {
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
ManualToken bool `json:"manual_token,omitempty"`
|
||||
// Revision changes on every credential publication. Runtime token caches
|
||||
// use it as a cheap cross-process invalidation signal without reading the
|
||||
// platform keychain on every request.
|
||||
Revision string `json:"revision,omitempty"`
|
||||
}
|
||||
|
||||
// WriteTokenMarker writes a token.json marker containing only an updated_at
|
||||
// timestamp. The host application uses this file's presence and mtime to
|
||||
// decide whether it needs to trigger a new auth exchange.
|
||||
func WriteTokenMarker(configDir string) error {
|
||||
marker := TokenMarker{UpdatedAt: time.Now().Format(time.RFC3339)}
|
||||
return writeTokenMarker(configDir, false)
|
||||
}
|
||||
|
||||
// WriteManualTokenMarker marks the legacy global keychain slot as an explicit
|
||||
// `auth login --token` credential. The additive field keeps older hosts, which
|
||||
// only inspect token.json presence and mtime, fully compatible.
|
||||
func WriteManualTokenMarker(configDir string) error {
|
||||
return writeTokenMarker(configDir, true)
|
||||
}
|
||||
|
||||
func writeTokenMarker(configDir string, manual bool) error {
|
||||
marker := TokenMarker{
|
||||
UpdatedAt: time.Now().Format(time.RFC3339),
|
||||
ManualToken: manual,
|
||||
Revision: uuid.NewString(),
|
||||
}
|
||||
data, _ := tokenJSONMarshalIndent(marker, "", " ")
|
||||
if err := tokenMkdirAll(configDir, 0o700); err != nil {
|
||||
return err
|
||||
@@ -128,6 +164,44 @@ func WriteTokenMarker(configDir string) error {
|
||||
return tokenRename(tmp, filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// ReadTokenMarkerRevision returns the current credential publication revision.
|
||||
// Existing markers without a revision remain readable, but callers must avoid
|
||||
// caching them because they cannot prove that the credential is unchanged.
|
||||
func ReadTokenMarkerRevision(configDir string) (revision string, present bool, err error) {
|
||||
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", false, nil
|
||||
}
|
||||
return "", false, fmt.Errorf("read token marker: %w", err)
|
||||
}
|
||||
var marker TokenMarker
|
||||
if err := json.Unmarshal(data, &marker); err != nil {
|
||||
// The marker is only a cache-coherency hint. A malformed historical or
|
||||
// externally modified marker must disable caching, not make an otherwise
|
||||
// valid credential unusable.
|
||||
return "", true, nil
|
||||
}
|
||||
return strings.TrimSpace(marker.Revision), true, nil
|
||||
}
|
||||
|
||||
func manualTokenMarkerActive(configDir string) (bool, error) {
|
||||
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return false, nil
|
||||
}
|
||||
return false, err
|
||||
}
|
||||
var marker TokenMarker
|
||||
if err := json.Unmarshal(data, &marker); err != nil {
|
||||
// Historical hosts only require the marker's presence. A malformed old
|
||||
// marker must not make profile authentication unusable.
|
||||
return false, nil
|
||||
}
|
||||
return marker.ManualToken, nil
|
||||
}
|
||||
|
||||
// DeleteTokenMarker removes the token.json marker file.
|
||||
func DeleteTokenMarker(configDir string) error {
|
||||
if err := tokenRemove(filepath.Join(configDir, tokenJSONFile)); err != nil && !os.IsNotExist(err) {
|
||||
@@ -136,13 +210,10 @@ func DeleteTokenMarker(configDir string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to the default keychain-based storage.
|
||||
// SaveTokenData persists TokenData under the auth dual lock. When an edition
|
||||
// hook (SaveToken) is registered, the locked write delegates to that hook;
|
||||
// otherwise it falls back to the default keychain-based storage.
|
||||
func SaveTokenData(configDir string, data *TokenData) error {
|
||||
if h := edition.Get(); h.SaveToken != nil {
|
||||
return saveTokenViaHook(h, configDir, data)
|
||||
}
|
||||
return withProfilesLock(configDir, func() error {
|
||||
return saveTokenDataLocked(configDir, data)
|
||||
})
|
||||
@@ -158,26 +229,126 @@ func saveTokenDataLocked(configDir string, data *TokenData) error {
|
||||
return saveTokenViaHook(h, configDir, data)
|
||||
}
|
||||
if data != nil && strings.TrimSpace(data.CorpID) != "" {
|
||||
if err := tokenSaveKeychainForCorpID(data.CorpID, data); err != nil {
|
||||
corpID := strings.TrimSpace(data.CorpID)
|
||||
userID := strings.TrimSpace(data.UserID)
|
||||
cfg, err := tokenLoadProfiles(configDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
makeCurrent := strings.TrimSpace(RuntimeProfile()) == ""
|
||||
if err := tokenUpsertProfile(configDir, data, makeCurrent); err != nil {
|
||||
if err := ensureProfilesWritable(cfg); err != nil {
|
||||
return err
|
||||
}
|
||||
runtimeSelector := strings.TrimSpace(RuntimeProfile())
|
||||
makeCurrent := runtimeSelector == ""
|
||||
exactSelector := profileSelector(corpID, userID)
|
||||
mirrorOrg := makeCurrent ||
|
||||
exactProfileSelectorForCorp(cfg, corpID, cfg.OrgCurrentProfiles[corpID]) == exactSelector
|
||||
existingIdentity := profileIndexByIdentity(cfg, corpID, userID) >= 0
|
||||
upgradesLegacyProfile := !existingIdentity && userID != "" && legacyProfileIndexByCorpID(cfg, corpID) >= 0
|
||||
logging.AuthDebug(
|
||||
"auth.token.persist.plan",
|
||||
"corp_id", corpID,
|
||||
"user_id", userID,
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"identity_selector", exactSelector,
|
||||
"existing_identity", existingIdentity,
|
||||
"upgrades_legacy_profile", upgradesLegacyProfile,
|
||||
"profiles_before", len(cfg.Profiles),
|
||||
"runtime_profile", runtimeSelector,
|
||||
"write_identity_slot", userID != "",
|
||||
"write_org_mirror", mirrorOrg,
|
||||
"write_global_mirror", makeCurrent,
|
||||
)
|
||||
snapshot, err := snapshotTokenPersistence(configDir, cfg, corpID, userID, mirrorOrg)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
preserveManualDefault := !makeCurrent &&
|
||||
snapshot.marker.known &&
|
||||
snapshot.marker.exists &&
|
||||
snapshot.marker.manual
|
||||
rollback := func(operationErr error) error {
|
||||
if rollbackErr := restoreTokenPersistence(configDir, snapshot); rollbackErr != nil {
|
||||
return errors.Join(operationErr, fmt.Errorf("rollback token persistence: %w", rollbackErr))
|
||||
}
|
||||
return operationErr
|
||||
}
|
||||
if userID != "" {
|
||||
if err := tokenSaveKeychainForIdentity(corpID, userID, data); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
} else {
|
||||
for _, profile := range cfg.Profiles {
|
||||
if strings.TrimSpace(profile.CorpID) == corpID && strings.TrimSpace(profile.UserID) != "" {
|
||||
return fmt.Errorf("cannot store profile for corpId %q without userId because account identities already exist", corpID)
|
||||
}
|
||||
}
|
||||
}
|
||||
if err := tokenUpsertProfile(configDir, data, makeCurrent); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
if mirrorOrg {
|
||||
if err := tokenSaveKeychainForCorpID(corpID, data); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
}
|
||||
if makeCurrent {
|
||||
if err := tokenSaveKeychain(data); err != nil {
|
||||
return err
|
||||
return rollback(err)
|
||||
}
|
||||
} else if !preserveManualDefault {
|
||||
if err := tokenSyncLegacyMirror(configDir); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
} else if err := tokenSyncLegacyMirror(configDir); err != nil {
|
||||
return err
|
||||
}
|
||||
return tokenWriteMarker(configDir)
|
||||
if preserveManualDefault {
|
||||
if err := tokenWriteManualMarker(configDir); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
} else if err := tokenWriteMarker(configDir); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
logging.AuthDebug(
|
||||
"auth.token.persist.done",
|
||||
"corp_id", corpID,
|
||||
"user_id", userID,
|
||||
"user_name", strings.TrimSpace(data.UserName),
|
||||
"identity_selector", exactSelector,
|
||||
"write_identity_slot", userID != "",
|
||||
"write_org_mirror", mirrorOrg,
|
||||
"write_global_mirror", makeCurrent,
|
||||
)
|
||||
return nil
|
||||
}
|
||||
legacySnapshot, err := snapshotTokenSlot(tokenLoadKeychain)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
markerSnapshot, err := snapshotTokenMarker(configDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tokenSaveKeychain(data); err != nil {
|
||||
return err
|
||||
}
|
||||
return tokenWriteMarker(configDir)
|
||||
if err := tokenWriteManualMarker(configDir); err != nil {
|
||||
var rollbackErr error
|
||||
if restoreErr := restoreTokenSlot(
|
||||
legacySnapshot,
|
||||
tokenSaveKeychain,
|
||||
tokenDeleteKeychain,
|
||||
); restoreErr != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, restoreErr)
|
||||
}
|
||||
if restoreErr := restoreTokenMarker(configDir, markerSnapshot); restoreErr != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, restoreErr)
|
||||
}
|
||||
if rollbackErr != nil {
|
||||
return errors.Join(err, fmt.Errorf("rollback manual token persistence: %w", rollbackErr))
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func saveTokenViaHook(h *edition.Hooks, configDir string, data *TokenData) error {
|
||||
@@ -213,13 +384,38 @@ func LoadTokenDataForProfile(configDir, profile string) (*TokenData, error) {
|
||||
return &td, nil
|
||||
}
|
||||
|
||||
var result *TokenData
|
||||
err := withProfilesLock(configDir, func() error {
|
||||
var loadErr error
|
||||
result, loadErr = loadTokenDataForProfileLocked(configDir, profile)
|
||||
return loadErr
|
||||
})
|
||||
return result, err
|
||||
}
|
||||
|
||||
func loadTokenDataForProfileLocked(configDir, profile string) (*TokenData, error) {
|
||||
// Default: keychain with legacy .data migration
|
||||
if strings.TrimSpace(profile) == "" {
|
||||
manual, err := manualTokenMarkerActive(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if manual {
|
||||
data, loadErr := tokenLoadKeychain()
|
||||
if loadErr == nil && data != nil && strings.TrimSpace(data.CorpID) == "" {
|
||||
return data, nil
|
||||
}
|
||||
if loadErr != nil && !errors.Is(loadErr, ErrTokenDataNotFound) {
|
||||
return nil, loadErr
|
||||
}
|
||||
}
|
||||
}
|
||||
selected, err := tokenResolveProfile(configDir, profile)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if selected != nil {
|
||||
data, err := tokenLoadKeychainForCorpID(selected.CorpID)
|
||||
data, err := tokenLoadProfileIdentity(*selected)
|
||||
if err == nil {
|
||||
return data, nil
|
||||
}
|
||||
@@ -231,13 +427,21 @@ func LoadTokenDataForProfile(configDir, profile string) (*TokenData, error) {
|
||||
// the SAME org; otherwise surface the error instead of silently acting
|
||||
// as a different organization (the legacy mirror may have drifted).
|
||||
if legacy, lerr := tokenLoadKeychain(); lerr == nil && legacy != nil &&
|
||||
strings.TrimSpace(legacy.CorpID) == strings.TrimSpace(selected.CorpID) {
|
||||
strings.TrimSpace(legacy.CorpID) == strings.TrimSpace(selected.CorpID) &&
|
||||
(strings.TrimSpace(selected.UserID) == "" || strings.TrimSpace(legacy.UserID) == strings.TrimSpace(selected.UserID)) {
|
||||
return legacy, nil
|
||||
} else if lerr != nil && !errors.Is(lerr, ErrTokenDataNotFound) {
|
||||
return nil, lerr
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
cfg, err := tokenLoadProfiles(configDir)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if cfg != nil && cfg.Version >= profilesVersion {
|
||||
return nil, ErrTokenDataNotFound
|
||||
}
|
||||
if tokenKeychainExists() {
|
||||
return tokenLoadKeychain()
|
||||
}
|
||||
@@ -253,9 +457,38 @@ func LoadTokenDataForProfile(configDir, profile string) (*TokenData, error) {
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
|
||||
// registered, it delegates entirely to the hook; otherwise it falls back
|
||||
// to keychain + legacy cleanup.
|
||||
func tokenLoadProfileIdentity(profile Profile) (*TokenData, error) {
|
||||
if strings.TrimSpace(profile.UserID) == "" {
|
||||
return tokenLoadKeychainForCorpID(profile.CorpID)
|
||||
}
|
||||
data, err := tokenLoadKeychainIdentity(profile.CorpID, profile.UserID)
|
||||
if err == nil {
|
||||
return data, nil
|
||||
}
|
||||
if !errors.Is(err, ErrTokenDataNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
orgData, orgErr := tokenLoadKeychainForCorpID(profile.CorpID)
|
||||
if orgErr != nil {
|
||||
if errors.Is(orgErr, ErrTokenDataNotFound) {
|
||||
return nil, err
|
||||
}
|
||||
return nil, orgErr
|
||||
}
|
||||
if strings.TrimSpace(orgData.UserID) == "" {
|
||||
return nil, fmt.Errorf("organization token mirror for corpId %q has no userId; cannot use it for profile %q", profile.CorpID, ProfileSelector(profile))
|
||||
}
|
||||
if strings.TrimSpace(orgData.UserID) != strings.TrimSpace(profile.UserID) {
|
||||
return nil, err
|
||||
}
|
||||
if saveErr := tokenSaveKeychainForIdentity(profile.CorpID, profile.UserID, orgData); saveErr != nil {
|
||||
return nil, saveErr
|
||||
}
|
||||
return orgData, nil
|
||||
}
|
||||
|
||||
// DeleteTokenData removes token data. Edition hooks and the default keychain
|
||||
// path are both serialized with refresh through the auth dual lock.
|
||||
func DeleteTokenData(configDir string) error {
|
||||
return DeleteTokenDataForProfile(configDir, RuntimeProfile())
|
||||
}
|
||||
@@ -267,7 +500,9 @@ func DeleteTokenDataForProfile(configDir, profile string) error {
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return fmt.Errorf("profile selection is not supported by the current auth backend")
|
||||
}
|
||||
return h.DeleteToken(configDir)
|
||||
return withProfilesLock(configDir, func() error {
|
||||
return h.DeleteToken(configDir)
|
||||
})
|
||||
}
|
||||
return withProfilesLock(configDir, func() error {
|
||||
return deleteTokenDataForProfileLocked(configDir, profile)
|
||||
@@ -275,25 +510,114 @@ func DeleteTokenDataForProfile(configDir, profile string) error {
|
||||
}
|
||||
|
||||
func deleteTokenDataForProfileLocked(configDir, profile string) error {
|
||||
selected, err := tokenResolveProfile(configDir, profile)
|
||||
if strings.TrimSpace(profile) == "" {
|
||||
manual, err := manualTokenMarkerActive(configDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if manual {
|
||||
return deleteManualTokenDataLocked(configDir)
|
||||
}
|
||||
}
|
||||
if err := profilesEnsureMigration(configDir); err != nil {
|
||||
return err
|
||||
}
|
||||
cfg, err := tokenLoadProfiles(configDir)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
effectiveSelector := strings.TrimSpace(profile)
|
||||
if effectiveSelector == "" {
|
||||
effectiveSelector = strings.TrimSpace(cfg.CurrentProfile)
|
||||
}
|
||||
var selected *Profile
|
||||
exact := false
|
||||
if effectiveSelector != "" {
|
||||
selected, exact, err = tokenResolveDeletion(cfg, effectiveSelector)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if selected != nil {
|
||||
keychainErr := tokenDeleteKeychainForCorpID(selected.CorpID)
|
||||
_, removeErr := tokenRemoveProfile(configDir, selected.CorpID)
|
||||
legacyErr := tokenSyncLegacyMirror(configDir)
|
||||
secureErr := tokenDeleteSecure(configDir)
|
||||
if keychainErr != nil {
|
||||
return keychainErr
|
||||
removed := *selected
|
||||
originalCfg := cloneProfilesConfig(cfg)
|
||||
identitySnapshots, err := snapshotDeletionIdentities(cfg, removed, exact)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if removeErr != nil {
|
||||
return removeErr
|
||||
orgSnapshot := snapshotTokenSlotForDeletion(func() (*TokenData, error) {
|
||||
return tokenLoadKeychainForCorpID(removed.CorpID)
|
||||
})
|
||||
legacySnapshot := snapshotTokenSlotForDeletion(tokenLoadKeychain)
|
||||
markerSnapshot := snapshotTokenMarkerForDeletion(configDir)
|
||||
|
||||
// Clean the deprecated secure-store copy before changing the profile
|
||||
// transaction. A cleanup failure therefore leaves all current metadata
|
||||
// and keychain slots untouched.
|
||||
if err := tokenDeleteSecure(configDir); err != nil {
|
||||
return err
|
||||
}
|
||||
if legacyErr != nil {
|
||||
return legacyErr
|
||||
|
||||
removeSelector := removed.CorpID
|
||||
orgCurrent := false
|
||||
if exact {
|
||||
removeSelector = ProfileSelector(removed)
|
||||
orgCurrent = exactProfileSelectorForCorp(
|
||||
cfg,
|
||||
removed.CorpID,
|
||||
cfg.OrgCurrentProfiles[removed.CorpID],
|
||||
) == ProfileSelector(removed)
|
||||
}
|
||||
return secureErr
|
||||
if _, err := tokenRemoveProfile(configDir, removeSelector); err != nil {
|
||||
return err
|
||||
}
|
||||
rollback := func(operationErr error) error {
|
||||
if rollbackErr := restoreProfileDeletion(
|
||||
configDir,
|
||||
originalCfg,
|
||||
identitySnapshots,
|
||||
removed.CorpID,
|
||||
orgSnapshot,
|
||||
legacySnapshot,
|
||||
markerSnapshot,
|
||||
); rollbackErr != nil {
|
||||
return errors.Join(operationErr, fmt.Errorf("rollback profile deletion: %w", rollbackErr))
|
||||
}
|
||||
return operationErr
|
||||
}
|
||||
|
||||
if !exact || orgCurrent {
|
||||
updated, loadErr := tokenLoadProfiles(configDir)
|
||||
if loadErr != nil {
|
||||
return rollback(loadErr)
|
||||
}
|
||||
replacementSelector := updated.OrgCurrentProfiles[removed.CorpID]
|
||||
if exact && replacementSelector != "" {
|
||||
replacement, _, resolveErr := tokenResolveSelection(configDir, updated, replacementSelector)
|
||||
if resolveErr != nil {
|
||||
return rollback(resolveErr)
|
||||
}
|
||||
if err := tokenSyncOrganizationMirror(*replacement); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
} else if err := tokenDeleteKeychainForCorpID(removed.CorpID); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
}
|
||||
preserveManualDefault := markerSnapshot.known &&
|
||||
markerSnapshot.exists &&
|
||||
markerSnapshot.manual
|
||||
if !preserveManualDefault {
|
||||
if err := tokenSyncLegacyMirror(configDir); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
}
|
||||
for _, snapshot := range identitySnapshots {
|
||||
if err := tokenDeleteKeychainIdentity(snapshot.profile.CorpID, snapshot.profile.UserID); err != nil {
|
||||
return rollback(err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
keychainErr := tokenDeleteKeychain()
|
||||
@@ -308,21 +632,321 @@ func deleteTokenDataForProfileLocked(configDir, profile string) error {
|
||||
return markerErr
|
||||
}
|
||||
|
||||
func deleteManualTokenDataLocked(configDir string) error {
|
||||
legacySnapshot := snapshotTokenSlotForDeletion(tokenLoadKeychain)
|
||||
if err := tokenDeleteSecure(configDir); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tokenDeleteKeychain(); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tokenDeleteMarker(configDir); err != nil {
|
||||
if legacySnapshot.known {
|
||||
if rollbackErr := restoreTokenSlot(
|
||||
legacySnapshot,
|
||||
tokenSaveKeychain,
|
||||
tokenDeleteKeychain,
|
||||
); rollbackErr != nil {
|
||||
return errors.Join(err, fmt.Errorf("rollback manual token deletion: %w", rollbackErr))
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
type deletionIdentitySnapshot struct {
|
||||
profile Profile
|
||||
token *TokenData
|
||||
}
|
||||
|
||||
type tokenSlotSnapshot struct {
|
||||
token *TokenData
|
||||
known bool
|
||||
exists bool
|
||||
}
|
||||
|
||||
type tokenMarkerSnapshot struct {
|
||||
known bool
|
||||
exists bool
|
||||
manual bool
|
||||
}
|
||||
|
||||
type tokenPersistenceSnapshot struct {
|
||||
profiles *ProfilesConfig
|
||||
corpID string
|
||||
userID string
|
||||
identity tokenSlotSnapshot
|
||||
org tokenSlotSnapshot
|
||||
legacy tokenSlotSnapshot
|
||||
marker tokenMarkerSnapshot
|
||||
}
|
||||
|
||||
func cloneProfilesConfig(cfg *ProfilesConfig) *ProfilesConfig {
|
||||
if cfg == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *cfg
|
||||
cloned.Profiles = append([]Profile(nil), cfg.Profiles...)
|
||||
for i := range cloned.Profiles {
|
||||
cloned.Profiles[i].AuthorizedDomains = append(
|
||||
[]string(nil),
|
||||
cloned.Profiles[i].AuthorizedDomains...,
|
||||
)
|
||||
}
|
||||
if cfg.OrgCurrentProfiles != nil {
|
||||
cloned.OrgCurrentProfiles = make(map[string]string, len(cfg.OrgCurrentProfiles))
|
||||
for corpID, selector := range cfg.OrgCurrentProfiles {
|
||||
cloned.OrgCurrentProfiles[corpID] = selector
|
||||
}
|
||||
}
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func snapshotDeletionIdentities(cfg *ProfilesConfig, removed Profile, exact bool) ([]deletionIdentitySnapshot, error) {
|
||||
var profiles []Profile
|
||||
if exact {
|
||||
profiles = []Profile{removed}
|
||||
} else {
|
||||
for _, candidate := range cfg.Profiles {
|
||||
if strings.TrimSpace(candidate.CorpID) == strings.TrimSpace(removed.CorpID) {
|
||||
profiles = append(profiles, candidate)
|
||||
}
|
||||
}
|
||||
}
|
||||
snapshots := make([]deletionIdentitySnapshot, 0, len(profiles))
|
||||
for _, candidate := range profiles {
|
||||
if strings.TrimSpace(candidate.UserID) == "" {
|
||||
continue
|
||||
}
|
||||
data, err := tokenLoadKeychainIdentity(candidate.CorpID, candidate.UserID)
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrTokenDataNotFound) {
|
||||
snapshots = append(snapshots, deletionIdentitySnapshot{profile: candidate})
|
||||
continue
|
||||
}
|
||||
// A damaged target slot must remain removable. It cannot be restored
|
||||
// during rollback, but every readable slot in the same transaction
|
||||
// still is.
|
||||
snapshots = append(snapshots, deletionIdentitySnapshot{profile: candidate})
|
||||
continue
|
||||
}
|
||||
snapshots = append(snapshots, deletionIdentitySnapshot{profile: candidate, token: data})
|
||||
}
|
||||
return snapshots, nil
|
||||
}
|
||||
|
||||
func snapshotTokenSlot(load func() (*TokenData, error)) (tokenSlotSnapshot, error) {
|
||||
data, err := load()
|
||||
if err != nil {
|
||||
if errors.Is(err, ErrTokenDataNotFound) {
|
||||
return tokenSlotSnapshot{known: true}, nil
|
||||
}
|
||||
return tokenSlotSnapshot{}, err
|
||||
}
|
||||
return tokenSlotSnapshot{token: data, known: true, exists: data != nil}, nil
|
||||
}
|
||||
|
||||
func snapshotTokenSlotForDeletion(load func() (*TokenData, error)) tokenSlotSnapshot {
|
||||
snapshot, err := snapshotTokenSlot(load)
|
||||
if err != nil {
|
||||
return tokenSlotSnapshot{}
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func snapshotTokenMarker(configDir string) (tokenMarkerSnapshot, error) {
|
||||
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return tokenMarkerSnapshot{known: true}, nil
|
||||
}
|
||||
return tokenMarkerSnapshot{}, err
|
||||
}
|
||||
var marker TokenMarker
|
||||
if err := json.Unmarshal(data, &marker); err != nil {
|
||||
return tokenMarkerSnapshot{known: true, exists: true}, nil
|
||||
}
|
||||
return tokenMarkerSnapshot{known: true, exists: true, manual: marker.ManualToken}, nil
|
||||
}
|
||||
|
||||
func snapshotTokenMarkerForDeletion(configDir string) tokenMarkerSnapshot {
|
||||
snapshot, err := snapshotTokenMarker(configDir)
|
||||
if err != nil {
|
||||
return tokenMarkerSnapshot{}
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func restoreProfileDeletion(
|
||||
configDir string,
|
||||
cfg *ProfilesConfig,
|
||||
identities []deletionIdentitySnapshot,
|
||||
corpID string,
|
||||
org tokenSlotSnapshot,
|
||||
legacy tokenSlotSnapshot,
|
||||
marker tokenMarkerSnapshot,
|
||||
) error {
|
||||
var rollbackErr error
|
||||
for _, snapshot := range identities {
|
||||
if snapshot.token == nil {
|
||||
continue
|
||||
}
|
||||
if err := tokenSaveKeychainForIdentity(
|
||||
snapshot.profile.CorpID,
|
||||
snapshot.profile.UserID,
|
||||
snapshot.token,
|
||||
); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if err := tokenSaveProfiles(configDir, cloneProfilesConfig(cfg)); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
if org.known {
|
||||
if org.exists {
|
||||
if err := tokenSaveKeychainForCorpID(corpID, org.token); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
} else if err := tokenDeleteKeychainForCorpID(corpID); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if legacy.known {
|
||||
if legacy.exists {
|
||||
if err := tokenSaveKeychain(legacy.token); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
} else if err := tokenDeleteKeychain(); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if marker.known {
|
||||
if err := restoreTokenMarker(configDir, marker); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
return rollbackErr
|
||||
}
|
||||
|
||||
func snapshotTokenPersistence(
|
||||
configDir string,
|
||||
cfg *ProfilesConfig,
|
||||
corpID, userID string,
|
||||
includeOrganization bool,
|
||||
) (tokenPersistenceSnapshot, error) {
|
||||
snapshot := tokenPersistenceSnapshot{
|
||||
profiles: cloneProfilesConfig(cfg),
|
||||
corpID: corpID,
|
||||
userID: userID,
|
||||
}
|
||||
var err error
|
||||
if strings.TrimSpace(userID) != "" {
|
||||
snapshot.identity, err = snapshotTokenSlot(func() (*TokenData, error) {
|
||||
return tokenLoadKeychainIdentity(corpID, userID)
|
||||
})
|
||||
if err != nil {
|
||||
return tokenPersistenceSnapshot{}, err
|
||||
}
|
||||
}
|
||||
if includeOrganization {
|
||||
snapshot.org, err = snapshotTokenSlot(func() (*TokenData, error) {
|
||||
return tokenLoadKeychainForCorpID(corpID)
|
||||
})
|
||||
if err != nil {
|
||||
return tokenPersistenceSnapshot{}, err
|
||||
}
|
||||
}
|
||||
snapshot.legacy, err = snapshotTokenSlot(tokenLoadKeychain)
|
||||
if err != nil {
|
||||
return tokenPersistenceSnapshot{}, err
|
||||
}
|
||||
snapshot.marker, err = snapshotTokenMarker(configDir)
|
||||
if err != nil {
|
||||
return tokenPersistenceSnapshot{}, err
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
|
||||
func restoreTokenPersistence(configDir string, snapshot tokenPersistenceSnapshot) error {
|
||||
var rollbackErr error
|
||||
if err := tokenSaveProfiles(configDir, cloneProfilesConfig(snapshot.profiles)); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
if strings.TrimSpace(snapshot.userID) != "" && snapshot.identity.known {
|
||||
if err := restoreTokenSlot(
|
||||
snapshot.identity,
|
||||
func(data *TokenData) error {
|
||||
return tokenSaveKeychainForIdentity(snapshot.corpID, snapshot.userID, data)
|
||||
},
|
||||
func() error {
|
||||
return tokenDeleteKeychainIdentity(snapshot.corpID, snapshot.userID)
|
||||
},
|
||||
); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if snapshot.org.known {
|
||||
if err := restoreTokenSlot(
|
||||
snapshot.org,
|
||||
func(data *TokenData) error {
|
||||
return tokenSaveKeychainForCorpID(snapshot.corpID, data)
|
||||
},
|
||||
func() error {
|
||||
return tokenDeleteKeychainForCorpID(snapshot.corpID)
|
||||
},
|
||||
); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if snapshot.legacy.known {
|
||||
if err := restoreTokenSlot(snapshot.legacy, tokenSaveKeychain, tokenDeleteKeychain); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
if snapshot.marker.known {
|
||||
if err := restoreTokenMarker(configDir, snapshot.marker); err != nil {
|
||||
rollbackErr = errors.Join(rollbackErr, err)
|
||||
}
|
||||
}
|
||||
return rollbackErr
|
||||
}
|
||||
|
||||
func restoreTokenSlot(
|
||||
snapshot tokenSlotSnapshot,
|
||||
save func(*TokenData) error,
|
||||
remove func() error,
|
||||
) error {
|
||||
if snapshot.exists {
|
||||
return save(snapshot.token)
|
||||
}
|
||||
return remove()
|
||||
}
|
||||
|
||||
func restoreTokenMarker(configDir string, marker tokenMarkerSnapshot) error {
|
||||
switch {
|
||||
case !marker.exists:
|
||||
return tokenDeleteMarker(configDir)
|
||||
case marker.manual:
|
||||
return tokenWriteManualMarker(configDir)
|
||||
default:
|
||||
return tokenWriteMarker(configDir)
|
||||
}
|
||||
}
|
||||
|
||||
// DeleteAllTokenData removes all profile-scoped and legacy token data.
|
||||
func DeleteAllTokenData(configDir string) error {
|
||||
if h := edition.Get(); h.DeleteToken != nil {
|
||||
return h.DeleteToken(configDir)
|
||||
return withProfilesLock(configDir, func() error {
|
||||
return h.DeleteToken(configDir)
|
||||
})
|
||||
}
|
||||
return withProfilesLock(configDir, func() error {
|
||||
var firstErr error
|
||||
// Best-effort: even if profiles.json is unreadable, still clear every
|
||||
// other slot so the user can always self-heal via auth reset / logout.
|
||||
if cfg, err := tokenLoadProfiles(configDir); err == nil {
|
||||
for _, profile := range cfg.Profiles {
|
||||
if e := tokenDeleteKeychainForCorpID(profile.CorpID); e != nil && firstErr == nil {
|
||||
firstErr = e
|
||||
}
|
||||
}
|
||||
// Sweep the complete auth-token namespace so orphan identity slots that
|
||||
// are not present in profiles.json cannot survive reset/logout --all.
|
||||
if e := tokenRemoveAuthTokenEntries(keychain.Service); e != nil {
|
||||
firstErr = e
|
||||
}
|
||||
if e := tokenRemove(ProfilesPath(configDir)); e != nil && !os.IsNotExist(e) && firstErr == nil {
|
||||
firstErr = e
|
||||
@@ -335,15 +959,19 @@ func DeleteAllTokenData(configDir string) error {
|
||||
}
|
||||
}
|
||||
}
|
||||
if e := tokenDeleteKeychain(); e != nil && firstErr == nil {
|
||||
firstErr = e
|
||||
}
|
||||
if e := tokenDeleteSecure(configDir); e != nil && firstErr == nil {
|
||||
firstErr = e
|
||||
}
|
||||
if e := tokenDeleteMarker(configDir); e != nil && firstErr == nil {
|
||||
firstErr = e
|
||||
}
|
||||
if firstErr != nil {
|
||||
// Preserve an explicit v2 empty registry so any stale mirror that
|
||||
// could not be removed is never imported on a later read.
|
||||
if e := profilesSave(configDir, &ProfilesConfig{Version: profilesVersion}); e != nil {
|
||||
return fmt.Errorf("%v; save logged-out profile tombstone: %w", firstErr, e)
|
||||
}
|
||||
}
|
||||
return firstErr
|
||||
})
|
||||
}
|
||||
@@ -353,18 +981,42 @@ func DeleteAllTokenData(configDir string) error {
|
||||
// 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)
|
||||
tokenData, err := tokenLoadData(tokenDefaultConfigDir())
|
||||
if err != nil || tokenData == nil {
|
||||
return nil
|
||||
}
|
||||
// Direct mode: use DingTalk logout endpoint
|
||||
// Historical token records may not have Source. Preserve the legacy
|
||||
// process-wide MCP decision only for those records.
|
||||
if strings.TrimSpace(tokenData.Source) == "" && IsClientIDFromMCP() {
|
||||
copy := *tokenData
|
||||
copy.Source = "mcp"
|
||||
tokenData = ©
|
||||
}
|
||||
return RevokeTokenRemoteForData(ctx, tokenData)
|
||||
}
|
||||
|
||||
// RevokeTokenRemoteForData revokes the supplied account token using the
|
||||
// credential source and client ID persisted with that exact identity.
|
||||
func RevokeTokenRemoteForData(ctx context.Context, tokenData *TokenData) error {
|
||||
if tokenData == nil {
|
||||
return nil
|
||||
}
|
||||
clientID := strings.TrimSpace(tokenData.ClientID)
|
||||
if clientID == "" {
|
||||
clientID = ClientID()
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(tokenData.Source), "mcp") {
|
||||
return revokeTokenViaMCP(ctx, tokenData, clientID)
|
||||
}
|
||||
|
||||
// Direct mode: use DingTalk logout endpoint.
|
||||
logoutURL, err := tokenParseURL(tokenLogoutURL)
|
||||
if err != nil {
|
||||
return fmt.Errorf("parsing logout URL: %w", err)
|
||||
}
|
||||
|
||||
q := logoutURL.Query()
|
||||
q.Set("client_id", ClientID())
|
||||
q.Set("client_id", clientID)
|
||||
q.Set("continue", tokenLogoutContinueURL)
|
||||
logoutURL.RawQuery = q.Encode()
|
||||
|
||||
@@ -388,20 +1040,10 @@ func RevokeTokenRemote(ctx context.Context) error {
|
||||
}
|
||||
|
||||
// revokeTokenViaMCP revokes token via MCP endpoint.
|
||||
func revokeTokenViaMCP(ctx context.Context) error {
|
||||
revokeURL := tokenRevokeURL()
|
||||
if revokeURL == "" {
|
||||
return nil // No revoke endpoint available
|
||||
}
|
||||
|
||||
// Load current token to get accessToken
|
||||
tokenData, err := tokenLoadData(tokenDefaultConfigDir())
|
||||
if err != nil || tokenData == nil {
|
||||
return nil // No token to revoke
|
||||
}
|
||||
|
||||
func revokeTokenViaMCP(ctx context.Context, tokenData *TokenData, clientID string) error {
|
||||
revokeURL := tokenMCPBaseURL() + MCPRevokeTokenPath
|
||||
body := map[string]string{
|
||||
"clientId": ClientID(),
|
||||
"clientId": clientID,
|
||||
"accessToken": tokenData.AccessToken,
|
||||
}
|
||||
bodyBytes, err := tokenJSONMarshal(body)
|
||||
|
||||
@@ -18,6 +18,7 @@ package auth
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -92,7 +93,7 @@ func TestLoadTokenDataFallsBackToLegacyOnlyWhenCurrentSlotIsMissing(t *testing.T
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadTokenDataDoesNotHideUnreadableCurrentSlotWithLegacyFallback(t *testing.T) {
|
||||
func TestLoadTokenDataUsesIdentitySlotWhenOrganizationMirrorIsUnreadable(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
configDir := t.TempDir()
|
||||
@@ -106,11 +107,11 @@ func TestLoadTokenDataDoesNotHideUnreadableCurrentSlotWithLegacyFallback(t *test
|
||||
}
|
||||
|
||||
loaded, err := LoadTokenData(configDir)
|
||||
if err == nil {
|
||||
t.Fatalf("LoadTokenData() = %#v, nil; want unreadable profile error", loaded)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenData() error = %v", err)
|
||||
}
|
||||
if loaded != nil {
|
||||
t.Fatalf("LoadTokenData() data = %#v, want nil", loaded)
|
||||
if loaded == nil || loaded.AccessToken != data.AccessToken || loaded.UserID != data.UserID {
|
||||
t.Fatalf("LoadTokenData() = %#v, want identity token %#v", loaded, data)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -150,6 +151,65 @@ func TestPreflightTokenPersistenceRejectsUnreadableProfileSlot(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestExactOrgCurrentRefreshRejectsUnreadableOrgMirror(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
configDir := t.TempDir()
|
||||
data := testToken("at_exact_refresh", "corp_exact", "Exact Org")
|
||||
data.UserID = "user_exact"
|
||||
if err := SaveTokenData(configDir, data); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(profileCiphertextPathForTest(data.CorpID), []byte("corrupt ciphertext"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile(profile ciphertext) error = %v", err)
|
||||
}
|
||||
|
||||
SetRuntimeProfile("corp_exact:user_exact")
|
||||
defer SetRuntimeProfile("")
|
||||
if err := preflightTokenRefreshPersistence(configDir, data); err == nil ||
|
||||
!strings.Contains(err.Error(), "profile token slot") {
|
||||
t.Fatalf("preflightTokenRefreshPersistence(exact current) error = %v, want unreadable org mirror", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExactNonOrgCurrentRefreshIgnoresUnreadableOrgMirror(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
configDir := t.TempDir()
|
||||
first := testToken("at_first", "corp_exact", "Exact Org")
|
||||
first.UserID = "user_first"
|
||||
second := testToken("at_second", "corp_exact", "Exact Org")
|
||||
second.UserID = "user_second"
|
||||
if err := SaveTokenData(configDir, first); err != nil {
|
||||
t.Fatalf("SaveTokenData(first) error = %v", err)
|
||||
}
|
||||
if err := SaveTokenData(configDir, second); err != nil {
|
||||
t.Fatalf("SaveTokenData(second) error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(profileCiphertextPathForTest(first.CorpID), []byte("corrupt ciphertext"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile(profile ciphertext) error = %v", err)
|
||||
}
|
||||
|
||||
SetRuntimeProfile("corp_exact:user_first")
|
||||
defer SetRuntimeProfile("")
|
||||
if err := preflightTokenRefreshPersistence(configDir, first); err != nil {
|
||||
t.Fatalf("preflightTokenRefreshPersistence(exact non-current) error = %v", err)
|
||||
}
|
||||
updated := *first
|
||||
updated.AccessToken = "at_first_refreshed"
|
||||
updated.RefreshToken = "rt_first_refreshed"
|
||||
if err := SaveTokenData(configDir, &updated); err != nil {
|
||||
t.Fatalf("SaveTokenData(exact non-current refresh) error = %v", err)
|
||||
}
|
||||
loaded, err := LoadTokenDataForProfile(configDir, "corp_exact:user_first")
|
||||
if err != nil {
|
||||
t.Fatalf("LoadTokenDataForProfile(refreshed) error = %v", err)
|
||||
}
|
||||
if loaded.AccessToken != updated.AccessToken || loaded.RefreshToken != updated.RefreshToken {
|
||||
t.Fatalf("refreshed exact token = %#v, want %#v", loaded, updated)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExchangeAuthCodePreflightsOrphanProfileCiphertextBeforeHTTP(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
@@ -232,7 +292,7 @@ func TestRefreshPreflightIgnoresUnreadableUnrelatedProfile(t *testing.T) {
|
||||
t.Fatalf("WriteFile(A profile ciphertext) error = %v", err)
|
||||
}
|
||||
|
||||
if err := preflightTokenRefreshPersistence(dataB); err != nil {
|
||||
if err := preflightTokenRefreshPersistence(configDir, dataB); err != nil {
|
||||
t.Fatalf("preflightTokenRefreshPersistence(B) error = %v", err)
|
||||
}
|
||||
loaded, err := NewOAuthProvider(configDir, nil).Login(context.Background(), false)
|
||||
@@ -343,6 +403,45 @@ func TestLockedRefreshPreflightsLegacyMirrorBeforeHTTP(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestLockedRefreshRejectsFutureProfilesVersionBeforeHTTP(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
setPreflightTestCredentials(t)
|
||||
configDir := t.TempDir()
|
||||
data := testToken("at_future_refresh", "corp_future", "Future Org")
|
||||
data.ExpiresAt = time.Now().Add(-time.Hour)
|
||||
data.RefreshExpAt = time.Now().Add(time.Hour)
|
||||
if err := SaveTokenData(configDir, data); err != nil {
|
||||
t.Fatalf("SaveTokenData() error = %v", err)
|
||||
}
|
||||
cfg, err := LoadProfiles(configDir)
|
||||
if err != nil {
|
||||
t.Fatalf("LoadProfiles() error = %v", err)
|
||||
}
|
||||
cfg.Version = profilesVersion + 1
|
||||
raw, err := json.Marshal(cfg)
|
||||
if err != nil {
|
||||
t.Fatalf("json.Marshal() error = %v", err)
|
||||
}
|
||||
if err := os.WriteFile(ProfilesPath(configDir), raw, 0o600); err != nil {
|
||||
t.Fatalf("write future profiles: %v", err)
|
||||
}
|
||||
|
||||
var calls atomic.Int32
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
provider.httpClient = &http.Client{Transport: preflightRoundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
calls.Add(1)
|
||||
return nil, errors.New("unexpected refresh request")
|
||||
})}
|
||||
_, err = provider.Login(context.Background(), false)
|
||||
if err == nil || !strings.Contains(err.Error(), "newer than supported") {
|
||||
t.Fatalf("Login() error = %v, want future profiles rejection", err)
|
||||
}
|
||||
if got := calls.Load(); got != 0 {
|
||||
t.Fatalf("refresh HTTP calls = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExchangeAuthCodeAllowsFirstLogin(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
@@ -374,3 +473,37 @@ func TestExchangeAuthCodeAllowsFirstLogin(t *testing.T) {
|
||||
t.Fatalf("HTTP calls = %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExchangeAuthCodeExplicitUIDSkipsIdentityOverride(t *testing.T) {
|
||||
cleanupKeychain(t)
|
||||
t.Setenv(keychain.DisableKeychainEnv, "1")
|
||||
setPreflightTestCredentials(t)
|
||||
configDir := t.TempDir()
|
||||
|
||||
var identityCalls atomic.Int32
|
||||
provider := NewOAuthProvider(configDir, nil)
|
||||
provider.IdentityEnricher = func(context.Context, *TokenData) error {
|
||||
identityCalls.Add(1)
|
||||
return errors.New("identity lookup should not run for explicit uid")
|
||||
}
|
||||
provider.httpClient = &http.Client{Transport: preflightRoundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(
|
||||
`{"accessToken":"new-access","refreshToken":"new-refresh","expiresIn":7200,"corpId":"corp_new"}`,
|
||||
)),
|
||||
}, nil
|
||||
})}
|
||||
|
||||
data, err := provider.ExchangeAuthCode(context.Background(), "new-code", "explicit-user")
|
||||
if err != nil {
|
||||
t.Fatalf("ExchangeAuthCode() error = %v", err)
|
||||
}
|
||||
if data.UserID != "explicit-user" {
|
||||
t.Fatalf("ExchangeAuthCode() userId = %q, want explicit-user", data.UserID)
|
||||
}
|
||||
if got := identityCalls.Load(); got != 0 {
|
||||
t.Fatalf("IdentityEnricher calls = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user