Compare commits
60
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 | ||
|
|
99e5a3cceb | ||
|
|
7e31043875 | ||
|
|
55d7fbf59a | ||
|
|
c0f4d21c05 | ||
|
|
660c908585 | ||
|
|
377ebc5e85 | ||
|
|
2eca203e74 |
+699
-135
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
|
||||
|
||||
@@ -6,6 +6,60 @@ 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.
|
||||
@@ -13,6 +67,7 @@ The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and th
|
||||
|
||||
### 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.
|
||||
|
||||
+27
-15
@@ -108,21 +108,33 @@ 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.
|
||||
|
||||
If the immutable GitHub Release and npm package were delivered but a 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 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. It then downloads and re-verifies the immutable assets before invoking
|
||||
only the selected mirror. 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.
|
||||
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
|
||||
|
||||
|
||||
+73
-21
@@ -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` workflow 的九个 context:`Lint`、`Test`、`Coverage`、`Policy`、`Edition`、`Interface Integrity`、`AI Behavior`、`CLI Smoke`、`Mock MCP`,操作机已安装并登录 `gh`。本地脚本会在封 tag 前通过 API 检查 immutable releases、当前 SHA 的全部九个 context 和在途 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 一起发布。
|
||||
|
||||
## 预发发布
|
||||
|
||||
@@ -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,20 +130,20 @@ 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 分发失败且 GitHub immutable Release、npm 已交付时,从受保护的默认分支触发
|
||||
已启用的 OSS 或 Gitee 分发失败且 GitHub immutable Release、npm 已交付时,从受保护的默认分支触发
|
||||
Release workflow,并且只填写 `repair_oss_version` 或 `repair_gitee_version` 之一。channel
|
||||
repair 会精确绑定失败 tag run 的最新 attempt;contract、构建、Developer ID 签名、
|
||||
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 成功,
|
||||
@@ -110,25 +154,27 @@ workflow 和本地直发脚本已停用,避免绕开 publication queue 或用
|
||||
|
||||
## 既有 tag 的紧急恢复
|
||||
|
||||
tag push 已成功、但 Release workflow 失败且 GitHub Release 尚未公开时,不要新建临时 workflow、移动 tag 或跳过门禁。在最新且干净的 `main` worktree 运行:
|
||||
云端封板或本地 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-push run;也可以用 `--failed-run <run-id>` 精确指定。确认完整版本号后,它从默认分支触发受保护的恢复模式并等待完成。恢复模式必须满足:
|
||||
命令会自动解析 annotated tag object、peeled commit,以及 tag 绑定的失败云端 run 或最近一次匹配的失败 tag-push run;也可以用 `--failed-run <run-id>` 精确指定。确认完整版本号后,它从默认分支触发受保护的恢复模式并等待完成。恢复模式必须满足:
|
||||
|
||||
- 输入精确绑定原 annotated tag object、commit 和失败的 exact-tag `Release` run;commit 必须仍在 `main` 历史中。
|
||||
- 目标只允许不存在 GitHub Release 或仍为 Draft;已经公开的版本只能走对应的 channel repair,不能全量重建。
|
||||
- 输入精确绑定原 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 或门禁跳过。
|
||||
- 恢复复用正常的 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 中的固定证据。
|
||||
|
||||
OSS 的 `latest.txt` / `beta.txt` 当前是镜像频道元数据;仓库内安装器仍从 GitHub/Gitee 解析版本,不能把 OSS pointer 当成已接入的安装通道。
|
||||
云端 seal 后不要使用 GitHub 的 “Re-run failed jobs” 作为交付修复:annotated tag 永久绑定最初的 run attempt,普通 rerun 不会成为可接受的交付证据。GitHub Release 尚未公开时走上述 protected recovery;已经公开且仅 npm/OSS/Gitee 某一渠道失败时走对应 repair;版本内容本身有问题时走 withdrawal。
|
||||
|
||||
Homebrew 当前只属于本机预检/手工公式通道:预检会在当前 macOS 架构真实安装,但 Release workflow 不发布 tap,CI 生成的单主机公式也不应当作 Darwin 双架构正式交付。正式自动交付范围是 GitHub Release、npm、OSS,以及显式开启时的 Gitee fallback;Homebrew 双架构 tap 发布需另立需求。
|
||||
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 日志承担审计/续跑依据。
|
||||
|
||||
## 平台治理前置
|
||||
|
||||
@@ -136,8 +182,14 @@ Homebrew 当前只属于本机预检/手工公式通道:预检会在当前 mac
|
||||
|
||||
- `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,或任一 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
@@ -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) {
|
||||
|
||||
@@ -56,6 +56,7 @@ var (
|
||||
eventNewEventSource = newEventSource
|
||||
eventNewDingtalkSource = source.New
|
||||
eventResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
eventForceRefreshRejected = forceRefreshRejectedAccessToken
|
||||
eventBusRun = bus.Run
|
||||
eventReadyFDFromEnv = busctl.ReadyFDFromEnv
|
||||
eventResolvePersonal = resolvePersonalEventIdentity
|
||||
@@ -436,6 +437,9 @@ func newEventSource(_ context.Context, configDir, clientID, clientSecret string,
|
||||
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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -131,6 +131,7 @@ var (
|
||||
personalFindProcess = os.FindProcess
|
||||
personalSignalProcess = (*os.Process).Signal
|
||||
personalResolveAuxiliaryAccessToken = ResolveAuxiliaryAccessToken
|
||||
personalForceRefreshRejectedToken = forceRefreshRejectedAccessToken
|
||||
personalLoadTokenData = authpkg.LoadTokenData
|
||||
personalClientID = authpkg.ClientID
|
||||
personalResolveAppCredentialsStrict = authpkg.ResolveAppCredentialsStrict
|
||||
@@ -832,6 +833,9 @@ func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptio
|
||||
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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+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)
|
||||
}
|
||||
|
||||
|
||||
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
|
||||
}
|
||||
|
||||
+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 {
|
||||
|
||||
+37
-3
@@ -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)
|
||||
@@ -485,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
|
||||
@@ -766,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,
|
||||
@@ -778,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
|
||||
@@ -798,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(
|
||||
|
||||
@@ -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,6 +337,16 @@ 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, err := r.resolveAuthToken(context.Background()); err != nil || got != "explicit" {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -722,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"))
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -18,6 +18,7 @@ import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"html"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
@@ -301,7 +302,7 @@ 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:
|
||||
@@ -642,6 +643,14 @@ continueLogin:
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
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.
|
||||
@@ -665,7 +674,12 @@ func (p *OAuthProvider) GetTokenSnapshot(ctx context.Context) (*TokenData, error
|
||||
if rErr == nil {
|
||||
return refreshed, nil
|
||||
}
|
||||
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
|
||||
// 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)
|
||||
}
|
||||
if p.logger != nil {
|
||||
p.logger.Warn(i18n.T("refresh_token 刷新失败"), "error", rErr)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -304,3 +304,14 @@ func splitSchemaPathTokens(raw string) []string {
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// normalizeSchemaQueryCLIPath accepts the historical query spellings while
|
||||
// keeping authored Registry CLI paths strict and space-separated. Canonical
|
||||
// identity lookup still runs before this compatibility normalization.
|
||||
func normalizeSchemaQueryCLIPath(path string) string {
|
||||
parts := splitSchemaPathTokens(strings.TrimSpace(path))
|
||||
if len(parts) > 0 && parts[0] == "dws" {
|
||||
parts = parts[1:]
|
||||
}
|
||||
return strings.Join(parts, " ")
|
||||
}
|
||||
|
||||
@@ -290,7 +290,7 @@ func schemaPayloadFromLoadedCatalog(loaded loadedSchemaCatalog, args []string) (
|
||||
return payload, nil
|
||||
}
|
||||
raw := strings.TrimSpace(args[0])
|
||||
if tool, ok := loaded.Index.Resolve(raw); ok {
|
||||
if tool, ok := loaded.Index.ResolveQuery(raw); ok {
|
||||
return schemaToolForResolvedPath(tool, raw).ToPayload()
|
||||
}
|
||||
tokens := splitSchemaPathTokens(raw)
|
||||
|
||||
@@ -194,6 +194,23 @@ func TestCrossPlatformCoverageSchemaCatalogLookupAndConversionEdges(t *testing.T
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmbeddedSchemaLookupAcceptsCompatibleCLIPathSeparators(t *testing.T) {
|
||||
loaded := embeddedSchemaCatalog()
|
||||
for _, path := range []string{
|
||||
"dev app list",
|
||||
"dev.app.list",
|
||||
"dev/app/list",
|
||||
} {
|
||||
payload, err := schemaPayloadFromLoadedCatalog(loaded, []string{path})
|
||||
if err != nil {
|
||||
t.Fatalf("schemaPayloadFromLoadedCatalog(%q) error = %v", path, err)
|
||||
}
|
||||
if got := schemaString(payload["canonical_path"]); got != "dev.list_dev_app" {
|
||||
t.Fatalf("schemaPayloadFromLoadedCatalog(%q) canonical_path = %q, want %q", path, got, "dev.list_dev_app")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type catalogHookSnapshot struct {
|
||||
parameterBindings func(BoundCommandRegistry, SchemaRegistry) error
|
||||
dryRun func(SchemaRegistry) error
|
||||
|
||||
@@ -617,6 +617,22 @@ func (i SchemaIndex) Resolve(path string) (ToolSpec, bool) {
|
||||
return i.registry.Products[location.product].Tools[location.tool], true
|
||||
}
|
||||
|
||||
// ResolveQuery adds compatibility for dotted and slash-separated CLI paths at
|
||||
// the user-facing query boundary. Resolve remains strict because Registry
|
||||
// validation uses it to detect missing canonical identities without falling
|
||||
// through to a similarly spelled CLI path.
|
||||
func (i SchemaIndex) ResolveQuery(path string) (ToolSpec, bool) {
|
||||
if tool, ok := i.Resolve(path); ok {
|
||||
return tool, true
|
||||
}
|
||||
canonical, ok := i.byCLIPath[normalizeSchemaQueryCLIPath(path)]
|
||||
if !ok {
|
||||
return ToolSpec{}, false
|
||||
}
|
||||
location := i.byCanonical[canonical]
|
||||
return i.registry.Products[location.product].Tools[location.tool], true
|
||||
}
|
||||
|
||||
// CanonicalPaths returns the complete tool identity set in stable order.
|
||||
func (i SchemaIndex) CanonicalPaths() []string {
|
||||
paths := make([]string, 0, len(i.byCanonical))
|
||||
|
||||
@@ -514,6 +514,19 @@ func TestSchemaRegistryIndexResolvesCanonicalCLIAndAlias(t *testing.T) {
|
||||
t.Fatalf("Resolve(%q) = %#v, %v", path, resolved.Identity, ok)
|
||||
}
|
||||
}
|
||||
for _, path := range []string{
|
||||
"calendar.attendee.delete",
|
||||
"calendar/attendee/delete",
|
||||
"dws.calendar.attendee.delete",
|
||||
} {
|
||||
resolved, ok := index.ResolveQuery(path)
|
||||
if !ok || resolved.Identity.CanonicalPath != "calendar.attendee_delete" {
|
||||
t.Fatalf("ResolveQuery(%q) = %#v, %v", path, resolved.Identity, ok)
|
||||
}
|
||||
}
|
||||
if _, ok := index.ResolveQuery("calendar.attendee.unknown"); ok {
|
||||
t.Fatal("unknown dotted CLI path unexpectedly resolved")
|
||||
}
|
||||
if got := index.CanonicalPaths(); !reflect.DeepEqual(got, []string{"calendar.attendee_delete"}) {
|
||||
t.Fatalf("CanonicalPaths() = %#v", got)
|
||||
}
|
||||
|
||||
@@ -154,6 +154,35 @@ func TestSelectionExplicitEmptyListSurvivesFinalDelivery(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompatibleSchemaAliasSeparatorsSurviveFinalDelivery(t *testing.T) {
|
||||
snapshot := schemaDeliveryTestSnapshot(schemaDeliveryTestTool{
|
||||
Canonical: "sample.run",
|
||||
CLIPath: "sample category run",
|
||||
Aliases: []string{"sample legacy execute"},
|
||||
})
|
||||
encoded, err := json.Marshal(snapshot)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loaded, err := decodeSchemaCatalogSnapshot(encoded)
|
||||
if err != nil {
|
||||
t.Fatalf("decodeSchemaCatalogSnapshot(): %v", err)
|
||||
}
|
||||
canonical, err := schemaPayloadFromLoadedCatalog(loaded, []string{"sample.run"})
|
||||
if err != nil {
|
||||
t.Fatalf("canonical query: %v", err)
|
||||
}
|
||||
for _, path := range []string{"sample legacy execute", "sample.legacy.execute", "sample/legacy/execute"} {
|
||||
alias, aliasErr := schemaPayloadFromLoadedCatalog(loaded, []string{path})
|
||||
if aliasErr != nil {
|
||||
t.Fatalf("alias query %q: %v", path, aliasErr)
|
||||
}
|
||||
if problem := schemaAliasViewProblem(canonical, alias, "sample legacy execute"); problem != "" {
|
||||
t.Fatalf("alias projection for %q: %s", path, problem)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidateSchemaDeliveryInvariantsAllowsOnlyEnvelopeHashes(t *testing.T) {
|
||||
snapshot := schemaDeliveryTestSnapshot(schemaDeliveryTestTool{Canonical: "sample.run", CLIPath: "sample run"})
|
||||
snapshot.SurfaceHash = "sha256:reviewed-command-registry"
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
// Go's normal per-package coverage accounting attributes the exercised Schema
|
||||
// assembly code to internal/cli.
|
||||
func TestCrossPlatformCoverageProductionSchemaSourcePipeline(t *testing.T) {
|
||||
root := app.NewRootCommand()
|
||||
root := app.NewSchemaSourceRootCommand()
|
||||
resolved, err := cli.ResolveSchemaBuild(root)
|
||||
if err != nil {
|
||||
t.Fatalf("ResolveSchemaBuild() error = %v", err)
|
||||
@@ -33,17 +33,17 @@ func TestCrossPlatformCoverageProductionSchemaSourcePipeline(t *testing.T) {
|
||||
if len(snapshot.Tools) == 0 {
|
||||
t.Fatal("production Schema snapshot contains no tools")
|
||||
}
|
||||
registry, err := cli.AssembleSchemaRegistry(app.NewRootCommand())
|
||||
registry, err := cli.AssembleSchemaRegistry(app.NewSchemaSourceRootCommand())
|
||||
if err != nil {
|
||||
t.Fatalf("AssembleSchemaRegistry() error = %v", err)
|
||||
}
|
||||
if len(registry.Products) == 0 {
|
||||
t.Fatal("assembled production Schema registry contains no products")
|
||||
}
|
||||
if err := cli.ValidateEmbeddedRuntimeSchemaCompleteness(app.NewRootCommand()); err != nil {
|
||||
if err := cli.ValidateEmbeddedRuntimeSchemaCompleteness(app.NewSchemaSourceRootCommand()); err != nil {
|
||||
t.Fatalf("ValidateEmbeddedRuntimeSchemaCompleteness() error = %v", err)
|
||||
}
|
||||
root = app.NewRootCommand()
|
||||
root = app.NewSchemaSourceRootCommand()
|
||||
if _, err := cli.ApplyEmbeddedManualSchemaHints(root); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -284,7 +284,7 @@ func runtimeSchemaPayloadFromRegistry(registry SchemaRegistry, args []string) (m
|
||||
}
|
||||
|
||||
raw := strings.TrimSpace(args[0])
|
||||
if tool, ok := index.Resolve(raw); ok {
|
||||
if tool, ok := index.ResolveQuery(raw); ok {
|
||||
tool = schemaToolForResolvedPath(tool, raw)
|
||||
return renderRegistryToolPayload(tool)
|
||||
}
|
||||
@@ -338,7 +338,7 @@ func runtimeSchemaAllPayloadFromRegistry(registry SchemaRegistry) (map[string]an
|
||||
}
|
||||
|
||||
func schemaToolForResolvedPath(tool ToolSpec, raw string) ToolSpec {
|
||||
normalized := normalizeSchemaCLIPath(raw)
|
||||
normalized := normalizeSchemaQueryCLIPath(raw)
|
||||
if normalized == "" || normalized == tool.Identity.CLIPath || normalized == tool.Identity.PrimaryCLIPath {
|
||||
return tool
|
||||
}
|
||||
|
||||
@@ -414,7 +414,7 @@ func TestCrossPlatformCoveragePortalStartEndToEndAndFailures(t *testing.T) {
|
||||
return s
|
||||
}
|
||||
network := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) { return nil, errSourceInjected })}
|
||||
s := makeSource(&PortalTicketConfig{TicketURL: "https://x", AccessToken: "t", SourceID: "s", HTTPClient: network})
|
||||
s := makeSource(&PortalTicketConfig{TicketURL: "https://x", AccessToken: "t", SourceID: "s", HTTPClient: network, DisableReconnect: true})
|
||||
if err := s.Start(context.Background(), func(*dwsevent.RawEvent) {}); err == nil {
|
||||
t.Fatal("ticket failure expected")
|
||||
}
|
||||
@@ -483,7 +483,7 @@ func TestCrossPlatformCoveragePortalStartHandshakeReadAndAckErrors(t *testing.T)
|
||||
makeSource := func(endpoint string) *DingtalkSource {
|
||||
s, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: "https://ticket", AccessToken: "token", SourceID: "source",
|
||||
HTTPClient: staticPersonalTicketClient(endpoint, ""),
|
||||
HTTPClient: staticPersonalTicketClient(endpoint, ""), DisableReconnect: true,
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
|
||||
@@ -27,6 +27,7 @@ import (
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/config"
|
||||
"github.com/gorilla/websocket"
|
||||
@@ -43,6 +44,7 @@ const (
|
||||
type PersonalConfig struct {
|
||||
AccessToken string
|
||||
AccessTokenProvider AccessTokenProvider
|
||||
ForceRefreshToken ForceRefreshTokenFn
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
SourceID string
|
||||
@@ -57,6 +59,13 @@ type PersonalConfig struct {
|
||||
|
||||
type AccessTokenProvider func(context.Context) (string, error)
|
||||
|
||||
// ForceRefreshTokenFn rotates an access token that the server has just
|
||||
// rejected (HTTP 401). It receives the exact rejected token so the caller's
|
||||
// compare-and-refresh logic can skip the refresh when another goroutine has
|
||||
// already rotated it, and returns the fresh token to retry with. Optional:
|
||||
// when nil a 401 stays fatal, matching the previous behavior.
|
||||
type ForceRefreshTokenFn func(ctx context.Context, rejectedToken string) (string, error)
|
||||
|
||||
type PersonalSource struct {
|
||||
cfg PersonalConfig
|
||||
machine *Machine
|
||||
@@ -197,8 +206,29 @@ func (s *PersonalSource) runAttempt(ctx context.Context, emit dwsevent.EmitFn) (
|
||||
func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, error) {
|
||||
accessToken, err := resolveSourceAccessToken(ctx, s.cfg.AccessTokenProvider, s.cfg.AccessToken, "personal source")
|
||||
if err != nil {
|
||||
// Transient provider failures (network, 429, 5xx) must not kill a
|
||||
// long-running source; the reconnect loop retries after backoff.
|
||||
if authpkg.ClassifyRefreshFailure(err) == authpkg.RefreshFailureTransient {
|
||||
return nil, retryPersonal(err)
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
ticket, status, err := s.fetchTicketAttempt(ctx, accessToken)
|
||||
if status == http.StatusUnauthorized && s.cfg.ForceRefreshToken != nil {
|
||||
refreshed, refreshErr := refreshRejectedSourceToken(ctx, s.cfg.ForceRefreshToken, accessToken, "personal source", err)
|
||||
if refreshErr != nil {
|
||||
if authpkg.ClassifyRefreshFailure(refreshErr) == authpkg.RefreshFailureTransient {
|
||||
return nil, retryPersonal(refreshErr)
|
||||
}
|
||||
return nil, refreshErr
|
||||
}
|
||||
// Retry once with the freshly rotated token; a second 401 stays fatal.
|
||||
ticket, _, err = s.fetchTicketAttempt(ctx, refreshed)
|
||||
}
|
||||
return ticket, err
|
||||
}
|
||||
|
||||
func (s *PersonalSource) fetchTicketAttempt(ctx context.Context, accessToken string) (*ticketResponse, int, error) {
|
||||
body := map[string]any{
|
||||
"sourceId": s.cfg.SourceID,
|
||||
"mode": s.cfg.TicketMode,
|
||||
@@ -210,7 +240,7 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
|
||||
b, _ := json.Marshal(body)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, s.cfg.TicketURL, bytes.NewReader(b))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("personal source: create ticket request: %w", err)
|
||||
return nil, 0, fmt.Errorf("personal source: create ticket request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
@@ -221,28 +251,50 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
|
||||
|
||||
resp, err := s.cfg.HTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, retryPersonal(fmt.Errorf("personal source: fetch ticket: %w", err))
|
||||
return nil, 0, retryPersonal(fmt.Errorf("personal source: fetch ticket: %w", err))
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if err != nil {
|
||||
return nil, retryPersonal(fmt.Errorf("personal source: read ticket response: %w", err))
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
// Classify by status before touching the body: a truncated error body
|
||||
// must not upgrade a fatal status (notably 401) into a retryable
|
||||
// error, or the outer reconnect loop would bypass the single
|
||||
// refresh-retry guard. The body is not used here, so drain it only
|
||||
// best-effort for connection reuse.
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
err := fmt.Errorf("personal source: ticket HTTP %d", resp.StatusCode)
|
||||
if retryableTicketStatus(resp.StatusCode) {
|
||||
return nil, retryPersonal(err)
|
||||
return nil, resp.StatusCode, retryPersonal(err)
|
||||
}
|
||||
return nil, err
|
||||
return nil, resp.StatusCode, err
|
||||
}
|
||||
data, err := io.ReadAll(io.LimitReader(resp.Body, config.MaxResponseBodySize))
|
||||
if err != nil {
|
||||
return nil, resp.StatusCode, retryPersonal(fmt.Errorf("personal source: read ticket response: %w", err))
|
||||
}
|
||||
ticket, err := decodeTicket(data)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, resp.StatusCode, err
|
||||
}
|
||||
if ticket.Endpoint == "" || ticket.Ticket == "" {
|
||||
return nil, errors.New("personal source: ticket response missing endpoint or ticket")
|
||||
return nil, resp.StatusCode, errors.New("personal source: ticket response missing endpoint or ticket")
|
||||
}
|
||||
return ticket, nil
|
||||
return ticket, resp.StatusCode, nil
|
||||
}
|
||||
|
||||
// refreshRejectedSourceToken funnels a server-side 401 into the optional
|
||||
// force-refresh callback. It hands the actual rejected token to the caller's
|
||||
// compare-and-refresh logic and returns the rotated token for an immediate
|
||||
// retry. Refresh failures keep the original 401 as context instead of being
|
||||
// dropped.
|
||||
func refreshRejectedSourceToken(ctx context.Context, refresh ForceRefreshTokenFn, rejectedToken, component string, cause error) (string, error) {
|
||||
token, err := refresh(ctx, rejectedToken)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: refresh rejected access token: %w", component, errors.Join(cause, err))
|
||||
}
|
||||
if token = strings.TrimSpace(token); token == "" {
|
||||
return "", fmt.Errorf("%s: refresh rejected access token returned empty token: %w", component, cause)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func resolveSourceAccessToken(ctx context.Context, provider AccessTokenProvider, fallback, component string) (string, error) {
|
||||
@@ -381,6 +433,13 @@ func isRetryablePersonalError(err error) bool {
|
||||
func personalRetryLogError(err error) string {
|
||||
message := err.Error()
|
||||
switch {
|
||||
case strings.Contains(message, "resolve access token"), strings.Contains(message, "refresh rejected access token"):
|
||||
// Token resolution/refresh errors may carry provider details; log
|
||||
// only the structured HTTP status.
|
||||
if status := refreshHTTPStatus(err); status != 0 {
|
||||
return fmt.Sprintf("personal source: token refresh HTTP %d", status)
|
||||
}
|
||||
return "personal source: token refresh: temporary network error"
|
||||
case strings.Contains(message, "ticket HTTP"):
|
||||
return message
|
||||
case strings.Contains(message, "fetch ticket"):
|
||||
@@ -398,6 +457,14 @@ func personalRetryLogError(err error) string {
|
||||
}
|
||||
}
|
||||
|
||||
func refreshHTTPStatus(err error) int {
|
||||
var statusErr *authpkg.HTTPStatusError
|
||||
if !errors.As(err, &statusErr) || statusErr == nil {
|
||||
return 0
|
||||
}
|
||||
return statusErr.StatusCode
|
||||
}
|
||||
|
||||
func retryableTicketStatus(status int) bool {
|
||||
return status == http.StatusRequestTimeout ||
|
||||
status == http.StatusTooManyRequests ||
|
||||
|
||||
@@ -20,11 +20,13 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
|
||||
@@ -33,6 +35,8 @@ import (
|
||||
const (
|
||||
PortalTicketModeNormal = "normal"
|
||||
PortalTicketModeCustom = "custom"
|
||||
portalReconnectMin = time.Second
|
||||
portalReconnectMax = 30 * time.Second
|
||||
)
|
||||
|
||||
// PortalTicketConfig describes the portal-managed user Stream ticket flow.
|
||||
@@ -42,12 +46,45 @@ type PortalTicketConfig struct {
|
||||
TicketURL string
|
||||
AccessToken string
|
||||
AccessTokenProvider AccessTokenProvider
|
||||
ForceRefreshToken ForceRefreshTokenFn
|
||||
SourceID string
|
||||
Mode string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
UserAgent string
|
||||
HTTPClient *http.Client
|
||||
WebSocketDialer *websocket.Dialer
|
||||
ReconnectMin time.Duration
|
||||
ReconnectMax time.Duration
|
||||
DisableReconnect bool
|
||||
}
|
||||
|
||||
// portalStageError tags a portal stream failure with the stage it happened
|
||||
// in and whether the reconnect loop may retry it. Error() stays free of
|
||||
// untrusted response content so it is safe to log on every reconnect.
|
||||
type portalStageError struct {
|
||||
stage string
|
||||
status int
|
||||
retryable bool
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *portalStageError) Error() string {
|
||||
if e == nil {
|
||||
return "source: portal stream failed"
|
||||
}
|
||||
message := "source: portal " + strings.ReplaceAll(strings.TrimSpace(e.stage), "_", " ") + " failed"
|
||||
if e.status != 0 {
|
||||
message += fmt.Sprintf(" (HTTP %d)", e.status)
|
||||
}
|
||||
return message
|
||||
}
|
||||
|
||||
func (e *portalStageError) Unwrap() error {
|
||||
if e == nil {
|
||||
return nil
|
||||
}
|
||||
return e.cause
|
||||
}
|
||||
|
||||
var portalWriteMessage = func(conn *websocket.Conn, messageType int, data []byte) error {
|
||||
@@ -97,48 +134,102 @@ func normalizePortalTicketMode(mode string) string {
|
||||
|
||||
func (s *DingtalkSource) startPortalTicket(ctx context.Context, emit dwsevent.EmitFn) error {
|
||||
s.machine.OnConnecting()
|
||||
defer s.machine.OnStopped()
|
||||
|
||||
minBackoff := s.cfg.PortalTicket.ReconnectMin
|
||||
if minBackoff <= 0 {
|
||||
minBackoff = portalReconnectMin
|
||||
}
|
||||
maxBackoff := s.cfg.PortalTicket.ReconnectMax
|
||||
if maxBackoff <= 0 {
|
||||
maxBackoff = portalReconnectMax
|
||||
}
|
||||
if maxBackoff < minBackoff {
|
||||
maxBackoff = minBackoff
|
||||
}
|
||||
backoff := minBackoff
|
||||
for {
|
||||
acked, err := s.runPortalTicketAttempt(ctx, emit)
|
||||
if ctx.Err() != nil {
|
||||
return ctx.Err()
|
||||
}
|
||||
var stageErr *portalStageError
|
||||
if !errors.As(err, &stageErr) || stageErr == nil || !stageErr.retryable || s.cfg.PortalTicket.DisableReconnect {
|
||||
return err
|
||||
}
|
||||
if acked {
|
||||
backoff = minBackoff
|
||||
}
|
||||
s.machine.OnReconnect()
|
||||
slog.Warn("portal source reconnecting",
|
||||
"stage", stageErr.stage,
|
||||
"http_status", stageErr.status,
|
||||
"error_type", fmt.Sprintf("%T", stageErr.cause),
|
||||
"retry_in", backoff,
|
||||
"reconnect_count", s.machine.Snapshot().ReconnectCount,
|
||||
)
|
||||
if err := waitPersonalReconnect(ctx, backoff); err != nil {
|
||||
return err
|
||||
}
|
||||
backoff = nextPersonalBackoff(backoff, maxBackoff)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *DingtalkSource) runPortalTicketAttempt(ctx context.Context, emit dwsevent.EmitFn) (bool, error) {
|
||||
ticket, err := requestPortalTicket(ctx, s.cfg.PortalTicket)
|
||||
if err != nil {
|
||||
s.machine.OnStopped()
|
||||
return err
|
||||
return false, err
|
||||
}
|
||||
wsURL, err := websocketURL(ticket)
|
||||
if err != nil {
|
||||
s.machine.OnStopped()
|
||||
return err
|
||||
return false, err
|
||||
}
|
||||
|
||||
userAgent := strings.TrimSpace(s.cfg.PortalTicket.UserAgent)
|
||||
if userAgent == "" {
|
||||
userAgent = "dws-event-consume"
|
||||
}
|
||||
conn, resp, err := (&websocket.Dialer{HandshakeTimeout: 20 * time.Second}).DialContext(ctx, wsURL, http.Header{
|
||||
dialer := s.cfg.PortalTicket.WebSocketDialer
|
||||
if dialer == nil {
|
||||
dialer = &websocket.Dialer{HandshakeTimeout: 20 * time.Second}
|
||||
}
|
||||
conn, resp, err := dialer.DialContext(ctx, wsURL, http.Header{
|
||||
"User-Agent": []string{userAgent},
|
||||
})
|
||||
if err != nil {
|
||||
s.machine.OnStopped()
|
||||
status := 0
|
||||
cause := fmt.Errorf("source: portal stream connect: %w", err)
|
||||
if resp != nil {
|
||||
defer resp.Body.Close()
|
||||
status = resp.StatusCode
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
return fmt.Errorf("source: portal stream connect HTTP %d: %s: %w",
|
||||
cause = fmt.Errorf("source: portal stream connect HTTP %d: %s: %w",
|
||||
resp.StatusCode, truncatePortalTicketLog(string(raw), 300), err)
|
||||
}
|
||||
return fmt.Errorf("source: portal stream connect: %w", err)
|
||||
return false, &portalStageError{
|
||||
stage: "stream_connect",
|
||||
status: status,
|
||||
retryable: status == 0 || retryableTicketStatus(status),
|
||||
cause: cause,
|
||||
}
|
||||
}
|
||||
defer conn.Close()
|
||||
attemptCtx, cancel := context.WithCancel(ctx)
|
||||
defer func() {
|
||||
cancel()
|
||||
_ = conn.Close()
|
||||
}()
|
||||
closeOnContext(attemptCtx, conn)
|
||||
s.machine.OnConnected()
|
||||
|
||||
closeOnContext(ctx, conn)
|
||||
handler := s.makeHandler(emit)
|
||||
acked := false
|
||||
for {
|
||||
_, message, err := conn.ReadMessage()
|
||||
if err != nil {
|
||||
s.machine.OnStopped()
|
||||
if isContextDone(ctx) {
|
||||
return ctx.Err()
|
||||
return acked, ctx.Err()
|
||||
}
|
||||
return fmt.Errorf("source: portal stream read: %w", err)
|
||||
return acked, &portalStageError{stage: "stream_read", retryable: true, cause: fmt.Errorf("source: portal stream read: %w", err)}
|
||||
}
|
||||
df, err := payload.DecodeDataFrame(message)
|
||||
if err != nil {
|
||||
@@ -147,12 +238,12 @@ func (s *DingtalkSource) startPortalTicket(ctx context.Context, emit dwsevent.Em
|
||||
resp, _ := handler(ctx, df)
|
||||
ensurePortalAckHeaders(resp, df)
|
||||
if err := portalWriteMessage(conn, websocket.TextMessage, resp.Encode()); err != nil {
|
||||
s.machine.OnStopped()
|
||||
if isContextDone(ctx) {
|
||||
return ctx.Err()
|
||||
return acked, ctx.Err()
|
||||
}
|
||||
return fmt.Errorf("source: portal stream ack: %w", err)
|
||||
return acked, &portalStageError{stage: "stream_ack", retryable: true, cause: fmt.Errorf("source: portal stream ack: %w", err)}
|
||||
}
|
||||
acked = true
|
||||
}
|
||||
}
|
||||
|
||||
@@ -164,12 +255,33 @@ type portalStreamTicket struct {
|
||||
func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalStreamTicket, error) {
|
||||
accessToken, err := resolveSourceAccessToken(ctx, cfg.AccessTokenProvider, cfg.AccessToken, "source: portal ticket")
|
||||
if err != nil {
|
||||
// Transient provider failures (network, 429, 5xx) must not kill a
|
||||
// long-running source; the reconnect loop retries after backoff.
|
||||
if authpkg.ClassifyRefreshFailure(err) == authpkg.RefreshFailureTransient {
|
||||
return portalStreamTicket{}, &portalStageError{stage: "ticket_auth", status: refreshHTTPStatus(err), retryable: true, cause: err}
|
||||
}
|
||||
return portalStreamTicket{}, err
|
||||
}
|
||||
httpClient := cfg.HTTPClient
|
||||
if httpClient == nil {
|
||||
httpClient = &http.Client{Timeout: 20 * time.Second}
|
||||
}
|
||||
ticket, status, err := requestPortalTicketAttempt(ctx, cfg, httpClient, accessToken)
|
||||
if status == http.StatusUnauthorized && cfg.ForceRefreshToken != nil {
|
||||
refreshed, refreshErr := refreshRejectedSourceToken(ctx, cfg.ForceRefreshToken, accessToken, "source: portal ticket", err)
|
||||
if refreshErr != nil {
|
||||
if authpkg.ClassifyRefreshFailure(refreshErr) == authpkg.RefreshFailureTransient {
|
||||
return portalStreamTicket{}, &portalStageError{stage: "ticket_auth_refresh", status: refreshHTTPStatus(refreshErr), retryable: true, cause: refreshErr}
|
||||
}
|
||||
return portalStreamTicket{}, refreshErr
|
||||
}
|
||||
// Retry once with the freshly rotated token; a second 401 stays fatal.
|
||||
ticket, _, err = requestPortalTicketAttempt(ctx, cfg, httpClient, refreshed)
|
||||
}
|
||||
return ticket, err
|
||||
}
|
||||
|
||||
func requestPortalTicketAttempt(ctx context.Context, cfg *PortalTicketConfig, httpClient *http.Client, accessToken string) (portalStreamTicket, int, error) {
|
||||
body := map[string]string{
|
||||
"sourceId": strings.TrimSpace(cfg.SourceID),
|
||||
"channelType": strings.TrimSpace(cfg.SourceID),
|
||||
@@ -182,7 +294,7 @@ func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalSt
|
||||
rawBody, _ := json.Marshal(body)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, strings.TrimSpace(cfg.TicketURL), bytes.NewReader(rawBody))
|
||||
if err != nil {
|
||||
return portalStreamTicket{}, err
|
||||
return portalStreamTicket{}, 0, err
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
@@ -193,18 +305,34 @@ func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalSt
|
||||
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
return portalStreamTicket{}, fmt.Errorf("source: portal ticket request: %w", err)
|
||||
return portalStreamTicket{}, 0, &portalStageError{stage: "ticket_request", retryable: true, cause: fmt.Errorf("source: portal ticket request: %w", err)}
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
if resp.StatusCode >= 400 {
|
||||
return portalStreamTicket{}, fmt.Errorf("source: portal ticket HTTP %d: %s",
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
// Preserve HTTP status semantics before touching the body. A truncated
|
||||
// 401 body must not become a retryable read error that bypasses the
|
||||
// single token-refresh guard; the body is only best-effort diagnostics.
|
||||
raw, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
httpErr := fmt.Errorf("source: portal ticket HTTP %d: %s",
|
||||
resp.StatusCode, truncatePortalTicketLog(string(raw), 300))
|
||||
if retryableTicketStatus(resp.StatusCode) {
|
||||
return portalStreamTicket{}, resp.StatusCode, &portalStageError{stage: "ticket_request", status: resp.StatusCode, retryable: true, cause: httpErr}
|
||||
}
|
||||
return portalStreamTicket{}, resp.StatusCode, httpErr
|
||||
}
|
||||
raw, readErr := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
|
||||
if readErr != nil {
|
||||
return portalStreamTicket{}, resp.StatusCode, &portalStageError{
|
||||
stage: "ticket_request",
|
||||
status: resp.StatusCode,
|
||||
retryable: true,
|
||||
cause: fmt.Errorf("source: portal ticket read: %w", readErr),
|
||||
}
|
||||
}
|
||||
|
||||
var direct portalStreamTicket
|
||||
if err := json.Unmarshal(raw, &direct); err == nil && direct.Endpoint != "" && direct.Ticket != "" {
|
||||
return direct, nil
|
||||
return direct, resp.StatusCode, nil
|
||||
}
|
||||
|
||||
var envelope struct {
|
||||
@@ -214,16 +342,16 @@ func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalSt
|
||||
ErrorMsg string `json:"errorMsg"`
|
||||
}
|
||||
if err := json.Unmarshal(raw, &envelope); err != nil {
|
||||
return portalStreamTicket{}, fmt.Errorf("source: portal ticket parse: %w", err)
|
||||
return portalStreamTicket{}, resp.StatusCode, fmt.Errorf("source: portal ticket parse: %w", err)
|
||||
}
|
||||
if !envelope.Success {
|
||||
return portalStreamTicket{}, fmt.Errorf("source: portal ticket failed: %s %s",
|
||||
return portalStreamTicket{}, resp.StatusCode, fmt.Errorf("source: portal ticket failed: %s %s",
|
||||
envelope.ErrorCode, envelope.ErrorMsg)
|
||||
}
|
||||
if envelope.Result.Endpoint == "" || envelope.Result.Ticket == "" {
|
||||
return portalStreamTicket{}, errors.New("source: portal ticket result missing endpoint/ticket")
|
||||
return portalStreamTicket{}, resp.StatusCode, errors.New("source: portal ticket result missing endpoint/ticket")
|
||||
}
|
||||
return envelope.Result, nil
|
||||
return envelope.Result, resp.StatusCode, nil
|
||||
}
|
||||
|
||||
func websocketURL(ticket portalStreamTicket) (string, error) {
|
||||
|
||||
@@ -0,0 +1,485 @@
|
||||
package source
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
|
||||
)
|
||||
|
||||
// TestCrossPlatformCoveragePortalStart401RefreshRetryEndToEnd drives the full production chain
|
||||
// DingtalkSource.Start → startPortalTicket → requestPortalTicket: the first
|
||||
// ticket request is rejected with 401, ForceRefreshToken rotates the token,
|
||||
// the in-chain retry succeeds with the fresh token and a WebSocket event is
|
||||
// delivered to emit.
|
||||
func TestCrossPlatformCoveragePortalStart401RefreshRetryEndToEnd(t *testing.T) {
|
||||
var ticketCalls, refreshCalls atomic.Int64
|
||||
var rejectedSeen atomic.Value
|
||||
|
||||
upgrader := websocket.Upgrader{}
|
||||
var wsURL string
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/ticket", func(w http.ResponseWriter, r *http.Request) {
|
||||
ticketCalls.Add(1)
|
||||
switch r.Header.Get("x-user-access-token") {
|
||||
case "fresh-token":
|
||||
_ = json.NewEncoder(w).Encode(map[string]any{
|
||||
"success": true,
|
||||
"result": map[string]string{"endpoint": wsURL, "ticket": "ticket-1"},
|
||||
})
|
||||
default:
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
_, _ = io.WriteString(w, "token expired")
|
||||
}
|
||||
})
|
||||
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
df := payload.DataFrame{Type: "event", Headers: payload.DataFrameHeader{payload.DataFrameHeaderKMessageId: "msg-1"}, Data: `{}`}
|
||||
_ = conn.WriteJSON(df)
|
||||
_, _, _ = conn.ReadMessage()
|
||||
})
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
wsURL = "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
|
||||
|
||||
s, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: srv.URL + "/ticket",
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
return "stale-token", nil
|
||||
},
|
||||
ForceRefreshToken: func(_ context.Context, rejectedToken string) (string, error) {
|
||||
refreshCalls.Add(1)
|
||||
rejectedSeen.Store(rejectedToken)
|
||||
return "fresh-token", nil
|
||||
},
|
||||
SourceID: "source",
|
||||
HTTPClient: srv.Client(),
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
emitted := make(chan struct{}, 1)
|
||||
done := make(chan error, 1)
|
||||
go func() { done <- s.Start(ctx, func(*dwsevent.RawEvent) { emitted <- struct{}{} }) }()
|
||||
select {
|
||||
case <-emitted:
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("portal event timeout after 401 refresh retry")
|
||||
}
|
||||
cancel()
|
||||
select {
|
||||
case err := <-done:
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("portal stop = %v", err)
|
||||
}
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("portal stop timeout")
|
||||
}
|
||||
|
||||
if got := ticketCalls.Load(); got != 2 {
|
||||
t.Fatalf("ticket calls = %d, want 2", got)
|
||||
}
|
||||
if got := refreshCalls.Load(); got != 1 {
|
||||
t.Fatalf("refresh calls = %d, want 1", got)
|
||||
}
|
||||
if got, _ := rejectedSeen.Load().(string); got != "stale-token" {
|
||||
t.Fatalf("rejected token = %q, want %q", got, "stale-token")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoverageRequestPortalTicketRetryUsesRotatedTokenDirectly asserts the in-chain
|
||||
// retry sends the token returned by ForceRefreshToken instead of re-invoking
|
||||
// the provider (which could still serve the stale token).
|
||||
func TestCrossPlatformCoverageRequestPortalTicketRetryUsesRotatedTokenDirectly(t *testing.T) {
|
||||
providerCalls := 0
|
||||
var attemptTokens []string
|
||||
client := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
token := req.Header.Get("x-user-access-token")
|
||||
attemptTokens = append(attemptTokens, token)
|
||||
if token != "rotated" {
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"endpoint":"wss://x","ticket":"t"}`)), Header: make(http.Header)}, nil
|
||||
})}
|
||||
ticket, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://x",
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
providerCalls++
|
||||
return "stale", nil
|
||||
},
|
||||
ForceRefreshToken: func(_ context.Context, rejectedToken string) (string, error) {
|
||||
if rejectedToken != "stale" {
|
||||
t.Fatalf("rejected token = %q, want %q", rejectedToken, "stale")
|
||||
}
|
||||
return "rotated", nil
|
||||
},
|
||||
SourceID: "s",
|
||||
HTTPClient: client,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("requestPortalTicket = %v", err)
|
||||
}
|
||||
if ticket.Endpoint != "wss://x" || ticket.Ticket != "t" {
|
||||
t.Fatalf("ticket = %#v", ticket)
|
||||
}
|
||||
if providerCalls != 1 {
|
||||
t.Fatalf("provider calls = %d, want 1", providerCalls)
|
||||
}
|
||||
if len(attemptTokens) != 2 || attemptTokens[0] != "stale" || attemptTokens[1] != "rotated" {
|
||||
t.Fatalf("attempt tokens = %v", attemptTokens)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoverageRequestPortalTicketRefreshFailureKeepsBothErrors asserts a failing
|
||||
// refresh neither retries nor drops the refresh error or the original 401.
|
||||
func TestCrossPlatformCoverageRequestPortalTicketRefreshFailureKeepsBothErrors(t *testing.T) {
|
||||
refreshErr := errors.New("refresh_token exchange failed")
|
||||
attempts := 0
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
attempts++
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
})}
|
||||
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://x",
|
||||
AccessToken: "stale",
|
||||
ForceRefreshToken: func(context.Context, string) (string, error) {
|
||||
return "", refreshErr
|
||||
},
|
||||
SourceID: "s",
|
||||
HTTPClient: client,
|
||||
})
|
||||
if !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("error should wrap refresh error, got %v", err)
|
||||
}
|
||||
if err == nil || !strings.Contains(err.Error(), "HTTP 401") {
|
||||
t.Fatalf("error should keep original 401, got %v", err)
|
||||
}
|
||||
if attempts != 1 {
|
||||
t.Fatalf("attempts = %d, want 1 (no retry after failed refresh)", attempts)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoverageRequestPortalTicketWithoutRefreshCallback401StaysFatal covers backward
|
||||
// compatibility: nil ForceRefreshToken keeps the single-attempt fatal 401.
|
||||
func TestCrossPlatformCoverageRequestPortalTicketWithoutRefreshCallback401StaysFatal(t *testing.T) {
|
||||
attempts := 0
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
attempts++
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
})}
|
||||
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://x", AccessToken: "stale", SourceID: "s", HTTPClient: client,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "HTTP 401") {
|
||||
t.Fatalf("fatal 401 expected, got %v", err)
|
||||
}
|
||||
if attempts != 1 {
|
||||
t.Fatalf("attempts = %d, want 1", attempts)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoverageRequestPortalTicketSecond401IsFatal guards against refresh loops: the
|
||||
// controlled retry happens exactly once even if the rotated token is also
|
||||
// rejected.
|
||||
func TestCrossPlatformCoverageRequestPortalTicketSecond401IsFatal(t *testing.T) {
|
||||
attempts := 0
|
||||
refreshCalls := 0
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
attempts++
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
})}
|
||||
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://x",
|
||||
AccessToken: "stale",
|
||||
ForceRefreshToken: func(context.Context, string) (string, error) {
|
||||
refreshCalls++
|
||||
return "rotated-but-still-rejected", nil
|
||||
},
|
||||
SourceID: "s",
|
||||
HTTPClient: client,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "HTTP 401") {
|
||||
t.Fatalf("fatal 401 expected after single retry, got %v", err)
|
||||
}
|
||||
if attempts != 2 || refreshCalls != 1 {
|
||||
t.Fatalf("attempts = %d refreshCalls = %d, want 2/1", attempts, refreshCalls)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoverageRequestPortalTicketRefreshEmptyTokenIsFatal asserts an empty rotated
|
||||
// token is rejected instead of being sent to the server.
|
||||
func TestCrossPlatformCoverageRequestPortalTicketRefreshEmptyTokenIsFatal(t *testing.T) {
|
||||
attempts := 0
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
attempts++
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
})}
|
||||
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://x",
|
||||
AccessToken: "stale",
|
||||
ForceRefreshToken: func(context.Context, string) (string, error) {
|
||||
return " ", nil
|
||||
},
|
||||
SourceID: "s",
|
||||
HTTPClient: client,
|
||||
})
|
||||
if err == nil || !strings.Contains(err.Error(), "empty token") {
|
||||
t.Fatalf("empty rotated token error expected, got %v", err)
|
||||
}
|
||||
if attempts != 1 {
|
||||
t.Fatalf("attempts = %d, want 1", attempts)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoveragePersonalFetchTicket401RefreshRetry mirrors the portal behavior for the
|
||||
// personal stream ticket path.
|
||||
func TestCrossPlatformCoveragePersonalFetchTicket401RefreshRetry(t *testing.T) {
|
||||
var attemptTokens []string
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessTokenProvider: func(context.Context) (string, error) { return "stale", nil },
|
||||
ForceRefreshToken: func(_ context.Context, rejectedToken string) (string, error) {
|
||||
if rejectedToken != "stale" {
|
||||
t.Fatalf("rejected token = %q, want %q", rejectedToken, "stale")
|
||||
}
|
||||
return "rotated", nil
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
|
||||
token := req.Header.Get("x-user-access-token")
|
||||
attemptTokens = append(attemptTokens, token)
|
||||
if token != "rotated" {
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
}
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"endpoint":"wss://stream.test","ticket":"ticket"}`)), Header: make(http.Header)}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
ticket, err := src.fetchTicket(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("fetchTicket = %v", err)
|
||||
}
|
||||
if ticket.Endpoint != "wss://stream.test" || ticket.Ticket != "ticket" {
|
||||
t.Fatalf("ticket = %#v", ticket)
|
||||
}
|
||||
if len(attemptTokens) != 2 || attemptTokens[0] != "stale" || attemptTokens[1] != "rotated" {
|
||||
t.Fatalf("attempt tokens = %v", attemptTokens)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoveragePersonalFetchTicket401RefreshFailureStaysFatal asserts a failed refresh
|
||||
// keeps the 401 fatal (not retryable) and wraps the refresh error.
|
||||
func TestCrossPlatformCoveragePersonalFetchTicket401RefreshFailureStaysFatal(t *testing.T) {
|
||||
refreshErr := errors.New("refresh_token exchange failed")
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessToken: "stale",
|
||||
ForceRefreshToken: func(context.Context, string) (string, error) {
|
||||
return "", refreshErr
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: io.NopCloser(strings.NewReader("expired")), Header: make(http.Header)}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = src.fetchTicket(context.Background())
|
||||
if !errors.Is(err, refreshErr) {
|
||||
t.Fatalf("error should wrap refresh error, got %v", err)
|
||||
}
|
||||
if isRetryablePersonalError(err) {
|
||||
t.Fatalf("failed refresh should stay fatal, got retryable %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// brokenBody simulates a response body that fails mid-read, e.g. the server
|
||||
// closing the connection before the error payload is fully written.
|
||||
type brokenBody struct{}
|
||||
|
||||
func (brokenBody) Read([]byte) (int, error) { return 0, io.ErrUnexpectedEOF }
|
||||
func (brokenBody) Close() error { return nil }
|
||||
|
||||
func TestCrossPlatformCoveragePortalTicketNon2xxTruncatedBodyKeepsStatus(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
status int
|
||||
retryable bool
|
||||
}{
|
||||
{status: http.StatusUnauthorized},
|
||||
{status: http.StatusTooManyRequests, retryable: true},
|
||||
{status: http.StatusServiceUnavailable, retryable: true},
|
||||
} {
|
||||
t.Run(http.StatusText(tc.status), func(t *testing.T) {
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: tc.status, Body: brokenBody{}, Header: make(http.Header)}, nil
|
||||
})}
|
||||
|
||||
_, status, err := requestPortalTicketAttempt(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://ticket.test",
|
||||
SourceID: "source",
|
||||
}, client, "token")
|
||||
if status != tc.status || err == nil || !strings.Contains(err.Error(), fmt.Sprintf("HTTP %d", tc.status)) {
|
||||
t.Fatalf("requestPortalTicketAttempt() status=%d err=%v, want HTTP %d", status, err, tc.status)
|
||||
}
|
||||
if errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
t.Fatalf("non-2xx status error should not expose diagnostic body read failure: %v", err)
|
||||
}
|
||||
var stageErr *portalStageError
|
||||
if tc.retryable {
|
||||
if !errors.As(err, &stageErr) || stageErr.status != tc.status || !stageErr.retryable {
|
||||
t.Fatalf("retryable status should return retryable portalStageError, got %T %v", err, err)
|
||||
}
|
||||
} else if errors.As(err, &stageErr) && stageErr.retryable {
|
||||
t.Fatalf("fatal status should not become retryable stage error: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalTicket200TruncatedBodyIsRetryable(t *testing.T) {
|
||||
client := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: http.StatusOK, Body: brokenBody{}, Header: make(http.Header)}, nil
|
||||
})}
|
||||
|
||||
_, status, err := requestPortalTicketAttempt(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://ticket.test",
|
||||
SourceID: "source",
|
||||
}, client, "token")
|
||||
if status != http.StatusOK || !errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
t.Fatalf("requestPortalTicketAttempt() status=%d err=%v, want 200 + io.ErrUnexpectedEOF", status, err)
|
||||
}
|
||||
var stageErr *portalStageError
|
||||
if !errors.As(err, &stageErr) || !stageErr.retryable || stageErr.stage != "ticket_request" {
|
||||
t.Fatalf("2xx truncated body should be retryable ticket_request stage, got %T %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoveragePersonalFetchTicket401TruncatedBodyStaysFatal guards the single
|
||||
// refresh-retry protection: a 401 whose body fails with unexpected EOF must
|
||||
// be classified by status (fatal) and never wrapped as retryable, otherwise
|
||||
// the outer reconnect loop would refresh again on every iteration.
|
||||
func TestCrossPlatformCoveragePersonalFetchTicket401TruncatedBodyStaysFatal(t *testing.T) {
|
||||
attempts := 0
|
||||
refreshCalls := 0
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessToken: "stale",
|
||||
ForceRefreshToken: func(context.Context, string) (string, error) {
|
||||
refreshCalls++
|
||||
return "rotated", nil
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
attempts++
|
||||
return &http.Response{StatusCode: http.StatusUnauthorized, Body: brokenBody{}, Header: make(http.Header)}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = src.fetchTicket(context.Background())
|
||||
if err == nil || !strings.Contains(err.Error(), "HTTP 401") {
|
||||
t.Fatalf("fatal 401 error expected, got %v", err)
|
||||
}
|
||||
if isRetryablePersonalError(err) {
|
||||
t.Fatalf("401 with truncated body must stay fatal, got retryable %v", err)
|
||||
}
|
||||
if refreshCalls != 1 {
|
||||
t.Fatalf("refresh calls = %d, want 1", refreshCalls)
|
||||
}
|
||||
if attempts != 2 {
|
||||
t.Fatalf("attempts = %d, want 2", attempts)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoveragePersonalFetchTicket200TruncatedBodyStaysRetryable pins the existing
|
||||
// behavior for success responses: a body read failure on 2xx is a transient
|
||||
// transport problem and remains retryable.
|
||||
func TestCrossPlatformCoveragePersonalFetchTicket200TruncatedBodyStaysRetryable(t *testing.T) {
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessToken: "token",
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: 200, Body: brokenBody{}, Header: make(http.Header)}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = src.fetchTicket(context.Background())
|
||||
if !errors.Is(err, io.ErrUnexpectedEOF) {
|
||||
t.Fatalf("error should wrap io.ErrUnexpectedEOF, got %v", err)
|
||||
}
|
||||
if !isRetryablePersonalError(err) {
|
||||
t.Fatalf("2xx body read failure should stay retryable, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCrossPlatformCoveragePersonalFetchTicketAttemptTransportAndPayloadEdges covers the
|
||||
// remaining fetchTicketAttempt branches: transport failures and retryable
|
||||
// statuses stay retryable, while a well-formed response missing the endpoint
|
||||
// or ticket fields stays fatal.
|
||||
func TestCrossPlatformCoveragePersonalFetchTicketAttemptTransportAndPayloadEdges(t *testing.T) {
|
||||
newSource := func(rt roundTripFunc) *PersonalSource {
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessToken: "token",
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: rt},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return src
|
||||
}
|
||||
|
||||
dialErr := errors.New("dial tcp: connection refused")
|
||||
src := newSource(func(*http.Request) (*http.Response, error) { return nil, dialErr })
|
||||
_, _, err := src.fetchTicketAttempt(context.Background(), "token")
|
||||
if !errors.Is(err, dialErr) || !isRetryablePersonalError(err) {
|
||||
t.Fatalf("transport failure should stay retryable and wrap cause, got %v", err)
|
||||
}
|
||||
|
||||
src = newSource(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: http.StatusServiceUnavailable, Body: io.NopCloser(strings.NewReader("busy")), Header: make(http.Header)}, nil
|
||||
})
|
||||
_, status, err := src.fetchTicketAttempt(context.Background(), "token")
|
||||
if status != http.StatusServiceUnavailable || !isRetryablePersonalError(err) {
|
||||
t.Fatalf("503 should stay retryable, got status %d err %v", status, err)
|
||||
}
|
||||
|
||||
src = newSource(func(*http.Request) (*http.Response, error) {
|
||||
return &http.Response{StatusCode: 200, Body: io.NopCloser(strings.NewReader(`{"endpoint":"","ticket":""}`)), Header: make(http.Header)}, nil
|
||||
})
|
||||
_, _, err = src.fetchTicketAttempt(context.Background(), "token")
|
||||
if err == nil || isRetryablePersonalError(err) || !strings.Contains(err.Error(), "missing endpoint or ticket") {
|
||||
t.Fatalf("missing ticket fields should stay fatal, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,295 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
|
||||
package source
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
dwsevent "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/event"
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/open-dingtalk/dingtalk-stream-sdk-go/payload"
|
||||
)
|
||||
|
||||
func transientRefreshError() error {
|
||||
return &authpkg.HTTPStatusError{StatusCode: http.StatusServiceUnavailable}
|
||||
}
|
||||
|
||||
func terminalRefreshError() error {
|
||||
return &authpkg.HTTPStatusError{StatusCode: http.StatusUnauthorized}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalRetryLogErrorReportsOnlySafeTransientStatus(t *testing.T) {
|
||||
cause := errors.New("personal source: ticket HTTP 401 secret detail")
|
||||
_, err := refreshRejectedSourceToken(context.Background(), func(context.Context, string) (string, error) {
|
||||
return "", transientRefreshError()
|
||||
}, "rejected", "personal source", cause)
|
||||
if got, want := personalRetryLogError(retryPersonal(err)), "personal source: token refresh HTTP 503"; got != want {
|
||||
t.Fatalf("personalRetryLogError() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalSourceRetriesTransientTokenResolutionFailure(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var calls atomic.Int32
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
if calls.Add(1) == 2 {
|
||||
cancel()
|
||||
}
|
||||
return "", transientRefreshError()
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "open",
|
||||
TicketURL: "https://ticket.invalid",
|
||||
ReconnectMin: time.Millisecond,
|
||||
ReconnectMax: time.Millisecond,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(ctx, func(*dwsevent.RawEvent) {})
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Start() error = %v, want context canceled after retry", err)
|
||||
}
|
||||
if calls.Load() != 2 || src.State().ReconnectCount != 1 {
|
||||
t.Fatalf("provider calls=%d reconnects=%d, want 2 calls and 1 reconnect", calls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalSourceDoesNotRetryTerminalTokenResolutionFailure(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
calls.Add(1)
|
||||
return "", terminalRefreshError()
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "open",
|
||||
TicketURL: "https://ticket.invalid",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(context.Background(), func(*dwsevent.RawEvent) {})
|
||||
if authpkg.ClassifyRefreshFailure(err) != authpkg.RefreshFailureTerminal {
|
||||
t.Fatalf("Start() error = %v, want terminal refresh failure", err)
|
||||
}
|
||||
if calls.Load() != 1 || src.State().ReconnectCount != 0 {
|
||||
t.Fatalf("provider calls=%d reconnects=%d, want 1 call and no reconnect", calls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalSourceRetriesTransientRejectedTokenRefresh(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var ticketCalls atomic.Int32
|
||||
var refreshCalls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
ticketCalls.Add(1)
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
src, err := NewPersonal(PersonalConfig{
|
||||
AccessTokenProvider: func(context.Context) (string, error) { return "old-token", nil },
|
||||
ForceRefreshToken: func(_ context.Context, rejected string) (string, error) {
|
||||
if rejected != "old-token" {
|
||||
t.Fatalf("rejected token = %q, want old-token", rejected)
|
||||
}
|
||||
if refreshCalls.Add(1) == 2 {
|
||||
cancel()
|
||||
}
|
||||
return "", transientRefreshError()
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "open",
|
||||
TicketURL: srv.URL,
|
||||
HTTPClient: srv.Client(),
|
||||
ReconnectMin: time.Millisecond,
|
||||
ReconnectMax: time.Millisecond,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(ctx, func(*dwsevent.RawEvent) {})
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Start() error = %v, want context canceled after retry", err)
|
||||
}
|
||||
if ticketCalls.Load() != 2 || refreshCalls.Load() != 2 || src.State().ReconnectCount != 1 {
|
||||
t.Fatalf("ticket calls=%d refresh calls=%d reconnects=%d", ticketCalls.Load(), refreshCalls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalRetryLogErrorFallsBackToNetworkMessage(t *testing.T) {
|
||||
err := retryPersonal(fmt.Errorf("personal source: resolve access token: %w", errors.New("dial tcp: lookup oauth.invalid")))
|
||||
if got, want := personalRetryLogError(err), "personal source: token refresh: temporary network error"; got != want {
|
||||
t.Fatalf("personalRetryLogError() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalStageErrorNilAndUnwrap(t *testing.T) {
|
||||
var nilErr *portalStageError
|
||||
if got, want := nilErr.Error(), "source: portal stream failed"; got != want {
|
||||
t.Fatalf("nil stage error = %q, want %q", got, want)
|
||||
}
|
||||
if nilErr.Unwrap() != nil {
|
||||
t.Fatal("nil stage error should unwrap to nil")
|
||||
}
|
||||
cause := errors.New("cause")
|
||||
stageErr := &portalStageError{stage: "stream_read", retryable: true, cause: cause}
|
||||
if !errors.Is(stageErr, cause) {
|
||||
t.Fatalf("stage error should unwrap to cause: %v", stageErr)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalSourceRetriesTransientTokenResolutionFailure(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var calls atomic.Int32
|
||||
src, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: "https://ticket.invalid",
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
if calls.Add(1) == 2 {
|
||||
cancel()
|
||||
}
|
||||
return "", transientRefreshError()
|
||||
},
|
||||
SourceID: "open",
|
||||
// Min above max exercises the reconnect clamp.
|
||||
ReconnectMin: 2 * time.Millisecond,
|
||||
ReconnectMax: time.Millisecond,
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(ctx, func(*dwsevent.RawEvent) {})
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Start() error = %v, want context canceled after retry", err)
|
||||
}
|
||||
if calls.Load() != 2 || src.State().ReconnectCount != 1 {
|
||||
t.Fatalf("provider calls=%d reconnects=%d, want 2 calls and 1 reconnect", calls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalSourceResetsBackoffAfterAckedAttempt(t *testing.T) {
|
||||
var ticketCalls atomic.Int32
|
||||
upgrader := websocket.Upgrader{}
|
||||
var wsURL string
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("/ticket", func(w http.ResponseWriter, _ *http.Request) {
|
||||
if ticketCalls.Add(1) > 1 {
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
_, _ = io.WriteString(w, `{"endpoint":`+strconvQuote(wsURL)+`,"ticket":"t"}`)
|
||||
})
|
||||
mux.HandleFunc("/ws", func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
defer conn.Close()
|
||||
df := payload.DataFrame{Type: "event", Headers: payload.DataFrameHeader{payload.DataFrameHeaderKMessageId: "m"}, Data: `{}`}
|
||||
_ = conn.WriteJSON(df)
|
||||
// Wait for the ACK, then close so the read fails retryably with an
|
||||
// acked attempt behind it, which resets the reconnect backoff.
|
||||
_, _, _ = conn.ReadMessage()
|
||||
})
|
||||
srv := httptest.NewServer(mux)
|
||||
defer srv.Close()
|
||||
wsURL = "ws" + strings.TrimPrefix(srv.URL, "http") + "/ws"
|
||||
|
||||
src, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: srv.URL + "/ticket",
|
||||
AccessToken: "t",
|
||||
SourceID: "open",
|
||||
HTTPClient: srv.Client(),
|
||||
ReconnectMin: time.Millisecond,
|
||||
ReconnectMax: time.Millisecond,
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var events atomic.Int32
|
||||
err = src.Start(context.Background(), func(*dwsevent.RawEvent) { events.Add(1) })
|
||||
if err == nil || !strings.Contains(err.Error(), "portal ticket HTTP 400") {
|
||||
t.Fatalf("Start() error = %v, want fatal ticket HTTP 400 after reconnect", err)
|
||||
}
|
||||
if events.Load() != 1 || ticketCalls.Load() != 2 || src.State().ReconnectCount != 1 {
|
||||
t.Fatalf("events=%d ticket calls=%d reconnects=%d, want 1/2/1", events.Load(), ticketCalls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalSourceDoesNotRetryTerminalTokenResolutionFailure(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
src, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: "https://ticket.invalid",
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
calls.Add(1)
|
||||
return "", terminalRefreshError()
|
||||
},
|
||||
SourceID: "open",
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(context.Background(), func(*dwsevent.RawEvent) {})
|
||||
if authpkg.ClassifyRefreshFailure(err) != authpkg.RefreshFailureTerminal {
|
||||
t.Fatalf("Start() error = %v, want terminal refresh failure", err)
|
||||
}
|
||||
if calls.Load() != 1 || src.State().ReconnectCount != 0 {
|
||||
t.Fatalf("provider calls=%d reconnects=%d, want 1 call and no reconnect", calls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalSourceRetriesTransientRejectedTokenRefresh(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
var ticketCalls atomic.Int32
|
||||
var refreshCalls atomic.Int32
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
ticketCalls.Add(1)
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer srv.Close()
|
||||
|
||||
src, err := New(Config{PortalTicket: &PortalTicketConfig{
|
||||
TicketURL: srv.URL,
|
||||
AccessTokenProvider: func(context.Context) (string, error) { return "old-token", nil },
|
||||
ForceRefreshToken: func(_ context.Context, rejected string) (string, error) {
|
||||
if rejected != "old-token" {
|
||||
t.Fatalf("rejected token = %q, want old-token", rejected)
|
||||
}
|
||||
if refreshCalls.Add(1) == 2 {
|
||||
cancel()
|
||||
}
|
||||
return "", transientRefreshError()
|
||||
},
|
||||
SourceID: "open",
|
||||
HTTPClient: srv.Client(),
|
||||
ReconnectMin: time.Millisecond,
|
||||
ReconnectMax: time.Millisecond,
|
||||
}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = src.Start(ctx, func(*dwsevent.RawEvent) {})
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Start() error = %v, want context canceled after retry", err)
|
||||
}
|
||||
if ticketCalls.Load() != 2 || refreshCalls.Load() != 2 || src.State().ReconnectCount != 1 {
|
||||
t.Fatalf("ticket calls=%d refresh calls=%d reconnects=%d", ticketCalls.Load(), refreshCalls.Load(), src.State().ReconnectCount)
|
||||
}
|
||||
}
|
||||
@@ -116,11 +116,14 @@ func NewWorkflowInvocation(legacyPath, workflowName string, steps []Invocation)
|
||||
func MergePayloads(jsonPayload, paramsPayload string, overrides map[string]any) (map[string]any, error) {
|
||||
merged := make(map[string]any)
|
||||
|
||||
for label, payload := range map[string]string{
|
||||
"--json": jsonPayload,
|
||||
"--params": paramsPayload,
|
||||
for _, payload := range []struct {
|
||||
label string
|
||||
raw string
|
||||
}{
|
||||
{label: "--json", raw: jsonPayload},
|
||||
{label: "--params", raw: paramsPayload},
|
||||
} {
|
||||
value, err := parseJSONObject(label, payload)
|
||||
value, err := parseJSONObject(payload.label, payload.raw)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -52,6 +52,10 @@ func TestCrossPlatformCoverageMergePayloadsAndToolCallRequest(t *testing.T) {
|
||||
if empty, err := MergePayloads(" ", "", nil); err != nil || len(empty) != 0 {
|
||||
t.Fatalf("empty MergePayloads() = %#v, %v", empty, err)
|
||||
}
|
||||
merged, err = MergePayloads(`{"same":"json"}`, `{"same":"params"}`, nil)
|
||||
if err != nil || merged["same"] != "params" {
|
||||
t.Fatalf("MergePayloads() precedence = %#v, %v; want --params to win", merged, err)
|
||||
}
|
||||
for _, input := range []string{`{`, `[]`, `null`} {
|
||||
if _, err := MergePayloads(input, "", nil); err == nil {
|
||||
t.Errorf("MergePayloads(%q) error = nil", input)
|
||||
|
||||
@@ -18,7 +18,7 @@ func TestCrossPlatformCoverageGenerateProductionAgentMetadataPipeline(t *testing
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
root := app.NewRootCommand()
|
||||
root := app.NewSchemaSourceRootCommand()
|
||||
if _, err := cli.ApplyEmbeddedManualSchemaHints(root); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -46,7 +46,7 @@ var (
|
||||
writeMetadataFileBytes = os.WriteFile
|
||||
writeMetadataJSON = writeJSON
|
||||
|
||||
newMetadataRoot = app.NewRootCommand
|
||||
newMetadataRoot = app.NewSchemaSourceRootCommand
|
||||
buildEffectiveMetadata = cli.BuildEffectiveCommandRegistry
|
||||
bindEffectiveMetadata = cli.BindEffectiveCommandRegistry
|
||||
loadSelectionMetadataHints = cli.LoadAgentHintsFromSelectionForValidation
|
||||
|
||||
@@ -48,7 +48,7 @@ func main() {
|
||||
fail(err)
|
||||
}
|
||||
|
||||
root := app.NewRootCommand()
|
||||
root := app.NewSchemaSourceRootCommand()
|
||||
if err := generateSchemaCatalog(root, resolvedSurfacePath, outputPath); err != nil {
|
||||
fail(err)
|
||||
}
|
||||
|
||||
@@ -131,7 +131,7 @@ func TestCrossPlatformCoverageGenerateSchemaCatalogFailureEdges(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageGenerateSchemaCatalogResolvesBuildExactlyOnce(t *testing.T) {
|
||||
root := app.NewRootCommand()
|
||||
root := app.NewSchemaSourceRootCommand()
|
||||
resolveCalls := 0
|
||||
resolvedRegistryHash := ""
|
||||
resolver := func(candidate *cobra.Command) (cli.ResolvedSchemaBuild, error) {
|
||||
|
||||
@@ -17,7 +17,7 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
newSmokeRoot = app.NewRootCommand
|
||||
newSmokeRoot = app.NewSchemaSourceRootCommand
|
||||
buildEffectiveSmokeRegistry = cli.BuildEffectiveCommandRegistry
|
||||
bindEffectiveSmokeRegistry = cli.BindEffectiveCommandRegistry
|
||||
buildSmokeRegistryData = buildSmokeRegistry
|
||||
|
||||
@@ -14,7 +14,9 @@
|
||||
package plugin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -24,6 +26,8 @@ import (
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/mcptypes"
|
||||
)
|
||||
|
||||
const maxPluginCLIOverlayBytes = 4 << 20
|
||||
|
||||
// UserContext holds the minimal user identity fields injected into
|
||||
// stdio plugin subprocesses via environment variables.
|
||||
type UserContext struct {
|
||||
@@ -111,23 +115,9 @@ func (p *Plugin) ToServerDescriptors() []mcptypes.ServerDescriptor {
|
||||
continue
|
||||
}
|
||||
|
||||
overlay := mcptypes.CLIOverlay{}
|
||||
if len(srv.CLI) > 0 {
|
||||
if err := json.Unmarshal(srv.CLI, &overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLIOverlay",
|
||||
"plugin", p.Manifest.Name,
|
||||
"server", key,
|
||||
"error", err,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// Ensure the overlay has an ID — fall back to server key.
|
||||
if overlay.ID == "" {
|
||||
overlay.ID = key
|
||||
}
|
||||
if overlay.Command == "" {
|
||||
overlay.Command = key
|
||||
overlay, ok := p.ResolveCLIOverlay(key)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
|
||||
source := "plugin"
|
||||
@@ -154,3 +144,73 @@ func (p *Plugin) ToServerDescriptors() []mcptypes.ServerDescriptor {
|
||||
}
|
||||
return descriptors
|
||||
}
|
||||
|
||||
// ResolveCLIOverlay resolves inline or external manifest CLI metadata exactly
|
||||
// once. External files are opened relative to the plugin root with os.Root so
|
||||
// absolute paths, parent traversal, and escaping symlinks fail closed.
|
||||
func (p *Plugin) ResolveCLIOverlay(serverKey string) (mcptypes.CLIOverlay, bool) {
|
||||
overlay := mcptypes.CLIOverlay{
|
||||
ID: serverKey,
|
||||
Command: serverKey,
|
||||
}
|
||||
server, ok := p.Manifest.MCPServers[serverKey]
|
||||
if !ok || len(server.CLI) == 0 {
|
||||
return overlay, true
|
||||
}
|
||||
|
||||
data := []byte(strings.TrimSpace(string(server.CLI)))
|
||||
if len(data) == 0 {
|
||||
return overlay, true
|
||||
}
|
||||
if data[0] == '"' {
|
||||
var relativePath string
|
||||
if err := json.Unmarshal(data, &relativePath); err != nil ||
|
||||
strings.TrimSpace(relativePath) == "" {
|
||||
slog.Warn("plugin: invalid external CLI overlay path",
|
||||
"plugin", p.Manifest.Name, "server", serverKey, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
root, err := os.OpenRoot(p.Root)
|
||||
if err != nil {
|
||||
slog.Warn("plugin: failed to open plugin root",
|
||||
"plugin", p.Manifest.Name, "server", serverKey, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
defer root.Close()
|
||||
file, err := root.Open(relativePath)
|
||||
if err != nil {
|
||||
slog.Warn("plugin: failed to open CLI overlay file",
|
||||
"plugin", p.Manifest.Name, "server", serverKey,
|
||||
"path", relativePath, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
defer file.Close()
|
||||
data, err = io.ReadAll(io.LimitReader(file, maxPluginCLIOverlayBytes+1))
|
||||
if err != nil || len(data) > maxPluginCLIOverlayBytes {
|
||||
slog.Warn("plugin: failed to read CLI overlay file",
|
||||
"plugin", p.Manifest.Name, "server", serverKey,
|
||||
"path", relativePath, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
}
|
||||
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(&overlay); err != nil {
|
||||
slog.Warn("plugin: failed to parse CLI overlay",
|
||||
"plugin", p.Manifest.Name, "server", serverKey, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
if err := decoder.Decode(&struct{}{}); err != io.EOF {
|
||||
slog.Warn("plugin: CLI overlay contains trailing JSON",
|
||||
"plugin", p.Manifest.Name, "server", serverKey, "error", err)
|
||||
return mcptypes.CLIOverlay{}, false
|
||||
}
|
||||
if strings.TrimSpace(overlay.ID) == "" {
|
||||
overlay.ID = serverKey
|
||||
}
|
||||
if strings.TrimSpace(overlay.Command) == "" {
|
||||
overlay.Command = serverKey
|
||||
}
|
||||
return overlay, true
|
||||
}
|
||||
|
||||
@@ -0,0 +1,143 @@
|
||||
// Copyright 2026 Alibaba Group
|
||||
// Licensed under the Apache License, Version 2.0
|
||||
|
||||
package plugin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestResolveCLIOverlayDefaultsAndInlineErrors(t *testing.T) {
|
||||
plugin := &Plugin{
|
||||
Root: t.TempDir(),
|
||||
Manifest: Manifest{MCPServers: map[string]*MCPServer{
|
||||
"empty": {CLI: nil},
|
||||
"whitespace": {CLI: json.RawMessage(" \n\t")},
|
||||
"fallback": {CLI: json.RawMessage(`{"id":" ","command":""}`)},
|
||||
"malformed": {CLI: json.RawMessage(`{`)},
|
||||
"unknown": {CLI: json.RawMessage(`{"id":"plugin","unknownField":true}`)},
|
||||
"unknownFlag": {CLI: json.RawMessage(`{
|
||||
"toolOverrides":{"tool":{"flags":{"value":{"unknownFlagField":true}}}}
|
||||
}`)},
|
||||
"trailing": {CLI: json.RawMessage(`{} {}`)},
|
||||
"legacyTools": {CLI: json.RawMessage(`{
|
||||
"tools":[{
|
||||
"name":"tool","cliName":"leaf","title":"Title","description":"Description",
|
||||
"isSensitive":true,"category":"read","hidden":true,
|
||||
"flags":{"value":{"alias":"value-alias","shorthand":"v"}}
|
||||
}]
|
||||
}`)},
|
||||
}},
|
||||
}
|
||||
|
||||
for _, serverKey := range []string{"missing", "empty", "whitespace", "fallback"} {
|
||||
t.Run(serverKey, func(t *testing.T) {
|
||||
overlay, ok := plugin.ResolveCLIOverlay(serverKey)
|
||||
if !ok || overlay.ID != serverKey || overlay.Command != serverKey {
|
||||
t.Fatalf("ResolveCLIOverlay(%q) = (%#v, %v), want default overlay", serverKey, overlay, ok)
|
||||
}
|
||||
})
|
||||
}
|
||||
legacy, ok := plugin.ResolveCLIOverlay("legacyTools")
|
||||
if !ok || len(legacy.Tools) != 1 || legacy.Tools[0].CLIName != "leaf" ||
|
||||
legacy.Tools[0].Flags["value"].Alias != "value-alias" {
|
||||
t.Fatalf("ResolveCLIOverlay(legacyTools) = (%#v, %v), want historical tool metadata", legacy, ok)
|
||||
}
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("malformed"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay(malformed) = (%#v, true), want failure", overlay)
|
||||
}
|
||||
for _, serverKey := range []string{"unknown", "unknownFlag", "trailing"} {
|
||||
if overlay, ok := plugin.ResolveCLIOverlay(serverKey); ok {
|
||||
t.Fatalf("ResolveCLIOverlay(%s) = (%#v, true), want strict failure", serverKey, overlay)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveCLIOverlayRejectsInvalidExternalFiles(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
plugin := &Plugin{
|
||||
Root: root,
|
||||
Manifest: Manifest{MCPServers: map[string]*MCPServer{
|
||||
"external": {},
|
||||
}},
|
||||
}
|
||||
server := plugin.Manifest.MCPServers["external"]
|
||||
|
||||
t.Run("malformed path JSON", func(t *testing.T) {
|
||||
server.CLI = json.RawMessage(`"unterminated`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want malformed path failure", overlay)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("blank path", func(t *testing.T) {
|
||||
server.CLI = json.RawMessage(`" "`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want blank path failure", overlay)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing root", func(t *testing.T) {
|
||||
plugin.Root = filepath.Join(root, "does-not-exist")
|
||||
server.CLI = json.RawMessage(`"overlay.json"`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want missing root failure", overlay)
|
||||
}
|
||||
plugin.Root = root
|
||||
})
|
||||
|
||||
t.Run("missing file", func(t *testing.T) {
|
||||
server.CLI = json.RawMessage(`"missing.json"`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want missing file failure", overlay)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("read error", func(t *testing.T) {
|
||||
if err := os.Mkdir(filepath.Join(root, "overlay-dir"), 0o700); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server.CLI = json.RawMessage(`"overlay-dir"`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want directory read failure", overlay)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("oversized file", func(t *testing.T) {
|
||||
path := filepath.Join(root, "oversized.json")
|
||||
if err := os.WriteFile(path, bytes.Repeat([]byte(" "), maxPluginCLIOverlayBytes+1), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server.CLI = json.RawMessage(`"oversized.json"`)
|
||||
if overlay, ok := plugin.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, true), want oversized file failure", overlay)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestResolveCLIOverlayExternalFileAppliesFallbackIdentity(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(root, "overlay.json"), []byte(`{
|
||||
"id":"",
|
||||
"command":" ",
|
||||
"aliases":["conference"]
|
||||
}`), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
plugin := &Plugin{
|
||||
Root: root,
|
||||
Manifest: Manifest{MCPServers: map[string]*MCPServer{
|
||||
"server-key": {CLI: json.RawMessage(`"overlay.json"`)},
|
||||
}},
|
||||
}
|
||||
|
||||
overlay, ok := plugin.ResolveCLIOverlay("server-key")
|
||||
if !ok || overlay.ID != "server-key" || overlay.Command != "server-key" ||
|
||||
len(overlay.Aliases) != 1 || overlay.Aliases[0] != "conference" {
|
||||
t.Fatalf("ResolveCLIOverlay() = (%#v, %v), want fallback identity with file metadata", overlay, ok)
|
||||
}
|
||||
}
|
||||
@@ -131,22 +131,19 @@ func TestCrossPlatformCoveragePluginConverterEdges(t *testing.T) {
|
||||
t.Fatalf("StdioClients(user) length = %d", len(got))
|
||||
}
|
||||
descriptors := p.ToServerDescriptors()
|
||||
if len(descriptors) != 2 {
|
||||
if len(descriptors) != 1 {
|
||||
t.Fatalf("descriptors length = %d", len(descriptors))
|
||||
}
|
||||
seenDefault := false
|
||||
seenOverlay := false
|
||||
for _, d := range descriptors {
|
||||
switch d.Key {
|
||||
case "http-default":
|
||||
seenDefault = d.CLI.ID == "http-default" && d.CLI.Command == "http-default" && d.HasCLIMeta
|
||||
case "http-overlay":
|
||||
seenOverlay = d.CLI.ID == "custom" && d.CLI.Command == "run" &&
|
||||
d.AuthHeaders["Authorization"] == "Bearer secret"
|
||||
}
|
||||
}
|
||||
if !seenDefault || !seenOverlay {
|
||||
t.Fatalf("descriptor defaults/overlay not covered: %#v", descriptors)
|
||||
if !seenOverlay {
|
||||
t.Fatalf("valid descriptor overlay not covered: %#v", descriptors)
|
||||
}
|
||||
dataDir := filepath.Join(filepath.Dir(filepath.Dir(root)), "data")
|
||||
if got := expandPluginVars("$"+"{DWS_PLUGIN_ROOT}|$"+"{DWS_PLUGIN_DATA}|$"+"{PLUGIN_TOKEN}", root); got != root+"|"+dataDir+"|secret" {
|
||||
@@ -300,6 +297,7 @@ func TestCrossPlatformCoverageManifestEdges(t *testing.T) {
|
||||
func TestCrossPlatformCoverageLoaderDiscoveryLifecycle(t *testing.T) {
|
||||
oldHome := pluginUserHomeDir
|
||||
t.Cleanup(func() { pluginUserHomeDir = oldHome })
|
||||
t.Setenv("DWS_CONFIG_DIR", "")
|
||||
home := t.TempDir()
|
||||
pluginUserHomeDir = func() (string, error) { return home, nil }
|
||||
defaultLoader := NewLoader("1.2.3")
|
||||
@@ -310,6 +308,12 @@ func TestCrossPlatformCoverageLoaderDiscoveryLifecycle(t *testing.T) {
|
||||
if got := NewLoader("dev").PluginsDir; got != filepath.Join(".dws", "plugins") {
|
||||
t.Fatalf("NewLoader error path = %q", got)
|
||||
}
|
||||
customConfig := t.TempDir()
|
||||
t.Setenv("DWS_CONFIG_DIR", customConfig)
|
||||
if got := NewLoader("dev").PluginsDir; got != filepath.Join(customConfig, "plugins") {
|
||||
t.Fatalf("NewLoader custom config path = %q", got)
|
||||
}
|
||||
t.Setenv("DWS_CONFIG_DIR", "")
|
||||
|
||||
root := t.TempDir()
|
||||
l := &Loader{PluginsDir: root, CLIVersion: "1.0.0"}
|
||||
|
||||
@@ -56,9 +56,13 @@ var (
|
||||
|
||||
// NewLoader creates a Loader with default paths.
|
||||
func NewLoader(cliVersion string) *Loader {
|
||||
home, _ := pluginUserHomeDir()
|
||||
configDir := strings.TrimSpace(os.Getenv("DWS_CONFIG_DIR"))
|
||||
if configDir == "" {
|
||||
home, _ := pluginUserHomeDir()
|
||||
configDir = filepath.Join(home, ".dws")
|
||||
}
|
||||
return &Loader{
|
||||
PluginsDir: filepath.Join(home, ".dws", "plugins"),
|
||||
PluginsDir: filepath.Join(configDir, "plugins"),
|
||||
CLIVersion: cliVersion,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -74,6 +74,56 @@ func TestParseManifest(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestResolveCLIOverlayExternalFileIsRootContained(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
overlayPath := filepath.Join(root, "overlay.json")
|
||||
if err := os.WriteFile(overlayPath, []byte(`{
|
||||
"id":"external-id",
|
||||
"command":"external-command",
|
||||
"toolOverrides":{"ping":{"cliName":"ping"}}
|
||||
}`), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loaded := &Plugin{
|
||||
Root: root,
|
||||
Manifest: Manifest{
|
||||
Name: "external-plugin",
|
||||
MCPServers: map[string]*MCPServer{
|
||||
"external": {
|
||||
Type: "streamable-http",
|
||||
Endpoint: "https://example.invalid/mcp",
|
||||
CLI: json.RawMessage(`"overlay.json"`),
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
overlay, ok := loaded.ResolveCLIOverlay("external")
|
||||
if !ok || overlay.ID != "external-id" || overlay.Command != "external-command" {
|
||||
t.Fatalf("external overlay = (%#v, %v)", overlay, ok)
|
||||
}
|
||||
descriptors := loaded.ToServerDescriptors()
|
||||
if len(descriptors) != 1 || descriptors[0].CLI.ID != "external-id" {
|
||||
t.Fatalf("external HTTP descriptors = %#v", descriptors)
|
||||
}
|
||||
|
||||
outside := filepath.Join(t.TempDir(), "outside-overlay.json")
|
||||
if err := os.WriteFile(outside, []byte(`{"id":"escaped"}`), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
loaded.Manifest.MCPServers["external"].CLI = json.RawMessage(`"../outside-overlay.json"`)
|
||||
if overlay, ok := loaded.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("parent traversal resolved overlay %#v", overlay)
|
||||
}
|
||||
|
||||
link := filepath.Join(root, "escaped-link.json")
|
||||
if err := os.Symlink(outside, link); err == nil {
|
||||
loaded.Manifest.MCPServers["external"].CLI = json.RawMessage(`"escaped-link.json"`)
|
||||
if overlay, ok := loaded.ResolveCLIOverlay("external"); ok {
|
||||
t.Fatalf("escaping symlink resolved overlay %#v", overlay)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestManifestValidate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -219,8 +269,23 @@ func TestLoadAllSuppressesOptionalPluginValidationWarnings(t *testing.T) {
|
||||
|
||||
func TestPluginToServerDescriptors(t *testing.T) {
|
||||
cliOverlay, _ := json.Marshal(map[string]any{
|
||||
"id": "conference",
|
||||
"command": "conference",
|
||||
"id": "conference",
|
||||
"command": "conference",
|
||||
"description": "video conference",
|
||||
"groups": map[string]any{
|
||||
"camera": map[string]any{"description": "camera control"},
|
||||
},
|
||||
"toolOverrides": map[string]any{
|
||||
"open_camera": map[string]any{
|
||||
"cliName": "open",
|
||||
"group": "camera",
|
||||
"description": "open camera",
|
||||
"isSensitive": true,
|
||||
"flags": map[string]any{
|
||||
"device_id": map[string]any{"description": "camera device"},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
p := &Plugin{
|
||||
@@ -262,6 +327,16 @@ func TestPluginToServerDescriptors(t *testing.T) {
|
||||
if d.CLI.ID != "conference" {
|
||||
t.Errorf("cli.id = %q, want conference", d.CLI.ID)
|
||||
}
|
||||
override := d.CLI.ToolOverrides["open_camera"]
|
||||
if d.CLI.Description != "video conference" ||
|
||||
d.CLI.Groups["camera"].Description != "camera control" ||
|
||||
override.CLIName != "open" ||
|
||||
override.Group != "camera" ||
|
||||
override.Description != "open camera" ||
|
||||
!override.IsSensitive ||
|
||||
override.Flags["device_id"].Description != "camera device" {
|
||||
t.Fatalf("CLI overlay command metadata was not preserved: %#v", d.CLI)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPluginToServerDescriptorsWithHeaders(t *testing.T) {
|
||||
|
||||
@@ -0,0 +1,33 @@
|
||||
// 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 builtin_test
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/shortcut/builtin"
|
||||
)
|
||||
|
||||
func TestBaseCommandsExposeSortedDistributionShortcuts(t *testing.T) {
|
||||
commands := builtin.BaseCommands()
|
||||
if len(commands) == 0 {
|
||||
t.Fatal("BaseCommands returned no distribution shortcuts")
|
||||
}
|
||||
for index, command := range commands {
|
||||
if index > 0 && commands[index-1].Name() > command.Name() {
|
||||
t.Fatalf("BaseCommands are not sorted: %q before %q", commands[index-1].Name(), command.Name())
|
||||
}
|
||||
children := command.Commands()
|
||||
if len(children) == 0 {
|
||||
t.Fatalf("base shortcut service %q has no commands", command.Name())
|
||||
}
|
||||
for _, child := range children {
|
||||
if !strings.HasPrefix(child.Name(), "+") {
|
||||
t.Fatalf("base shortcut %q under %q is not mounted as a +command", child.Name(), command.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -50,3 +50,9 @@ import (
|
||||
func Commands() []*cobra.Command {
|
||||
return shortcut.Commands()
|
||||
}
|
||||
|
||||
// BaseCommands returns only distribution-owned shortcuts for Schema and
|
||||
// interface generation.
|
||||
func BaseCommands() []*cobra.Command {
|
||||
return shortcut.BuiltInCommands()
|
||||
}
|
||||
|
||||
@@ -41,6 +41,17 @@ func Commands() []*cobra.Command {
|
||||
return build(allShortcuts)
|
||||
}
|
||||
|
||||
// BuiltInCommands compiles only distribution-owned shortcuts.
|
||||
func BuiltInCommands() []*cobra.Command {
|
||||
builtins := make([]Shortcut, 0, len(allShortcuts))
|
||||
for _, registered := range allShortcuts {
|
||||
if !registered.UserDefined {
|
||||
builtins = append(builtins, registered)
|
||||
}
|
||||
}
|
||||
return build(builtins)
|
||||
}
|
||||
|
||||
// All returns the registered shortcuts. Primarily for coverage tests that need
|
||||
// each shortcut's declared flags (types/enums/required) to synthesize inputs.
|
||||
func All() []Shortcut {
|
||||
|
||||
@@ -266,6 +266,26 @@ func TestCrossPlatformCoverageBuildGroupsByService(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuiltInCommandsExcludeUserDefinedShortcuts(t *testing.T) {
|
||||
previous := append([]Shortcut(nil), allShortcuts...)
|
||||
t.Cleanup(func() { allShortcuts = previous })
|
||||
allShortcuts = nil
|
||||
Register(
|
||||
Shortcut{Service: "calendar", Command: "+builtin", Execute: noop},
|
||||
Shortcut{Service: "calendar", Command: "+user", UserDefined: true, Execute: noop},
|
||||
)
|
||||
|
||||
all := Commands()
|
||||
if len(all) != 1 || len(all[0].Commands()) != 2 {
|
||||
t.Fatalf("all shortcut commands = %#v", all)
|
||||
}
|
||||
builtins := BuiltInCommands()
|
||||
if len(builtins) != 1 || len(builtins[0].Commands()) != 1 ||
|
||||
builtins[0].Commands()[0].Name() != "+builtin" {
|
||||
t.Fatalf("built-in shortcut commands = %#v", builtins)
|
||||
}
|
||||
}
|
||||
|
||||
func noop(_ *RuntimeContext) error { return nil }
|
||||
|
||||
func TestCrossPlatformCoverageCallMCPWriteDataRejectsDryRun(t *testing.T) {
|
||||
|
||||
@@ -127,6 +127,10 @@ type Shortcut struct {
|
||||
Tips []string
|
||||
// Hidden hides the command from listings while keeping it invocable.
|
||||
Hidden bool
|
||||
// UserDefined identifies shortcuts loaded from the user's config
|
||||
// directory. Distribution-owned Schema and interface snapshots exclude
|
||||
// these runtime extensions even if another root loaded them earlier.
|
||||
UserDefined bool
|
||||
|
||||
// Validate optionally checks resolved flag values before execution. Return a
|
||||
// non-nil error to abort with a validation message. Runs after built-in
|
||||
|
||||
@@ -196,6 +196,7 @@ func Compile(s Spec) shortcut.Shortcut {
|
||||
Intent: intent,
|
||||
Risk: risk,
|
||||
Flags: flags,
|
||||
UserDefined: true,
|
||||
Execute: func(rt *shortcut.RuntimeContext) error {
|
||||
params := map[string]any{}
|
||||
for key, tmpl := range bind {
|
||||
|
||||
@@ -69,6 +69,9 @@ func TestCrossPlatformCoverageCompileFlagsAndDefaults(t *testing.T) {
|
||||
if sc.Service != "chat" || sc.Command != "+notify-team" {
|
||||
t.Fatalf("bad identity: %+v", sc)
|
||||
}
|
||||
if !sc.UserDefined {
|
||||
t.Fatal("compiled user shortcut is missing user-defined provenance")
|
||||
}
|
||||
if sc.Risk != shortcut.RiskRead {
|
||||
t.Errorf("risk default = %q, want read", sc.Risk)
|
||||
}
|
||||
|
||||
@@ -23,15 +23,28 @@ const SourceAnnotation = "dws.source"
|
||||
// SourceEnvelope marks a command as authored by the runtime discovery envelope.
|
||||
const SourceEnvelope = "envelope"
|
||||
|
||||
// SourcePlugin marks a command as an installed plugin extension. Plugin
|
||||
// commands are part of the runtime CLI surface, not the embedded base Schema.
|
||||
const SourcePlugin = "plugin"
|
||||
|
||||
// MarkEnvelopeSource stamps cmd with runtime discovery provenance.
|
||||
func MarkEnvelopeSource(cmd *cobra.Command) {
|
||||
markSource(cmd, SourceEnvelope)
|
||||
}
|
||||
|
||||
// MarkPluginSource stamps cmd with installed-plugin provenance.
|
||||
func MarkPluginSource(cmd *cobra.Command) {
|
||||
markSource(cmd, SourcePlugin)
|
||||
}
|
||||
|
||||
func markSource(cmd *cobra.Command, source string) {
|
||||
if cmd == nil {
|
||||
return
|
||||
}
|
||||
if cmd.Annotations == nil {
|
||||
cmd.Annotations = map[string]string{}
|
||||
}
|
||||
cmd.Annotations[SourceAnnotation] = SourceEnvelope
|
||||
cmd.Annotations[SourceAnnotation] = source
|
||||
}
|
||||
|
||||
// IsEnvelopeSourced reports whether cmd was authored by the runtime discovery
|
||||
@@ -40,6 +53,11 @@ func IsEnvelopeSourced(cmd *cobra.Command) bool {
|
||||
return cmd != nil && cmd.Annotations[SourceAnnotation] == SourceEnvelope
|
||||
}
|
||||
|
||||
// IsPluginSourced reports whether cmd came from an installed plugin.
|
||||
func IsPluginSourced(cmd *cobra.Command) bool {
|
||||
return cmd != nil && cmd.Annotations[SourceAnnotation] == SourcePlugin
|
||||
}
|
||||
|
||||
// KindAnnotation is the annotation key for marking command kinds.
|
||||
const KindAnnotation = "dws.kind"
|
||||
|
||||
|
||||
@@ -43,3 +43,19 @@ func TestMarkEnvelopeSourceNilDoesNotPanic(t *testing.T) {
|
||||
t.Parallel()
|
||||
MarkEnvelopeSource(nil)
|
||||
}
|
||||
|
||||
func TestPluginSourceProvenance(t *testing.T) {
|
||||
t.Parallel()
|
||||
if IsPluginSourced(nil) {
|
||||
t.Fatal("nil command should not be plugin sourced")
|
||||
}
|
||||
cmd := &cobra.Command{Use: "conference"}
|
||||
MarkPluginSource(cmd)
|
||||
if !IsPluginSourced(cmd) || IsEnvelopeSourced(cmd) {
|
||||
t.Fatalf("plugin source annotation = %#v", cmd.Annotations)
|
||||
}
|
||||
if got := cmd.Annotations[SourceAnnotation]; got != SourcePlugin {
|
||||
t.Fatalf("SourceAnnotation = %q, want %q", got, SourcePlugin)
|
||||
}
|
||||
MarkPluginSource(nil)
|
||||
}
|
||||
|
||||
+98
-3
@@ -1,6 +1,10 @@
|
||||
package mcptypes
|
||||
|
||||
import "encoding/json"
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
)
|
||||
|
||||
type ServerDescriptor struct {
|
||||
Key string
|
||||
@@ -16,19 +20,110 @@ type ServerDescriptor struct {
|
||||
type CLIOverlay struct {
|
||||
ID string `json:"id"`
|
||||
Command string `json:"command"`
|
||||
Parent string `json:"parent,omitempty"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Aliases []string `json:"aliases"`
|
||||
Prefixes []string `json:"prefixes"`
|
||||
Group string `json:"group,omitempty"`
|
||||
Skip bool `json:"skip"`
|
||||
Hidden bool `json:"hidden,omitempty"`
|
||||
Tools []CLITool `json:"tools"`
|
||||
Groups map[string]CLIGroupDef `json:"groups,omitempty"`
|
||||
ToolOverrides map[string]CLIToolOverride `json:"toolOverrides,omitempty"`
|
||||
ServerDeps []string `json:"serverDeps,omitempty"`
|
||||
Hints map[string]json.RawMessage `json:"hintCommands,omitempty"`
|
||||
RedirectTo string `json:"redirectTo,omitempty"`
|
||||
}
|
||||
|
||||
type CLIGroupDef struct {
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
|
||||
type CLITool struct {
|
||||
Name string `json:"name"`
|
||||
Name string `json:"name"`
|
||||
CLIName string `json:"cliName,omitempty"`
|
||||
Title string `json:"title,omitempty"`
|
||||
Description string `json:"description,omitempty"`
|
||||
IsSensitive bool `json:"isSensitive,omitempty"`
|
||||
Category string `json:"category,omitempty"`
|
||||
Hidden bool `json:"hidden,omitempty"`
|
||||
Flags map[string]CLIFlagHint `json:"flags,omitempty"`
|
||||
}
|
||||
|
||||
type CLIFlagHint struct {
|
||||
Shorthand string `json:"shorthand,omitempty"`
|
||||
Alias string `json:"alias,omitempty"`
|
||||
}
|
||||
|
||||
type CLIToolOverride struct {
|
||||
ServerOverride string `json:"serverOverride,omitempty"`
|
||||
CLIName string `json:"cliName,omitempty"`
|
||||
CLIAliases []string `json:"cliAliases,omitempty"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Example string `json:"example,omitempty"`
|
||||
Group string `json:"group,omitempty"`
|
||||
IsSensitive bool `json:"isSensitive,omitempty"`
|
||||
Hidden bool `json:"hidden,omitempty"`
|
||||
Flags map[string]CLIFlagOverride `json:"flags,omitempty"`
|
||||
OutputFormat map[string]any `json:"outputFormat,omitempty"`
|
||||
ServerOverride string `json:"serverOverride,omitempty"`
|
||||
BodyWrapper string `json:"bodyWrapper,omitempty"`
|
||||
MutuallyExclusive [][]string `json:"mutuallyExclusive,omitempty"`
|
||||
RequireOneOf [][]string `json:"requireOneOf,omitempty"`
|
||||
RequireTogether [][]string `json:"requireTogether,omitempty"`
|
||||
RejectPositional bool `json:"rejectPositional,omitempty"`
|
||||
RedirectTo string `json:"redirectTo,omitempty"`
|
||||
Pipeline []json.RawMessage `json:"pipeline,omitempty"`
|
||||
}
|
||||
|
||||
type CLIFlagOverride struct {
|
||||
Alias string `json:"alias,omitempty"`
|
||||
Aliases []string `json:"aliases,omitempty"`
|
||||
MapsTo string `json:"mapsTo,omitempty"`
|
||||
Transform string `json:"transform,omitempty"`
|
||||
TransformArgs map[string]any `json:"transformArgs,omitempty"`
|
||||
EnvDefault string `json:"envDefault,omitempty"`
|
||||
Hidden bool `json:"hidden,omitempty"`
|
||||
Default string `json:"default,omitempty"`
|
||||
Shorthand string `json:"shorthand,omitempty"`
|
||||
Required bool `json:"required,omitempty"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Positional bool `json:"positional,omitempty"`
|
||||
PositionalIndex int `json:"positionalIndex,omitempty"`
|
||||
Type string `json:"type,omitempty"`
|
||||
OmitWhen string `json:"omitWhen,omitempty"`
|
||||
RuntimeDefault string `json:"runtimeDefault,omitempty"`
|
||||
PipelineLocal bool `json:"pipelineLocal,omitempty"`
|
||||
}
|
||||
|
||||
func (override *CLIFlagOverride) UnmarshalJSON(data []byte) error {
|
||||
type alias CLIFlagOverride
|
||||
aux := struct {
|
||||
Default json.RawMessage `json:"default,omitempty"`
|
||||
*alias
|
||||
}{alias: (*alias)(override)}
|
||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
||||
decoder.DisallowUnknownFields()
|
||||
if err := decoder.Decode(&aux); err != nil {
|
||||
return err
|
||||
}
|
||||
override.Default = coercePluginScalar(aux.Default)
|
||||
return nil
|
||||
}
|
||||
|
||||
func coercePluginScalar(raw json.RawMessage) string {
|
||||
value := strings.TrimSpace(string(raw))
|
||||
if value == "" || value == "null" ||
|
||||
strings.HasPrefix(value, "{") || strings.HasPrefix(value, "[") {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(value, `"`) {
|
||||
var decoded string
|
||||
if json.Unmarshal(raw, &decoded) == nil {
|
||||
return decoded
|
||||
}
|
||||
return ""
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func OverlayFromJSON(data json.RawMessage) CLIOverlay {
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package mcptypes
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCLIFlagOverrideUnmarshalJSONCoercesScalarDefaults(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
want string
|
||||
}{
|
||||
{name: "missing", raw: `{}`, want: ""},
|
||||
{name: "null", raw: `{"default":null}`, want: ""},
|
||||
{name: "object", raw: `{"default":{"nested":true}}`, want: ""},
|
||||
{name: "array", raw: `{"default":[1,2]}`, want: ""},
|
||||
{name: "string", raw: `{"default":"hello"}`, want: "hello"},
|
||||
{name: "escaped string", raw: `{"default":"line\nvalue"}`, want: "line\nvalue"},
|
||||
{name: "boolean", raw: `{"default":true}`, want: "true"},
|
||||
{name: "integer", raw: `{"default":42}`, want: "42"},
|
||||
{name: "number", raw: `{"default":-1.25}`, want: "-1.25"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
var got CLIFlagOverride
|
||||
if err := json.Unmarshal([]byte(tt.raw), &got); err != nil {
|
||||
t.Fatalf("json.Unmarshal() error = %v", err)
|
||||
}
|
||||
if got.Default != tt.want {
|
||||
t.Fatalf("Default = %q, want %q", got.Default, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
var got CLIFlagOverride
|
||||
if err := json.Unmarshal([]byte(`{"alias":`), &got); err == nil {
|
||||
t.Fatal("json.Unmarshal() error = nil, want malformed JSON error")
|
||||
}
|
||||
if err := got.UnmarshalJSON([]byte(`{"alias":`)); err == nil {
|
||||
t.Fatal("CLIFlagOverride.UnmarshalJSON() error = nil, want malformed JSON error")
|
||||
}
|
||||
if err := json.Unmarshal([]byte(`{"unknownFlagField":true}`), &got); err == nil {
|
||||
t.Fatal("CLIFlagOverride.UnmarshalJSON() accepted an unknown field")
|
||||
}
|
||||
if got := coercePluginScalar(json.RawMessage(`"unterminated`)); got != "" {
|
||||
t.Fatalf("coercePluginScalar(invalid string) = %q, want empty", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOverlayFromJSONHandlesEmptyValidAndMalformedInput(t *testing.T) {
|
||||
if got := OverlayFromJSON(nil); got.ID != "" || got.Command != "" {
|
||||
t.Fatalf("OverlayFromJSON(nil) = %#v, want zero overlay", got)
|
||||
}
|
||||
|
||||
valid := json.RawMessage(`{
|
||||
"id":"conference",
|
||||
"command":"meeting",
|
||||
"toolOverrides":{"create":{"flags":{"count":{"default":3}}}}
|
||||
}`)
|
||||
got := OverlayFromJSON(valid)
|
||||
if got.ID != "conference" || got.Command != "meeting" ||
|
||||
got.ToolOverrides["create"].Flags["count"].Default != "3" {
|
||||
t.Fatalf("OverlayFromJSON(valid) = %#v", got)
|
||||
}
|
||||
|
||||
if got := OverlayFromJSON(json.RawMessage(`{`)); got.ID != "" || got.Command != "" {
|
||||
t.Fatalf("OverlayFromJSON(malformed) = %#v, want zero overlay", got)
|
||||
}
|
||||
}
|
||||
@@ -128,13 +128,8 @@ require_remote() {
|
||||
}
|
||||
|
||||
sync_main_if_safe() {
|
||||
current_branch="$(git symbolic-ref --quiet --short HEAD 2>/dev/null || true)"
|
||||
[ "$current_branch" = "main" ] || {
|
||||
printf 'release validation must run from the main worktree (current: %s)\n' "${current_branch:-detached HEAD}" >&2
|
||||
exit 1
|
||||
}
|
||||
[ -z "$(git status --porcelain --untracked-files=all)" ] || {
|
||||
printf '%s\n' 'release main worktree must be clean before synchronization' >&2
|
||||
printf '%s\n' 'release worktree must be clean before synchronization' >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
@@ -146,6 +141,14 @@ sync_main_if_safe() {
|
||||
if [ "$head_commit" = "$remote_commit" ]; then
|
||||
return 0
|
||||
fi
|
||||
current_branch="$(git symbolic-ref --quiet --short HEAD 2>/dev/null || true)"
|
||||
if [ "$current_branch" != "main" ]; then
|
||||
if git merge-base --is-ancestor HEAD "$remote_main"; then
|
||||
return 0
|
||||
fi
|
||||
printf 'HEAD is not contained in %s/main history; merge it through a reviewed PR before release\n' "$REMOTE" >&2
|
||||
exit 1
|
||||
fi
|
||||
if git merge-base --is-ancestor HEAD "$remote_main"; then
|
||||
git merge --ff-only "$remote_main"
|
||||
return 0
|
||||
|
||||
Executable
+203
@@ -0,0 +1,203 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
SCRIPT_DIR="$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)"
|
||||
. "$SCRIPT_DIR/release-lib.sh"
|
||||
|
||||
ROOT="$(CDPATH= cd -- "$SCRIPT_DIR/../.." && pwd)"
|
||||
CHANNEL=""
|
||||
REQUESTED_CHANNEL=""
|
||||
BUMP="patch"
|
||||
|
||||
usage() {
|
||||
cat >&2 <<'EOF'
|
||||
usage: next-release-version.sh --channel <beta|prerelease|stable> [options]
|
||||
|
||||
Options:
|
||||
--bump <patch|minor|major> Core bump when starting a new beta line (default: patch)
|
||||
--repo-root <path> Override repository root (primarily for tests)
|
||||
|
||||
Output:
|
||||
release_version=<vX.Y.Z[-beta.N]>
|
||||
from_beta=<vX.Y.Z-beta.N or empty>
|
||||
channel=<prerelease|stable>
|
||||
base=<latest allocated stable tag or empty>
|
||||
|
||||
Both ordinary release tags and refs/tags/withdrawn/v... tombstones reserve a
|
||||
version permanently. The caller must create the returned release tag
|
||||
atomically; this script only calculates the next candidate.
|
||||
EOF
|
||||
}
|
||||
|
||||
while [ "$#" -gt 0 ]; do
|
||||
case "$1" in
|
||||
--channel)
|
||||
[ "$#" -ge 2 ] || { usage; exit 2; }
|
||||
REQUESTED_CHANNEL="$2"
|
||||
shift 2
|
||||
;;
|
||||
--bump)
|
||||
[ "$#" -ge 2 ] || { usage; exit 2; }
|
||||
BUMP="$2"
|
||||
shift 2
|
||||
;;
|
||||
--repo-root)
|
||||
[ "$#" -ge 2 ] || { usage; exit 2; }
|
||||
ROOT="$2"
|
||||
shift 2
|
||||
;;
|
||||
-h|--help)
|
||||
usage
|
||||
exit 0
|
||||
;;
|
||||
*)
|
||||
printf 'unknown argument: %s\n' "$1" >&2
|
||||
usage
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
done
|
||||
|
||||
case "$REQUESTED_CHANNEL" in
|
||||
beta|prerelease) CHANNEL="prerelease" ;;
|
||||
stable) CHANNEL="stable" ;;
|
||||
*)
|
||||
printf 'invalid release channel: %s (expected beta, prerelease, or stable)\n' "${REQUESTED_CHANNEL:-<empty>}" >&2
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
case "$BUMP" in
|
||||
patch|minor|major) ;;
|
||||
*)
|
||||
printf 'invalid version bump: %s (expected patch, minor, or major)\n' "$BUMP" >&2
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
|
||||
cd "$ROOT"
|
||||
git rev-parse --is-inside-work-tree >/dev/null 2>&1 || {
|
||||
printf 'not a Git worktree: %s\n' "$ROOT" >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
allocated_version_for_ref() {
|
||||
_avfr_ref="$1"
|
||||
case "$_avfr_ref" in
|
||||
refs/tags/withdrawn/v*) printf '%s\n' "${_avfr_ref#refs/tags/withdrawn/}" ;;
|
||||
refs/tags/v*) printf '%s\n' "${_avfr_ref#refs/tags/}" ;;
|
||||
*) return 1 ;;
|
||||
esac
|
||||
}
|
||||
|
||||
stable_core_is_allocated() {
|
||||
_scia_core="$1"
|
||||
git show-ref --verify --quiet "refs/tags/$_scia_core" ||
|
||||
git show-ref --verify --quiet "refs/tags/withdrawn/$_scia_core"
|
||||
}
|
||||
|
||||
version_is_withdrawn() {
|
||||
git show-ref --verify --quiet "refs/tags/withdrawn/$1"
|
||||
}
|
||||
|
||||
bump_core() {
|
||||
_bc_version="${1#v}"
|
||||
_bc_major="${_bc_version%%.*}"
|
||||
_bc_remainder="${_bc_version#*.}"
|
||||
_bc_minor="${_bc_remainder%%.*}"
|
||||
_bc_patch="${_bc_remainder#*.}"
|
||||
case "$2" in
|
||||
patch) _bc_patch=$((_bc_patch + 1)) ;;
|
||||
minor)
|
||||
_bc_minor=$((_bc_minor + 1))
|
||||
_bc_patch=0
|
||||
;;
|
||||
major)
|
||||
_bc_major=$((_bc_major + 1))
|
||||
_bc_minor=0
|
||||
_bc_patch=0
|
||||
;;
|
||||
esac
|
||||
printf 'v%s.%s.%s\n' "$_bc_major" "$_bc_minor" "$_bc_patch"
|
||||
}
|
||||
|
||||
tag_refs="$(git for-each-ref --format='%(refname)' refs/tags)"
|
||||
latest_stable=""
|
||||
for ref in $tag_refs; do
|
||||
version="$(allocated_version_for_ref "$ref" 2>/dev/null || true)"
|
||||
[ -n "$version" ] || continue
|
||||
if release_is_stable_version "$version"; then
|
||||
if [ -z "$latest_stable" ] || release_core_is_greater "$version" "$latest_stable"; then
|
||||
latest_stable="$version"
|
||||
fi
|
||||
fi
|
||||
done
|
||||
|
||||
highest_open_beta_core=""
|
||||
for ref in $tag_refs; do
|
||||
version="$(allocated_version_for_ref "$ref" 2>/dev/null || true)"
|
||||
[ -n "$version" ] || continue
|
||||
release_is_prerelease_version "$version" || continue
|
||||
core="$(release_core_tag "$version")"
|
||||
stable_core_is_allocated "$core" && continue
|
||||
if [ -n "$latest_stable" ] && ! release_core_is_greater "$core" "$latest_stable"; then
|
||||
continue
|
||||
fi
|
||||
if [ -z "$highest_open_beta_core" ] || release_core_is_greater "$core" "$highest_open_beta_core"; then
|
||||
highest_open_beta_core="$core"
|
||||
fi
|
||||
done
|
||||
|
||||
release_version=""
|
||||
from_beta=""
|
||||
base="$latest_stable"
|
||||
|
||||
if [ "$CHANNEL" = "prerelease" ]; then
|
||||
if [ -n "$highest_open_beta_core" ]; then
|
||||
latest_beta=""
|
||||
for ref in $tag_refs; do
|
||||
version="$(allocated_version_for_ref "$ref" 2>/dev/null || true)"
|
||||
[ -n "$version" ] || continue
|
||||
release_is_prerelease_version "$version" || continue
|
||||
[ "$(release_core_tag "$version")" = "$highest_open_beta_core" ] || continue
|
||||
if [ -z "$latest_beta" ] || release_version_is_greater "$version" "$latest_beta"; then
|
||||
latest_beta="$version"
|
||||
fi
|
||||
done
|
||||
next_beta_number=$(( $(release_beta_number "$latest_beta") + 1 ))
|
||||
release_version="$highest_open_beta_core-beta.$next_beta_number"
|
||||
else
|
||||
[ -n "$latest_stable" ] || {
|
||||
printf 'cannot start a beta line without an allocated stable baseline\n' >&2
|
||||
exit 1
|
||||
}
|
||||
next_core="$(bump_core "$latest_stable" "$BUMP")"
|
||||
release_version="$next_core-beta.1"
|
||||
fi
|
||||
else
|
||||
[ -n "$highest_open_beta_core" ] || {
|
||||
printf 'cannot create a stable release without an open beta line newer than the latest allocated stable\n' >&2
|
||||
exit 1
|
||||
}
|
||||
latest_beta=""
|
||||
for ref in $tag_refs; do
|
||||
version="$(allocated_version_for_ref "$ref" 2>/dev/null || true)"
|
||||
[ -n "$version" ] || continue
|
||||
release_is_prerelease_version "$version" || continue
|
||||
[ "$(release_core_tag "$version")" = "$highest_open_beta_core" ] || continue
|
||||
if [ -z "$latest_beta" ] || release_version_is_greater "$version" "$latest_beta"; then
|
||||
latest_beta="$version"
|
||||
fi
|
||||
done
|
||||
if version_is_withdrawn "$latest_beta" ||
|
||||
! git show-ref --verify --quiet "refs/tags/$latest_beta"; then
|
||||
printf 'latest beta %s is withdrawn; create the next beta before stable promotion\n' "$latest_beta" >&2
|
||||
exit 1
|
||||
fi
|
||||
release_version="$highest_open_beta_core"
|
||||
from_beta="$latest_beta"
|
||||
fi
|
||||
|
||||
printf 'release_version=%s\n' "$release_version"
|
||||
printf 'from_beta=%s\n' "$from_beta"
|
||||
printf 'channel=%s\n' "$CHANNEL"
|
||||
printf 'base=%s\n' "$base"
|
||||
@@ -120,12 +120,41 @@ git fetch --force --no-tags "$REMOTE" "+refs/tags/$VERSION:$recovery_ref"
|
||||
}
|
||||
tag_object="$(git rev-parse "$recovery_ref")"
|
||||
commit="$(git rev-parse "$recovery_ref^{commit}")"
|
||||
tag_message="$(git for-each-ref "$recovery_ref" --format='%(contents)')"
|
||||
cloud_run_id="$(printf '%s\n' "$tag_message" | awk -F ': ' '$1 == "Release-Run" { print $2 }')"
|
||||
cloud_run_attempt="$(printf '%s\n' "$tag_message" | awk -F ': ' '$1 == "Release-Run-Attempt" { print $2 }')"
|
||||
cloud_actor="$(printf '%s\n' "$tag_message" | awk -F ': ' '$1 == "Requested-By" { print $2 }')"
|
||||
cloud_actor_id="$(printf '%s\n' "$tag_message" | awk -F ': ' '$1 == "Requested-By-ID" { print $2 }')"
|
||||
cloud_sealed_commit="$(printf '%s\n' "$tag_message" | awk -F ': ' '$1 == "Sealed-Commit" { print $2 }')"
|
||||
cloud_marker_count=0
|
||||
for cloud_value in "$cloud_run_id" "$cloud_run_attempt" "$cloud_actor" "$cloud_actor_id" "$cloud_sealed_commit"; do
|
||||
[ -z "$cloud_value" ] || cloud_marker_count=$((cloud_marker_count + 1))
|
||||
done
|
||||
if [ "$cloud_marker_count" -ne 0 ] && [ "$cloud_marker_count" -ne 5 ]; then
|
||||
printf '%s contains incomplete cloud release metadata\n' "$VERSION" >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ "$cloud_marker_count" -eq 5 ]; then
|
||||
for cloud_number in "$cloud_run_id" "$cloud_run_attempt" "$cloud_actor_id"; do
|
||||
printf '%s\n' "$cloud_number" | grep -Eq '^[1-9][0-9]*$' || {
|
||||
printf '%s contains invalid cloud release identity\n' "$VERSION" >&2
|
||||
exit 1
|
||||
}
|
||||
done
|
||||
[ "$cloud_sealed_commit" = "$commit" ] || {
|
||||
printf '%s cloud release metadata is not bound to %s\n' "$VERSION" "$commit" >&2
|
||||
exit 1
|
||||
}
|
||||
fi
|
||||
git merge-base --is-ancestor "$commit" "refs/remotes/$REMOTE/main" || {
|
||||
printf '%s commit %s is not contained in %s/main\n' "$VERSION" "$commit" "$REMOTE" >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
if [ -z "$FAILED_RUN_ID" ]; then
|
||||
if [ -z "$FAILED_RUN_ID" ] && [ "$cloud_marker_count" -eq 5 ]; then
|
||||
FAILED_RUN_ID="$cloud_run_id"
|
||||
FAILED_RUN_ATTEMPT="$cloud_run_attempt"
|
||||
elif [ -z "$FAILED_RUN_ID" ]; then
|
||||
candidate_runs="$(
|
||||
gh api \
|
||||
-H 'Accept: application/vnd.github+json' \
|
||||
@@ -160,7 +189,7 @@ attempt_record="$(
|
||||
gh api \
|
||||
-H 'Accept: application/vnd.github+json' \
|
||||
"repos/$EXPECTED_REPOSITORY/actions/runs/$FAILED_RUN_ID/attempts/$FAILED_RUN_ATTEMPT" \
|
||||
--jq '[.id, .run_attempt, .repository.full_name, .path, .event, .status, .conclusion, .head_branch, .head_sha] | @tsv'
|
||||
--jq '[.id, .run_attempt, .repository.full_name, .path, .event, .status, .conclusion, .head_branch, .head_sha, .actor.login, .actor.id] | @tsv'
|
||||
)" || {
|
||||
printf 'could not query Release run %s attempt %s\n' "$FAILED_RUN_ID" "$FAILED_RUN_ATTEMPT" >&2
|
||||
exit 1
|
||||
@@ -174,6 +203,7 @@ attempt_status="$(printf '%s\n' "$attempt_record" | cut -f6)"
|
||||
attempt_conclusion="$(printf '%s\n' "$attempt_record" | cut -f7)"
|
||||
attempt_branch="$(printf '%s\n' "$attempt_record" | cut -f8)"
|
||||
attempt_commit="$(printf '%s\n' "$attempt_record" | cut -f9)"
|
||||
attempt_actor_id="$(printf '%s\n' "$attempt_record" | cut -f11)"
|
||||
case "$attempt_conclusion" in
|
||||
failure|cancelled|timed_out|startup_failure|stale) ;;
|
||||
*)
|
||||
@@ -182,18 +212,33 @@ case "$attempt_conclusion" in
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
if [ "$cloud_marker_count" -eq 5 ]; then
|
||||
expected_attempt_event="workflow_dispatch"
|
||||
expected_attempt_branch="main"
|
||||
else
|
||||
expected_attempt_event="push"
|
||||
expected_attempt_branch="$VERSION"
|
||||
fi
|
||||
if [ "$attempt_id" != "$FAILED_RUN_ID" ] ||
|
||||
[ "$attempt_number" != "$FAILED_RUN_ATTEMPT" ] ||
|
||||
[ "$attempt_repository" != "$EXPECTED_REPOSITORY" ] ||
|
||||
[ "$attempt_path" != ".github/workflows/release.yml" ] ||
|
||||
[ "$attempt_event" != "push" ] ||
|
||||
[ "$attempt_event" != "$expected_attempt_event" ] ||
|
||||
[ "$attempt_status" != "completed" ] ||
|
||||
[ "$attempt_branch" != "$VERSION" ] ||
|
||||
[ "$attempt_branch" != "$expected_attempt_branch" ] ||
|
||||
[ "$attempt_commit" != "$commit" ]; then
|
||||
printf 'Release run %s attempt %s does not match %s at %s\n' \
|
||||
"$FAILED_RUN_ID" "$FAILED_RUN_ATTEMPT" "$VERSION" "$commit" >&2
|
||||
exit 1
|
||||
fi
|
||||
if [ "$cloud_marker_count" -eq 5 ] &&
|
||||
{ [ "$cloud_run_id" != "$FAILED_RUN_ID" ] ||
|
||||
[ "$cloud_run_attempt" != "$FAILED_RUN_ATTEMPT" ] ||
|
||||
[ "$cloud_actor_id" != "$attempt_actor_id" ]; }; then
|
||||
printf 'Release run %s attempt %s is not bound by the cloud seal for %s\n' \
|
||||
"$FAILED_RUN_ID" "$FAILED_RUN_ATTEMPT" "$VERSION" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
printf 'Recovery target:\n'
|
||||
printf ' version: %s\n' "$VERSION"
|
||||
|
||||
@@ -71,27 +71,27 @@ git rev-parse --verify --quiet "$remote_main^{commit}" >/dev/null || {
|
||||
printf 'release branch is not available locally: %s/%s\n' "$REMOTE" "$BRANCH" >&2
|
||||
exit 1
|
||||
}
|
||||
remote_main_commit="$(git rev-parse "$remote_main^{commit}")"
|
||||
|
||||
if [ "$CONTEXT" = "local" ]; then
|
||||
[ -z "$(git status --porcelain --untracked-files=all)" ] || {
|
||||
printf 'release worktree must be clean (staged, unstaged, and untracked files are blocked)\n' >&2
|
||||
exit 1
|
||||
}
|
||||
current_branch="$(git symbolic-ref --quiet --short HEAD 2>/dev/null || true)"
|
||||
[ "$current_branch" = "$BRANCH" ] || {
|
||||
printf 'local release must run from branch %s (current: %s)\n' "$BRANCH" "${current_branch:-detached HEAD}" >&2
|
||||
exit 1
|
||||
}
|
||||
[ "$head_commit" = "$remote_main_commit" ] || {
|
||||
printf 'HEAD must exactly match %s/%s before release\n' "$REMOTE" "$BRANCH" >&2
|
||||
git merge-base --is-ancestor HEAD "$remote_main" || {
|
||||
printf 'HEAD must be contained in %s/%s history before release\n' "$REMOTE" "$BRANCH" >&2
|
||||
exit 1
|
||||
}
|
||||
if git rev-parse --verify --quiet "refs/tags/$VERSION" >/dev/null; then
|
||||
printf 'release tag already exists locally: %s\n' "$VERSION" >&2
|
||||
exit 1
|
||||
fi
|
||||
remote_tag="$(git ls-remote --tags "$REMOTE" "refs/tags/$VERSION" "refs/tags/$VERSION^{}")" || {
|
||||
if git rev-parse --verify --quiet "refs/tags/withdrawn/$VERSION" >/dev/null; then
|
||||
printf 'release version was withdrawn and can never be reused: %s\n' "$VERSION" >&2
|
||||
exit 1
|
||||
fi
|
||||
remote_tag="$(git ls-remote --tags "$REMOTE" \
|
||||
"refs/tags/$VERSION" "refs/tags/$VERSION^{}" \
|
||||
"refs/tags/withdrawn/$VERSION" "refs/tags/withdrawn/$VERSION^{}")" || {
|
||||
printf 'could not query release tags from remote: %s\n' "$REMOTE" >&2
|
||||
exit 1
|
||||
}
|
||||
@@ -100,6 +100,10 @@ if [ "$CONTEXT" = "local" ]; then
|
||||
exit 1
|
||||
}
|
||||
else
|
||||
if git rev-parse --verify --quiet "refs/tags/withdrawn/$VERSION" >/dev/null; then
|
||||
printf 'release version was withdrawn and can never be reused: %s\n' "$VERSION" >&2
|
||||
exit 1
|
||||
fi
|
||||
git rev-parse --verify --quiet "refs/tags/$VERSION^{commit}" >/dev/null || {
|
||||
printf 'CI release tag is not available: %s\n' "$VERSION" >&2
|
||||
exit 1
|
||||
@@ -120,32 +124,58 @@ fi
|
||||
|
||||
previous_stable=""
|
||||
previous_stable_commit=""
|
||||
for tag in $(git tag --list 'v*' --sort=-version:refname); do
|
||||
latest_allocated_stable=""
|
||||
for tag_ref in $(git for-each-ref --format='%(refname)' refs/tags); do
|
||||
case "$tag_ref" in
|
||||
refs/tags/withdrawn/*) tag="${tag_ref#refs/tags/withdrawn/}" ;;
|
||||
refs/tags/*) tag="${tag_ref#refs/tags/}" ;;
|
||||
*) continue ;;
|
||||
esac
|
||||
[ "$tag" = "$VERSION" ] && continue
|
||||
if release_is_stable_version "$tag"; then
|
||||
previous_stable="$tag"
|
||||
previous_stable_commit="$(git rev-parse "$tag^{commit}")"
|
||||
break
|
||||
release_is_stable_version "$tag" || continue
|
||||
if [ -z "$latest_allocated_stable" ] ||
|
||||
release_core_is_greater "$tag" "$latest_allocated_stable"; then
|
||||
latest_allocated_stable="$tag"
|
||||
fi
|
||||
done
|
||||
if [ -n "$previous_stable" ] && ! release_core_is_greater "$VERSION" "$previous_stable"; then
|
||||
printf 'release version %s must be greater than latest stable %s\n' "$VERSION" "$previous_stable" >&2
|
||||
if [ -n "$latest_allocated_stable" ] &&
|
||||
! release_core_is_greater "$VERSION" "$latest_allocated_stable"; then
|
||||
printf 'release version %s must be greater than latest allocated stable %s\n' \
|
||||
"$VERSION" "$latest_allocated_stable" >&2
|
||||
exit 1
|
||||
fi
|
||||
|
||||
for tag in $(git tag --list 'v*' --sort=-version:refname); do
|
||||
[ "$tag" = "$VERSION" ] && continue
|
||||
release_is_stable_version "$tag" || continue
|
||||
if git rev-parse --verify --quiet "refs/tags/withdrawn/$tag" >/dev/null; then
|
||||
continue
|
||||
fi
|
||||
previous_stable="$tag"
|
||||
previous_stable_commit="$(git rev-parse "$tag^{commit}")"
|
||||
break
|
||||
done
|
||||
|
||||
core_tag="$(release_core_tag "$VERSION")"
|
||||
if [ "$CHANNEL" = "prerelease" ]; then
|
||||
[ -z "$FROM_BETA" ] || { printf -- '--from-beta is only valid for stable releases\n' >&2; exit 1; }
|
||||
if git rev-parse --verify --quiet "refs/tags/$core_tag" >/dev/null; then
|
||||
if git rev-parse --verify --quiet "refs/tags/$core_tag" >/dev/null ||
|
||||
git rev-parse --verify --quiet "refs/tags/withdrawn/$core_tag" >/dev/null; then
|
||||
printf 'cannot publish prerelease after stable tag exists: %s\n' "$core_tag" >&2
|
||||
exit 1
|
||||
fi
|
||||
previous_beta=""
|
||||
for tag in $(git tag --list "$core_tag-beta.*" --sort=-version:refname); do
|
||||
for tag_ref in $(git for-each-ref --format='%(refname)' \
|
||||
"refs/tags/$core_tag-beta.*" "refs/tags/withdrawn/$core_tag-beta.*"); do
|
||||
case "$tag_ref" in
|
||||
refs/tags/withdrawn/*) tag="${tag_ref#refs/tags/withdrawn/}" ;;
|
||||
refs/tags/*) tag="${tag_ref#refs/tags/}" ;;
|
||||
*) continue ;;
|
||||
esac
|
||||
[ "$tag" = "$VERSION" ] && continue
|
||||
if release_is_prerelease_version "$tag"; then
|
||||
release_is_prerelease_version "$tag" || continue
|
||||
if [ -z "$previous_beta" ] || release_version_is_greater "$tag" "$previous_beta"; then
|
||||
previous_beta="$tag"
|
||||
break
|
||||
fi
|
||||
done
|
||||
beta_number="$(release_beta_number "$VERSION")"
|
||||
@@ -178,6 +208,10 @@ else
|
||||
printf 'stable version %s does not match beta baseline %s\n' "$VERSION" "$FROM_BETA" >&2
|
||||
exit 1
|
||||
}
|
||||
if git rev-parse --verify --quiet "refs/tags/withdrawn/$FROM_BETA" >/dev/null; then
|
||||
printf 'stable beta baseline was withdrawn: %s\n' "$FROM_BETA" >&2
|
||||
exit 1
|
||||
fi
|
||||
git rev-parse --verify --quiet "refs/tags/$FROM_BETA^{commit}" >/dev/null || {
|
||||
printf 'stable beta baseline is not available locally: %s\n' "$FROM_BETA" >&2
|
||||
exit 1
|
||||
@@ -187,11 +221,6 @@ else
|
||||
printf 'stable beta baseline is not an ancestor of HEAD: %s\n' "$FROM_BETA" >&2
|
||||
exit 1
|
||||
}
|
||||
if ! git diff --quiet "$FROM_BETA^{commit}" HEAD -- . ':(exclude)CHANGELOG.md'; then
|
||||
printf 'stable source drifted from %s; only CHANGELOG.md may differ\n' "$FROM_BETA" >&2
|
||||
git diff --name-only "$FROM_BETA^{commit}" HEAD -- . ':(exclude)CHANGELOG.md' >&2
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
semver="$(release_semver "$VERSION")"
|
||||
|
||||
Executable
+52
@@ -0,0 +1,52 @@
|
||||
#!/bin/sh
|
||||
set -eu
|
||||
|
||||
TAG="${1:-}"
|
||||
[ -n "$TAG" ] || {
|
||||
printf 'usage: release-tag-oss-mode.sh <tag>\n' >&2
|
||||
exit 2
|
||||
}
|
||||
|
||||
tag_ref="refs/tags/$TAG"
|
||||
git rev-parse --verify --quiet "$tag_ref" >/dev/null || {
|
||||
printf 'release tag is not available: %s\n' "$TAG" >&2
|
||||
exit 1
|
||||
}
|
||||
tag_object="$(git rev-parse "$tag_ref")"
|
||||
[ "$(git cat-file -t "$tag_object")" = tag ] || {
|
||||
printf 'release tag must be annotated: %s\n' "$TAG" >&2
|
||||
exit 1
|
||||
}
|
||||
tag_contents="$(git cat-file tag "$tag_object")" || {
|
||||
printf 'release tag could not be read: %s\n' "$TAG" >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
mode="$(
|
||||
printf '%s\n' "$tag_contents" |
|
||||
awk '
|
||||
{
|
||||
line = $0
|
||||
sub(/\r$/, "", line)
|
||||
}
|
||||
line ~ /^OSS-Mirror: / {
|
||||
count++
|
||||
value = substr(line, length("OSS-Mirror: ") + 1)
|
||||
}
|
||||
END {
|
||||
if (count > 1) exit 2
|
||||
if (count == 0) print "enabled"
|
||||
else print value
|
||||
}
|
||||
'
|
||||
)" || {
|
||||
printf 'release tag contains duplicate OSS-Mirror metadata: %s\n' "$TAG" >&2
|
||||
exit 1
|
||||
}
|
||||
case "$mode" in
|
||||
enabled|deferred) printf '%s\n' "$mode" ;;
|
||||
*)
|
||||
printf 'release tag contains invalid OSS-Mirror metadata: %s (%s)\n' "$TAG" "$mode" >&2
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
@@ -288,8 +288,12 @@ build_policy_binary() {
|
||||
}
|
||||
|
||||
fetch_release_tags() {
|
||||
git fetch --force "$REMOTE" '+refs/tags/v*:refs/tags/v*'
|
||||
git fetch --force --no-tags "$OFFICIAL_TAGS_URL" '+refs/tags/v*:refs/tags/v*'
|
||||
git fetch --force "$REMOTE" \
|
||||
'+refs/tags/v*:refs/tags/v*' \
|
||||
'+refs/tags/withdrawn/v*:refs/tags/withdrawn/v*'
|
||||
git fetch --force --no-tags "$OFFICIAL_TAGS_URL" \
|
||||
'+refs/tags/v*:refs/tags/v*' \
|
||||
'+refs/tags/withdrawn/v*:refs/tags/withdrawn/v*'
|
||||
}
|
||||
|
||||
printf '==> Refreshing %s/%s and release tags\n' "$REMOTE" "$BRANCH"
|
||||
@@ -347,7 +351,7 @@ else
|
||||
if [ -n "$previous_stable" ]; then
|
||||
printf '==> Comparing command tree with %s\n' "$previous_stable"
|
||||
"$ROOT/scripts/policy/check-command-compatibility.sh" \
|
||||
--base-ref "$REMOTE/$BRANCH" \
|
||||
--base-ref HEAD \
|
||||
--stable-ref "$previous_stable"
|
||||
fi
|
||||
|
||||
@@ -401,7 +405,7 @@ if [ "$previous_stable" != "$previous_stable_before_refresh" ]; then
|
||||
printf '==> Stable authority advanced from %s to %s; rechecking command compatibility\n' \
|
||||
"${previous_stable_before_refresh:-none}" "$previous_stable"
|
||||
"$ROOT/scripts/policy/check-command-compatibility.sh" \
|
||||
--base-ref "$REMOTE/$BRANCH" \
|
||||
--base-ref HEAD \
|
||||
--stable-ref "$previous_stable"
|
||||
fi
|
||||
|
||||
@@ -412,9 +416,9 @@ fi
|
||||
|
||||
# Delivery, compatibility, and publication checks above may take long enough
|
||||
# for main or stable authority to move. This last refresh must be followed only
|
||||
# by local proof/tag creation. The atomic push advertises main with the tag, so
|
||||
# an already-advanced remote main rejects the whole transaction; a later main
|
||||
# advance is safe because the sealed commit remains in protected main history.
|
||||
# by local proof/tag creation. Only the tag is pushed: the sealed commit is
|
||||
# already contained in protected main history, so a later main advance never
|
||||
# invalidates the release.
|
||||
printf '==> Settling final %s/%s and stable authority\n' "$REMOTE" "$BRANCH"
|
||||
git fetch --force "$REMOTE" "+refs/heads/$BRANCH:refs/remotes/$REMOTE/$BRANCH"
|
||||
fetch_release_tags
|
||||
@@ -443,8 +447,7 @@ else
|
||||
git tag -a "$VERSION" -m "Release $VERSION" -m 'Channel: prerelease'
|
||||
fi
|
||||
|
||||
if ! git push --atomic "$push_url" \
|
||||
"HEAD:refs/heads/$BRANCH" "refs/tags/$VERSION"; then
|
||||
if ! git push "$push_url" "refs/tags/$VERSION"; then
|
||||
set +e
|
||||
remote_refs="$(git ls-remote --tags "$push_url" "refs/tags/$VERSION" "refs/tags/$VERSION^{}")"
|
||||
query_status=$?
|
||||
|
||||
@@ -82,7 +82,7 @@ if [ "$delivered_by_push" -ne 1 ]; then
|
||||
if DWS_RELEASE_OFFICIAL_REPOSITORY="$REPOSITORY" \
|
||||
DWS_RELEASE_GITHUB_TOKEN="$API_TOKEN" \
|
||||
"$SCRIPT_DIR/verify-release-workflow-delivery.sh" "$TAG" "$EXPECTED_COMMIT" 2>/dev/null; then
|
||||
printf 'Delivered stable baseline verified through protected default-branch recovery: %s -> %s\n' \
|
||||
printf 'Delivered stable baseline verified through trusted default-branch delivery: %s -> %s\n' \
|
||||
"$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
@@ -53,8 +53,6 @@ while [ "$#" -gt 0 ]; do
|
||||
esac
|
||||
done
|
||||
|
||||
[ -z "$EXPECTED_VERSION" ] || need_cmd strings
|
||||
|
||||
TMP_ROOT="$(mktemp -d "${TMPDIR:-/tmp}/dws-package-verify-XXXXXX")"
|
||||
HOME_AGENT_PARENTS="
|
||||
.claude
|
||||
@@ -145,8 +143,8 @@ verify_npm() {
|
||||
if [ -n "$EXPECTED_VERSION" ]; then
|
||||
vendor_bin="$npm_prefix/lib/node_modules/dingtalk-workspace-cli/vendor/dws"
|
||||
need_file "$vendor_bin"
|
||||
strings "$vendor_bin" | grep -Fqx "v$EXPECTED_VERSION" || \
|
||||
err "npm-installed binary does not embed expected version v$EXPECTED_VERSION"
|
||||
LC_ALL=C grep -aFq "v$EXPECTED_VERSION" "$vendor_bin" || \
|
||||
err "npm-installed binary does not contain expected version marker v$EXPECTED_VERSION"
|
||||
EXPECTED_VERSION="$EXPECTED_VERSION" node -e '
|
||||
const pkg = require(process.argv[1]);
|
||||
if (pkg.version !== process.env.EXPECTED_VERSION) process.exit(1);
|
||||
@@ -193,8 +191,8 @@ verify_brew() {
|
||||
[ -x "$prefix/bin/dws" ] || err "brew install did not create $prefix/bin/dws"
|
||||
"$prefix/bin/dws" --help >/dev/null
|
||||
if [ -n "$EXPECTED_VERSION" ]; then
|
||||
strings "$prefix/bin/dws" | grep -Fqx "v$EXPECTED_VERSION" || \
|
||||
err "Homebrew-installed binary does not embed expected version v$EXPECTED_VERSION"
|
||||
LC_ALL=C grep -aFq "v$EXPECTED_VERSION" "$prefix/bin/dws" || \
|
||||
err "Homebrew-installed binary does not contain expected version marker v$EXPECTED_VERSION"
|
||||
fi
|
||||
need_file "$prefix/share/dingtalk-workspace-cli-local/skills/dws/SKILL.md"
|
||||
}
|
||||
|
||||
@@ -94,7 +94,7 @@ verify_binary_version() {
|
||||
printf '%s does not contain the expected dws binary\n' "$asset" >&2
|
||||
return 1
|
||||
}
|
||||
strings "$binary" | grep -Fqx "v$SEMVER" || {
|
||||
LC_ALL=C grep -aFq "v$SEMVER" "$binary" || {
|
||||
printf '%s binary does not embed expected version v%s\n' "$asset" "$SEMVER" >&2
|
||||
return 1
|
||||
}
|
||||
|
||||
@@ -18,6 +18,9 @@ if [ "${1:-}" = "--channel-repair" ]; then
|
||||
exit 2
|
||||
;;
|
||||
esac
|
||||
elif [ "${1:-}" = "--npm-repair" ]; then
|
||||
MODE="npm-repair"
|
||||
shift
|
||||
fi
|
||||
|
||||
TAG="${1:-}"
|
||||
@@ -25,7 +28,7 @@ EXPECTED_COMMIT="${2:-}"
|
||||
REPOSITORY="${DWS_RELEASE_OFFICIAL_REPOSITORY:-DingTalk-Real-AI/dingtalk-workspace-cli}"
|
||||
|
||||
[ -n "$TAG" ] && [ -n "$EXPECTED_COMMIT" ] || {
|
||||
printf 'usage: verify-release-workflow-delivery.sh [--channel-repair <oss|gitee>] <tag> <commit>\n' >&2
|
||||
printf 'usage: verify-release-workflow-delivery.sh [--channel-repair <oss|gitee> | --npm-repair] <tag> <commit>\n' >&2
|
||||
exit 2
|
||||
}
|
||||
if ! release_is_stable_version "$TAG" && ! release_is_prerelease_version "$TAG"; then
|
||||
@@ -81,6 +84,209 @@ for run in runs:
|
||||
done
|
||||
}
|
||||
|
||||
find_cloud_delivery_identity() {
|
||||
tag_ref="$(
|
||||
github_get "repos/$REPOSITORY/git/ref/tags/$TAG" \
|
||||
| python3 -c 'import json,sys
|
||||
ref=json.load(sys.stdin)
|
||||
obj=ref.get("object", {})
|
||||
if obj.get("type") == "tag" and obj.get("sha"):
|
||||
print(obj["sha"])'
|
||||
)" || return 1
|
||||
[ -n "$tag_ref" ] || return 1
|
||||
github_get "repos/$REPOSITORY/git/tags/$tag_ref" \
|
||||
| python3 -c 'import json,re,sys
|
||||
tag,commit=sys.argv[1:]
|
||||
payload=json.load(sys.stdin)
|
||||
if payload.get("tag") != tag or payload.get("object", {}).get("type") != "commit":
|
||||
raise SystemExit(1)
|
||||
if payload.get("object", {}).get("sha") != commit:
|
||||
raise SystemExit(1)
|
||||
fields={}
|
||||
for line in payload.get("message", "").splitlines():
|
||||
if ": " not in line:
|
||||
continue
|
||||
key,value=line.split(": ", 1)
|
||||
if key in fields:
|
||||
raise SystemExit(1)
|
||||
fields[key]=value
|
||||
required={
|
||||
"Channel", "Release-Run", "Release-Run-Attempt", "Requested-By",
|
||||
"Requested-By-ID", "Sealed-Commit", "Workflow-Commit",
|
||||
"Allocation-Fingerprint",
|
||||
}
|
||||
if not required.issubset(fields):
|
||||
raise SystemExit(1)
|
||||
if fields["Sealed-Commit"] != commit or fields["Workflow-Commit"] != commit:
|
||||
raise SystemExit(1)
|
||||
if not re.fullmatch(r"[1-9][0-9]*", fields["Release-Run"]):
|
||||
raise SystemExit(1)
|
||||
if not re.fullmatch(r"[1-9][0-9]*", fields["Release-Run-Attempt"]):
|
||||
raise SystemExit(1)
|
||||
if not re.fullmatch(r"[1-9][0-9]*", fields["Requested-By-ID"]):
|
||||
raise SystemExit(1)
|
||||
if not fields["Requested-By"] or not re.fullmatch(r"[0-9a-f]{64}", fields["Allocation-Fingerprint"]):
|
||||
raise SystemExit(1)
|
||||
is_beta="-beta." in tag
|
||||
if fields["Channel"] != ("prerelease" if is_beta else "stable"):
|
||||
raise SystemExit(1)
|
||||
from_beta=fields.get("From-Beta", "")
|
||||
if is_beta:
|
||||
if from_beta:
|
||||
raise SystemExit(1)
|
||||
else:
|
||||
core=re.escape(tag)
|
||||
if not re.fullmatch(core + r"-beta\.[1-9][0-9]*", from_beta):
|
||||
raise SystemExit(1)
|
||||
print("\t".join([
|
||||
fields["Release-Run"],
|
||||
fields["Release-Run-Attempt"],
|
||||
fields["Requested-By"],
|
||||
fields["Requested-By-ID"],
|
||||
fields["Workflow-Commit"],
|
||||
]))' "$TAG" "$EXPECTED_COMMIT"
|
||||
}
|
||||
|
||||
verify_cloud_delivery() {
|
||||
cloud_identity="$1"
|
||||
cloud_run_id="$(printf '%s\n' "$cloud_identity" | cut -f1)"
|
||||
cloud_run_attempt="$(printf '%s\n' "$cloud_identity" | cut -f2)"
|
||||
cloud_actor="$(printf '%s\n' "$cloud_identity" | cut -f3)"
|
||||
cloud_actor_id="$(printf '%s\n' "$cloud_identity" | cut -f4)"
|
||||
cloud_workflow_sha="$(printf '%s\n' "$cloud_identity" | cut -f5)"
|
||||
[ -n "$cloud_run_id" ] && [ -n "$cloud_run_attempt" ] &&
|
||||
[ -n "$cloud_actor" ] && [ -n "$cloud_actor_id" ] &&
|
||||
[ -n "$cloud_workflow_sha" ] || return 1
|
||||
|
||||
cloud_run_state="$(
|
||||
github_get "repos/$REPOSITORY/actions/runs/$cloud_run_id/attempts/$cloud_run_attempt" \
|
||||
| python3 -c 'import json,sys
|
||||
r=json.load(sys.stdin)
|
||||
print("\t".join(str(value) for value in (
|
||||
r.get("id", ""),
|
||||
r.get("run_attempt", ""),
|
||||
r.get("repository", {}).get("full_name", ""),
|
||||
r.get("path", ""),
|
||||
r.get("event", ""),
|
||||
r.get("status", ""),
|
||||
r.get("conclusion", ""),
|
||||
r.get("head_branch", ""),
|
||||
r.get("head_sha", ""),
|
||||
r.get("actor", {}).get("login", ""),
|
||||
r.get("actor", {}).get("id", ""),
|
||||
)))'
|
||||
)" || return 1
|
||||
expected_cloud_core="$(printf '%s\t%s\t%s\t.github/workflows/release.yml\tworkflow_dispatch\tcompleted\tsuccess\tmain\t%s' \
|
||||
"$cloud_run_id" "$cloud_run_attempt" "$REPOSITORY" "$cloud_workflow_sha")"
|
||||
[ "$(printf '%s\n' "$cloud_run_state" | cut -f1-9)" = "$expected_cloud_core" ] ||
|
||||
return 1
|
||||
[ -n "$(printf '%s\n' "$cloud_run_state" | cut -f10)" ] || return 1
|
||||
[ "$(printf '%s\n' "$cloud_run_state" | cut -f11)" = "$cloud_actor_id" ] ||
|
||||
return 1
|
||||
[ "$cloud_workflow_sha" = "$EXPECTED_COMMIT" ] || return 1
|
||||
|
||||
jobs_dir="$(mktemp -d "${TMPDIR:-/tmp}/dws-release-cloud-jobs.XXXXXX")"
|
||||
page=1
|
||||
while :; do
|
||||
jobs_page="$jobs_dir/jobs-$page.json"
|
||||
if ! github_get "repos/$REPOSITORY/actions/runs/$cloud_run_id/attempts/$cloud_run_attempt/jobs?per_page=100&page=$page" \
|
||||
>"$jobs_page"; then
|
||||
rm -rf "$jobs_dir"
|
||||
return 1
|
||||
fi
|
||||
page_count="$(
|
||||
python3 -c 'import json,sys; print(len(json.load(open(sys.argv[1])).get("jobs", [])))' \
|
||||
"$jobs_page"
|
||||
)" || {
|
||||
rm -rf "$jobs_dir"
|
||||
return 1
|
||||
}
|
||||
[ "$page_count" -eq 100 ] || break
|
||||
page=$((page + 1))
|
||||
done
|
||||
|
||||
result=0
|
||||
python3 - "$cloud_workflow_sha" "$jobs_dir"/jobs-*.json <<'PY' || result=$?
|
||||
import json
|
||||
import sys
|
||||
|
||||
workflow_sha, *pages = sys.argv[1:]
|
||||
jobs = []
|
||||
for page in pages:
|
||||
with open(page, encoding="utf-8") as handle:
|
||||
jobs.extend(json.load(handle).get("jobs", []))
|
||||
|
||||
required = (
|
||||
"Plan next cloud release",
|
||||
"Seal cloud release tag",
|
||||
"release-contract",
|
||||
"Build signed release artifacts",
|
||||
"Verify Apple Developer ID signatures",
|
||||
"Publish immutable GitHub Release",
|
||||
"Publish npm and mirrors",
|
||||
"Release delivery gate",
|
||||
)
|
||||
for name in required:
|
||||
matches = [job for job in jobs if job.get("name") == name]
|
||||
if len(matches) != 1:
|
||||
raise SystemExit(1)
|
||||
job = matches[0]
|
||||
if (
|
||||
job.get("head_sha") != workflow_sha
|
||||
or job.get("status") != "completed"
|
||||
or job.get("conclusion") != "success"
|
||||
):
|
||||
raise SystemExit(1)
|
||||
|
||||
seal = next(job for job in jobs if job.get("name") == "Seal cloud release tag")
|
||||
steps = [
|
||||
step for step in seal.get("steps", [])
|
||||
if step.get("name") == "Create one immutable annotated release tag"
|
||||
]
|
||||
if (
|
||||
len(steps) != 1
|
||||
or steps[0].get("status") != "completed"
|
||||
or steps[0].get("conclusion") != "success"
|
||||
):
|
||||
raise SystemExit(1)
|
||||
PY
|
||||
rm -rf "$jobs_dir"
|
||||
return "$result"
|
||||
}
|
||||
|
||||
verify_failed_cloud_delivery_identity() {
|
||||
cloud_identity="$1"
|
||||
cloud_run_id="$(printf '%s\n' "$cloud_identity" | cut -f1)"
|
||||
cloud_run_attempt="$(printf '%s\n' "$cloud_identity" | cut -f2)"
|
||||
cloud_actor="$(printf '%s\n' "$cloud_identity" | cut -f3)"
|
||||
cloud_actor_id="$(printf '%s\n' "$cloud_identity" | cut -f4)"
|
||||
cloud_workflow_sha="$(printf '%s\n' "$cloud_identity" | cut -f5)"
|
||||
[ "$cloud_workflow_sha" = "$EXPECTED_COMMIT" ] || return 1
|
||||
cloud_run_state="$(
|
||||
github_get "repos/$REPOSITORY/actions/runs/$cloud_run_id/attempts/$cloud_run_attempt" \
|
||||
| python3 -c 'import json,sys
|
||||
r=json.load(sys.stdin)
|
||||
print("\t".join(str(value) for value in (
|
||||
r.get("id", ""),
|
||||
r.get("run_attempt", ""),
|
||||
r.get("repository", {}).get("full_name", ""),
|
||||
r.get("path", ""),
|
||||
r.get("event", ""),
|
||||
r.get("status", ""),
|
||||
r.get("conclusion", ""),
|
||||
r.get("head_branch", ""),
|
||||
r.get("head_sha", ""),
|
||||
r.get("actor", {}).get("login", ""),
|
||||
r.get("actor", {}).get("id", ""),
|
||||
)))'
|
||||
)" || return 1
|
||||
expected_cloud_core="$(printf '%s\t%s\t%s\t.github/workflows/release.yml\tworkflow_dispatch\tcompleted\tfailure\tmain\t%s' \
|
||||
"$cloud_run_id" "$cloud_run_attempt" "$REPOSITORY" "$cloud_workflow_sha")"
|
||||
[ "$(printf '%s\n' "$cloud_run_state" | cut -f1-9)" = "$expected_cloud_core" ] &&
|
||||
[ -n "$(printf '%s\n' "$cloud_run_state" | cut -f10)" ] &&
|
||||
[ "$(printf '%s\n' "$cloud_run_state" | cut -f11)" = "$cloud_actor_id" ]
|
||||
}
|
||||
|
||||
find_failed_push_delivery() {
|
||||
matches=""
|
||||
page=1
|
||||
@@ -275,6 +481,112 @@ PY
|
||||
return "$result"
|
||||
}
|
||||
|
||||
verify_npm_repair_delivery() {
|
||||
run_id="$1"
|
||||
run_attempt="$2"
|
||||
require_cloud_seal="$3"
|
||||
jobs_dir="$(mktemp -d "${TMPDIR:-/tmp}/dws-release-npm-repair-jobs.XXXXXX")"
|
||||
page=1
|
||||
while :; do
|
||||
jobs_page="$jobs_dir/jobs-$page.json"
|
||||
if ! github_get "repos/$REPOSITORY/actions/runs/$run_id/attempts/$run_attempt/jobs?per_page=100&page=$page" \
|
||||
>"$jobs_page"; then
|
||||
rm -rf "$jobs_dir"
|
||||
return 1
|
||||
fi
|
||||
if ! page_count="$(
|
||||
python3 -c 'import json,sys; print(len(json.load(open(sys.argv[1])).get("jobs", [])))' \
|
||||
"$jobs_page"
|
||||
)"; then
|
||||
rm -rf "$jobs_dir"
|
||||
return 1
|
||||
fi
|
||||
[ "$page_count" -eq 100 ] || break
|
||||
page=$((page + 1))
|
||||
done
|
||||
|
||||
result=0
|
||||
python3 - "$EXPECTED_COMMIT" "$run_id" "$run_attempt" "$TAG" "$require_cloud_seal" "$jobs_dir"/jobs-*.json <<'PY' || result=$?
|
||||
import json
|
||||
import sys
|
||||
|
||||
commit, run_id, run_attempt, tag, require_cloud_seal, *pages = sys.argv[1:]
|
||||
jobs = []
|
||||
for page in pages:
|
||||
with open(page, encoding="utf-8") as handle:
|
||||
jobs.extend(json.load(handle).get("jobs", []))
|
||||
|
||||
def fail(message):
|
||||
print(
|
||||
f"Release run {run_id} attempt {run_attempt} is not safe npm-repair "
|
||||
f"authority for {tag}: {message}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
raise SystemExit(1)
|
||||
|
||||
def one_job(name):
|
||||
matches = [job for job in jobs if job.get("name") == name]
|
||||
if len(matches) != 1:
|
||||
fail(f"expected exactly one job {name!r}, found {len(matches)}")
|
||||
job = matches[0]
|
||||
if (
|
||||
job.get("head_sha") != commit
|
||||
or job.get("status") != "completed"
|
||||
or job.get("conclusion") != "success"
|
||||
):
|
||||
fail(f"required job {name!r} did not succeed at {commit}")
|
||||
return job
|
||||
|
||||
for name in (
|
||||
"release-contract",
|
||||
"Build signed release artifacts",
|
||||
"Verify Apple Developer ID signatures",
|
||||
):
|
||||
one_job(name)
|
||||
|
||||
if require_cloud_seal == "true":
|
||||
one_job("Plan next cloud release")
|
||||
seal = one_job("Seal cloud release tag")
|
||||
seal_steps = [
|
||||
step for step in seal.get("steps", [])
|
||||
if step.get("name") == "Create one immutable annotated release tag"
|
||||
]
|
||||
if (
|
||||
len(seal_steps) != 1
|
||||
or seal_steps[0].get("status") != "completed"
|
||||
or seal_steps[0].get("conclusion") != "success"
|
||||
):
|
||||
fail("cloud release seal step did not succeed")
|
||||
elif require_cloud_seal != "false":
|
||||
fail(f"invalid cloud seal requirement {require_cloud_seal!r}")
|
||||
|
||||
published = one_job("Publish immutable GitHub Release")
|
||||
steps = [
|
||||
step for step in published.get("steps", [])
|
||||
if step.get("name") == "Require immutable published GitHub Release"
|
||||
]
|
||||
if (
|
||||
len(steps) != 1
|
||||
or steps[0].get("status") != "completed"
|
||||
or steps[0].get("conclusion") != "success"
|
||||
):
|
||||
fail("immutable GitHub Release verification step did not succeed")
|
||||
|
||||
channels = [job for job in jobs if job.get("name") == "Publish npm and mirrors"]
|
||||
if len(channels) != 1:
|
||||
fail(f"expected exactly one npm publication job, found {len(channels)}")
|
||||
channel = channels[0]
|
||||
if (
|
||||
channel.get("head_sha") != commit
|
||||
or channel.get("status") != "completed"
|
||||
or channel.get("conclusion") not in {"success", "failure"}
|
||||
):
|
||||
fail("npm publication job is not a completed success/failure at the release commit")
|
||||
PY
|
||||
rm -rf "$jobs_dir"
|
||||
return "$result"
|
||||
}
|
||||
|
||||
push_delivery="$(find_push_delivery || true)"
|
||||
if [ -n "$push_delivery" ]; then
|
||||
printf 'Release workflow delivery verified through exact-tag push run %s: %s -> %s\n' \
|
||||
@@ -282,6 +594,15 @@ if [ -n "$push_delivery" ]; then
|
||||
exit 0
|
||||
fi
|
||||
|
||||
cloud_delivery_identity="$(find_cloud_delivery_identity || true)"
|
||||
if [ -n "$cloud_delivery_identity" ] &&
|
||||
verify_cloud_delivery "$cloud_delivery_identity"; then
|
||||
cloud_delivery_run="$(printf '%s\n' "$cloud_delivery_identity" | cut -f1)"
|
||||
printf 'Release workflow delivery verified through cloud release run %s: %s -> %s\n' \
|
||||
"$cloud_delivery_run" "$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
|
||||
if [ "$MODE" = "channel-repair" ]; then
|
||||
failed_push_identity="$(find_failed_push_delivery || true)"
|
||||
failed_push_delivery="$(printf '%s\n' "$failed_push_identity" | cut -f1)"
|
||||
@@ -292,6 +613,38 @@ if [ "$MODE" = "channel-repair" ]; then
|
||||
"$CHANNEL_REPAIR_TARGET" "$failed_push_delivery" "$failed_push_attempt" "$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
if [ -n "$cloud_delivery_identity" ] &&
|
||||
verify_failed_cloud_delivery_identity "$cloud_delivery_identity"; then
|
||||
failed_cloud_run="$(printf '%s\n' "$cloud_delivery_identity" | cut -f1)"
|
||||
failed_cloud_attempt="$(printf '%s\n' "$cloud_delivery_identity" | cut -f2)"
|
||||
if verify_channel_repair_delivery "$failed_cloud_run" "$failed_cloud_attempt"; then
|
||||
printf 'Release %s channel-repair authority verified through failed cloud release run %s attempt %s: %s -> %s\n' \
|
||||
"$CHANNEL_REPAIR_TARGET" "$failed_cloud_run" "$failed_cloud_attempt" "$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
if [ "$MODE" = "npm-repair" ]; then
|
||||
failed_push_identity="$(find_failed_push_delivery || true)"
|
||||
failed_push_run="$(printf '%s\n' "$failed_push_identity" | cut -f1)"
|
||||
failed_push_attempt="$(printf '%s\n' "$failed_push_identity" | cut -f2)"
|
||||
if [ -n "$failed_push_run" ] &&
|
||||
verify_npm_repair_delivery "$failed_push_run" "$failed_push_attempt" false; then
|
||||
printf 'Release npm-repair authority verified through failed exact-tag push run %s attempt %s: %s -> %s\n' \
|
||||
"$failed_push_run" "$failed_push_attempt" "$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
if [ -n "$cloud_delivery_identity" ] &&
|
||||
verify_failed_cloud_delivery_identity "$cloud_delivery_identity"; then
|
||||
failed_cloud_run="$(printf '%s\n' "$cloud_delivery_identity" | cut -f1)"
|
||||
failed_cloud_attempt="$(printf '%s\n' "$cloud_delivery_identity" | cut -f2)"
|
||||
if verify_npm_repair_delivery "$failed_cloud_run" "$failed_cloud_attempt" true; then
|
||||
printf 'Release npm-repair authority verified through failed cloud release run %s attempt %s: %s -> %s\n' \
|
||||
"$failed_cloud_run" "$failed_cloud_attempt" "$TAG" "$EXPECTED_COMMIT"
|
||||
exit 0
|
||||
fi
|
||||
fi
|
||||
fi
|
||||
|
||||
find_recovery_identity() {
|
||||
@@ -325,7 +678,7 @@ for run in runs:
|
||||
|
||||
recovery_identity="$(find_recovery_identity || true)"
|
||||
[ -n "$recovery_identity" ] || {
|
||||
printf 'Release workflow did not deliver %s at %s through a tag push or protected recovery\n' \
|
||||
printf 'Release workflow did not deliver %s at %s through a tag push, cloud release, or protected recovery\n' \
|
||||
"$TAG" "$EXPECTED_COMMIT" >&2
|
||||
exit 1
|
||||
}
|
||||
|
||||
Executable
+1215
File diff suppressed because it is too large
Load Diff
@@ -1,479 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Run every read shortcut against the real DWS backend.
|
||||
|
||||
This launches the built CLI once per read shortcut. It does not use --mock and
|
||||
does not use --dry-run. Inputs are synthetic but realistic where a shortcut
|
||||
expects a name, date, or query; resource identifiers remain test placeholders
|
||||
when no resource has been created for that command.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as dt
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import shlex
|
||||
import subprocess
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
from shortcut_real_result import (
|
||||
classify_failure,
|
||||
sanitize_result,
|
||||
summarize_failure_categories,
|
||||
classify_real_status,
|
||||
summarize_results,
|
||||
)
|
||||
|
||||
|
||||
ROOT = Path(__file__).resolve().parents[1]
|
||||
MATRIX_PATH = Path("/private/tmp/dws-shortcut-matrix.json")
|
||||
OUT_PATH = ROOT / "docs" / "shortcut-real-read-results.json"
|
||||
BIN = os.environ.get("DWS_REAL_TEST_BIN", "/private/tmp/dws-real-test")
|
||||
DEVAPP_FIXTURE_ID = "678f27ec-4339-49d8-9c49-371b284bf552"
|
||||
DEVAPP_FIXTURE_VERSION_ID = "4743accb-45e2-4bc9-8c96-74fc62ace2e8"
|
||||
CHAT_FIXTURE_OPEN_CONVERSATION_ID = "cid3Jijzhe2aqs9ysOXjhi05g=="
|
||||
CHAT_FIXTURE_GROUP_NAME = "浅曦-kida,Dennis,秋画"
|
||||
CHAT_FIXTURE_DM_OPEN_CONVERSATION_ID = "cidie1367hAfBxqipzE59k5sknHLrHmvYkw98NADhfnjPI="
|
||||
CHAT_FIXTURE_OPEN_MESSAGE_ID = "msgEuOor1PmFBNlx9M06N9z1Q=="
|
||||
CHAT_FIXTURE_OPEN_TASK_ID = "y/wM6Lo+9GbIqtILYPv1BZDcMW+2FgnqskgcpdOiMdM="
|
||||
CALENDAR_FIXTURE_EVENT_ID = "THN4YUtOTlplYU9sZzd2czE4YURLQT09_1784078100000"
|
||||
DOC_FIXTURE_NODE_ID = "P0MALyR8knpgo9GycY7ZlMxlJ3bzYmDO"
|
||||
DOC_FIXTURE_FOLDER_ID = "Amq4vjg89ZOAdqyaSMKpApXdW3kdP0wQ"
|
||||
DOC_FIXTURE_EXPORT_JOB_ID = "29346731713"
|
||||
SHEET_FIXTURE_NODE_ID = "mweZ92PV6O36dZbnsMZx70ylJxEKBD6p"
|
||||
DING_FIXTURE_OPEN_DING_ID = "5D73E3AC29C780072D1CD56C6874ACC2"
|
||||
CONTACT_FIXTURE_MOBILE = "13161187007"
|
||||
AITABLE_FIXTURE_BASE_ID = "gpG2NdyVXQyZ0OmoSbd1vbA6JMwvDqPk"
|
||||
AITABLE_FIXTURE_TABLE_ID = "hERWDMS"
|
||||
AITABLE_FIXTURE_VIEW_ID = "qvGDAH2"
|
||||
AITABLE_FIXTURE_FORM_VIEW_ID = "lmeV1cb"
|
||||
AITABLE_FIXTURE_FIELD_ID = "01ZM8y7"
|
||||
AITABLE_FIXTURE_RECORD_ID = "1015oH3OXy"
|
||||
AITABLE_FIXTURE_DASHBOARD_ID = "KY9tlWg5NEHgfs8WT6IO2"
|
||||
AITABLE_FIXTURE_CHART_ID = "widget-dlxFo0tNSImp4ITpDn5fQ"
|
||||
|
||||
HELD_CASES = {
|
||||
("devapp", "+credentials-get"):
|
||||
"该命令会读取真实应用凭证/密钥;不能用真实 app 自动执行。当前仅用占位 ID 验证负向路径,真实成功需人工在安全环境单独确认。",
|
||||
}
|
||||
|
||||
|
||||
def ensure_matrix() -> dict:
|
||||
env = os.environ.copy()
|
||||
env.setdefault("GOCACHE", "/private/tmp/dws_gocache")
|
||||
raw = subprocess.check_output(
|
||||
["go", "run", "./scripts/gen_shortcut_test_matrix.go"],
|
||||
cwd=ROOT,
|
||||
env=env,
|
||||
text=True,
|
||||
)
|
||||
MATRIX_PATH.write_text(raw, encoding="utf-8")
|
||||
return json.loads(raw)
|
||||
|
||||
|
||||
def replace_flag_values(args: list[str], service: str, command: str) -> list[str]:
|
||||
today = dt.date.today()
|
||||
start = (dt.datetime.now() + dt.timedelta(days=1)).replace(hour=10, minute=0, second=0, microsecond=0).isoformat() + "+08:00"
|
||||
end = (dt.datetime.now() + dt.timedelta(days=1)).replace(hour=11, minute=0, second=0, microsecond=0).isoformat() + "+08:00"
|
||||
no_id = "DWSREALREADNOSUCHID0000000000000"
|
||||
no_conv = "cidDWSREALREADNOSUCHCONV"
|
||||
day = str(today)
|
||||
yesterday = str(today - dt.timedelta(days=1))
|
||||
datetime_start = dt.datetime.now().replace(hour=10, minute=0, second=0, microsecond=0).strftime("%Y-%m-%d %H:%M:%S")
|
||||
datetime_end = (dt.datetime.now() + dt.timedelta(hours=1)).replace(minute=0, second=0, microsecond=0).strftime("%Y-%m-%d %H:%M:%S")
|
||||
replacements = {
|
||||
"name": "DWS shortcut 真实测试",
|
||||
"query": "测试",
|
||||
"keyword": "测试",
|
||||
"q": "测试",
|
||||
"text": "测试",
|
||||
"title": "测试",
|
||||
"phone": "13000000000",
|
||||
"mobile": "13000000000",
|
||||
"user": "冬翔",
|
||||
"users": "冬翔",
|
||||
"to": "冬翔",
|
||||
"with": "冬翔",
|
||||
"who": "冬翔",
|
||||
"dept": "模型算法",
|
||||
"dept-id": "842379556",
|
||||
"department-id": "842379556",
|
||||
"start": start,
|
||||
"end": end,
|
||||
"from": str(today - dt.timedelta(days=7)),
|
||||
"to-date": str(today),
|
||||
"date": str(today),
|
||||
"time": datetime_start,
|
||||
"days": "7",
|
||||
"types": "leave",
|
||||
"columns": "1001",
|
||||
"limit": "10",
|
||||
"page": "1",
|
||||
"page-size": "10",
|
||||
"cursor": "0",
|
||||
"calendar-id": "primary",
|
||||
"event": no_id,
|
||||
"type": "ALL",
|
||||
"types": "leave",
|
||||
"columns": "1001",
|
||||
"role-types": "executor",
|
||||
"status": "false",
|
||||
"artifacts": "basic",
|
||||
"direction": "older",
|
||||
"file-types": "alidoc",
|
||||
"order-by": "name",
|
||||
"order": "asc",
|
||||
"space-id": "1",
|
||||
"space": "测试",
|
||||
"base": no_id,
|
||||
"base-id": no_id,
|
||||
"table": no_id,
|
||||
"table-id": no_id,
|
||||
"view-id": no_id,
|
||||
"record-id": no_id,
|
||||
"record-ids": no_id,
|
||||
"field-id": no_id,
|
||||
"dashboard-id": no_id,
|
||||
"chart-id": no_id,
|
||||
"workflow-id": no_id,
|
||||
"node": no_id,
|
||||
"doc": no_id,
|
||||
"folder": no_id,
|
||||
"workspace": no_id,
|
||||
"group": no_conv,
|
||||
"conversation-id": no_conv,
|
||||
"open-conversation-id": no_conv,
|
||||
"message-id": no_id,
|
||||
"msg-id": no_id,
|
||||
"id": no_id,
|
||||
"session-id": no_id,
|
||||
"process-instance-id": no_id,
|
||||
"task-id": no_id,
|
||||
"template-id": no_id,
|
||||
"mail-id": no_id,
|
||||
"filters": "{}",
|
||||
"sort": "[]",
|
||||
}
|
||||
if service == "calendar" and command == "+free-slots":
|
||||
replacements["from"] = "9"
|
||||
replacements["to"] = "18"
|
||||
if service == "calendar":
|
||||
replacements["event"] = CALENDAR_FIXTURE_EVENT_ID
|
||||
replacements["cursor"] = ""
|
||||
if command == "+freebusy":
|
||||
replacements["users"] = "103262"
|
||||
if service == "doc" and command == "+comment-list":
|
||||
replacements["type"] = "global"
|
||||
if service == "doc":
|
||||
replacements["node"] = DOC_FIXTURE_NODE_ID
|
||||
replacements["doc"] = DOC_FIXTURE_NODE_ID
|
||||
replacements["folder"] = DOC_FIXTURE_FOLDER_ID
|
||||
replacements["job-id"] = DOC_FIXTURE_EXPORT_JOB_ID
|
||||
if service == "drive":
|
||||
replacements["node"] = DOC_FIXTURE_NODE_ID
|
||||
replacements["folder"] = DOC_FIXTURE_FOLDER_ID
|
||||
if service == "todo":
|
||||
replacements["task-id"] = "55119034912"
|
||||
replacements["size"] = "10"
|
||||
if service == "aitable":
|
||||
replacements["name"] = "Real共创版设备去向登记"
|
||||
replacements["base"] = AITABLE_FIXTURE_BASE_ID
|
||||
replacements["base-id"] = AITABLE_FIXTURE_BASE_ID
|
||||
replacements["table"] = AITABLE_FIXTURE_TABLE_ID
|
||||
replacements["table-id"] = AITABLE_FIXTURE_TABLE_ID
|
||||
replacements["view-id"] = AITABLE_FIXTURE_VIEW_ID
|
||||
replacements["view-ids"] = AITABLE_FIXTURE_VIEW_ID
|
||||
replacements["field-id"] = AITABLE_FIXTURE_FIELD_ID
|
||||
replacements["field-ids"] = AITABLE_FIXTURE_FIELD_ID
|
||||
replacements["record-id"] = AITABLE_FIXTURE_RECORD_ID
|
||||
replacements["record-ids"] = AITABLE_FIXTURE_RECORD_ID
|
||||
replacements["dashboard-id"] = AITABLE_FIXTURE_DASHBOARD_ID
|
||||
replacements["chart-id"] = AITABLE_FIXTURE_CHART_ID
|
||||
if command.startswith("+form-"):
|
||||
replacements["view-id"] = AITABLE_FIXTURE_FORM_VIEW_ID
|
||||
if command == "+resolve-table":
|
||||
replacements["name"] = "Mac Mini"
|
||||
if service == "ding" and command == "+list":
|
||||
replacements["type"] = "ALL"
|
||||
if service == "ding" and command == "+receiver-status":
|
||||
replacements["ding-id"] = DING_FIXTURE_OPEN_DING_ID
|
||||
if service == "contact" and command == "+list-sub-depts":
|
||||
replacements["dept"] = "842379556"
|
||||
if service == "contact":
|
||||
replacements["name"] = "冬翔"
|
||||
if command == "+by-mobile":
|
||||
replacements["mobile"] = CONTACT_FIXTURE_MOBILE
|
||||
if command == "+resolve-dept":
|
||||
replacements["name"] = "模型算法"
|
||||
if service == "oa":
|
||||
now_ms = int(dt.datetime.now().timestamp() * 1000)
|
||||
replacements["start"] = str(now_ms - 7 * 24 * 60 * 60 * 1000)
|
||||
replacements["end"] = str(now_ms)
|
||||
replacements["page"] = "1"
|
||||
replacements["limit"] = "10"
|
||||
if service == "report":
|
||||
replacements["start"] = (dt.datetime.now() - dt.timedelta(days=7)).replace(microsecond=0).isoformat() + "+08:00"
|
||||
replacements["end"] = dt.datetime.now().replace(microsecond=0).isoformat() + "+08:00"
|
||||
replacements["modified-start"] = replacements["start"]
|
||||
replacements["modified-end"] = replacements["end"]
|
||||
if service == "attendance":
|
||||
replacements.update({
|
||||
"user": "202397",
|
||||
"users": "202397",
|
||||
"staff-ids": "202397",
|
||||
"operator-staff-id": "202397",
|
||||
"leave-code": "731ed089-62ff-4734-a6c7-3c8fcc8294fc",
|
||||
"leave-names": "年假",
|
||||
"start": yesterday,
|
||||
"end": day,
|
||||
"from": yesterday,
|
||||
"to-date": day,
|
||||
"date": day,
|
||||
})
|
||||
if command in {"+get-checkin-record", "+query-report-data", "+query-report-leave"}:
|
||||
replacements["start"] = datetime_start
|
||||
replacements["end"] = datetime_end
|
||||
if command == "+get-approve-template":
|
||||
replacements["type"] = "leave"
|
||||
if command == "+search-group":
|
||||
replacements["type"] = "FIXED"
|
||||
if service == "devapp":
|
||||
replacements["unified-app-id"] = DEVAPP_FIXTURE_ID
|
||||
replacements["version-id"] = DEVAPP_FIXTURE_VERSION_ID
|
||||
replacements["cursor"] = ""
|
||||
if service == "mail":
|
||||
replacements["email"] = "xinyang.dxy@alibaba-inc.com"
|
||||
replacements["folder"] = "2"
|
||||
replacements["query"] = "subject:测试"
|
||||
replacements["keyword"] = "董鑫阳"
|
||||
replacements["employee-no"] = "202397"
|
||||
replacements["size"] = "10"
|
||||
replacements["cursor"] = ""
|
||||
if command == "+find-mail-user":
|
||||
replacements["query"] = "董鑫阳"
|
||||
if service == "chat":
|
||||
replacements["group"] = CHAT_FIXTURE_OPEN_CONVERSATION_ID
|
||||
replacements["conversation-id"] = CHAT_FIXTURE_OPEN_CONVERSATION_ID
|
||||
replacements["open-conversation-id"] = CHAT_FIXTURE_OPEN_CONVERSATION_ID
|
||||
replacements["msg-ids"] = CHAT_FIXTURE_OPEN_MESSAGE_ID
|
||||
replacements["message-id"] = CHAT_FIXTURE_OPEN_MESSAGE_ID
|
||||
replacements["msg-id"] = CHAT_FIXTURE_OPEN_MESSAGE_ID
|
||||
replacements["open-task-id"] = CHAT_FIXTURE_OPEN_TASK_ID
|
||||
if command == "+group-members":
|
||||
replacements["group"] = CHAT_FIXTURE_GROUP_NAME
|
||||
if command == "+messages-read-status":
|
||||
replacements["conversation-id"] = CHAT_FIXTURE_DM_OPEN_CONVERSATION_ID
|
||||
replacements["users"] = "冬翔"
|
||||
if command == "+messages-resource-url":
|
||||
replacements["type"] = "mediaId"
|
||||
if command == "+bot-find":
|
||||
replacements["cursor"] = ""
|
||||
if service == "sheet" and command == "+list-sheets":
|
||||
replacements["node"] = SHEET_FIXTURE_NODE_ID
|
||||
out = list(args)
|
||||
if "--format" not in out:
|
||||
out.extend(["--format", "json"])
|
||||
i = 0
|
||||
while i < len(out) - 1:
|
||||
if out[i].startswith("--"):
|
||||
key = out[i][2:]
|
||||
if key in replacements and not out[i + 1].startswith("--"):
|
||||
out[i + 1] = replacements[key]
|
||||
i += 2
|
||||
continue
|
||||
i += 1
|
||||
return out
|
||||
|
||||
|
||||
def shell_join(cmd: list[str]) -> str:
|
||||
return " ".join(shlex.quote(x) for x in cmd)
|
||||
|
||||
|
||||
def drop_flag(args: list[str], flag: str) -> list[str]:
|
||||
out: list[str] = []
|
||||
i = 0
|
||||
while i < len(args):
|
||||
if args[i] == flag:
|
||||
i += 2 if i + 1 < len(args) and not args[i + 1].startswith("--") else 1
|
||||
continue
|
||||
out.append(args[i])
|
||||
i += 1
|
||||
return out
|
||||
|
||||
|
||||
def adjust_command_args(args: list[str], service: str, command: str) -> list[str]:
|
||||
if service == "chat" and command == "+chat-messages":
|
||||
return drop_flag(args, "--user")
|
||||
if service == "chat" and command == "+messages-list-direct":
|
||||
out = drop_flag(args, "--open-dingtalk-id")
|
||||
for i in range(len(out) - 1):
|
||||
if out[i] == "--user":
|
||||
out[i + 1] = "103262"
|
||||
return out
|
||||
if service == "calendar" and command == "+freebusy":
|
||||
return drop_flag(args, "--rooms")
|
||||
if service == "devapp" and command == "+permission-list":
|
||||
out = drop_flag(args, "--scope-value")
|
||||
out = drop_flag(out, "--scope-type")
|
||||
out = drop_flag(out, "--api-status")
|
||||
for i in range(len(out) - 1):
|
||||
if out[i] == "--auth-status":
|
||||
out[i + 1] = "ALL"
|
||||
return out
|
||||
if service == "doc" and command == "+search":
|
||||
out = drop_flag(args, "--extensions")
|
||||
out = drop_flag(out, "--created-from")
|
||||
out = drop_flag(out, "--created-to")
|
||||
out = drop_flag(out, "--visited-from")
|
||||
out = drop_flag(out, "--visited-to")
|
||||
out = drop_flag(out, "--creator-uids")
|
||||
out = drop_flag(out, "--editor-uids")
|
||||
out = drop_flag(out, "--mentioned-uids")
|
||||
out = drop_flag(out, "--workspace-ids")
|
||||
return out
|
||||
if service == "doc" and command == "+comment-list":
|
||||
out = drop_flag(args, "--cursor")
|
||||
out = drop_flag(out, "--resolve-status")
|
||||
return out
|
||||
if service == "doc" and command == "+list":
|
||||
out = drop_flag(args, "--workspace")
|
||||
out = drop_flag(out, "--cursor")
|
||||
return out
|
||||
if service == "report" and command == "+outbox-list":
|
||||
return drop_flag(args, "--template-name")
|
||||
return args
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--service", action="append", help="Only run shortcuts from this service; may repeat")
|
||||
parser.add_argument("--command", action="append", help="Only run this shortcut command; may repeat")
|
||||
parser.add_argument("--failed-only", action="store_true", help="Only rerun shortcuts currently marked non-success in the output report")
|
||||
ns = parser.parse_args()
|
||||
services = set(ns.service or [])
|
||||
commands = set(ns.command or [])
|
||||
|
||||
matrix = ensure_matrix()
|
||||
rows = [r for r in matrix["results"] if r.get("risk") == "read"]
|
||||
if services:
|
||||
rows = [r for r in rows if r["service"] in services]
|
||||
if commands:
|
||||
rows = [r for r in rows if r["command"] in commands]
|
||||
if ns.failed_only and OUT_PATH.exists():
|
||||
existing = json.loads(OUT_PATH.read_text(encoding="utf-8"))
|
||||
failed = {
|
||||
(r.get("service"), r.get("command"))
|
||||
for r in existing.get("results", [])
|
||||
if r.get("status") != "real-ok"
|
||||
}
|
||||
rows = [r for r in rows if (r["service"], r["command"]) in failed]
|
||||
results = []
|
||||
summary = {"total": len(rows), "ok": 0, "error": 0, "timeout": 0, "held": 0}
|
||||
for idx, r in enumerate(rows, 1):
|
||||
key = (r["service"], r["command"])
|
||||
if key in HELD_CASES:
|
||||
summary["held"] += 1
|
||||
item = {
|
||||
"service": r["service"],
|
||||
"command": r["command"],
|
||||
"risk": r["risk"],
|
||||
"method": "held; sensitive credential read",
|
||||
"status": "held",
|
||||
"input": "",
|
||||
"args": [],
|
||||
"stdout": "",
|
||||
"stderr": HELD_CASES[key],
|
||||
"exit_code": None,
|
||||
"duration_ms": 0,
|
||||
}
|
||||
item = sanitize_result(item)
|
||||
category, fixability, note = classify_failure(item)
|
||||
item["failure_category"] = category
|
||||
item["fixability"] = fixability
|
||||
item["diagnosis"] = note
|
||||
results.append(item)
|
||||
continue
|
||||
args = adjust_command_args(replace_flag_values(r["args"], r["service"], r["command"]), r["service"], r["command"])
|
||||
cmd = [BIN] + args
|
||||
started = time.time()
|
||||
status = "real-error"
|
||||
stdout = ""
|
||||
stderr = ""
|
||||
exit_code = None
|
||||
try:
|
||||
p = subprocess.run(cmd, text=True, capture_output=True, timeout=30)
|
||||
stdout = p.stdout.strip()
|
||||
stderr = p.stderr.strip()
|
||||
exit_code = p.returncode
|
||||
status = classify_real_status(exit_code, stdout)
|
||||
if status == "real-ok":
|
||||
status = "real-ok"
|
||||
summary["ok"] += 1
|
||||
else:
|
||||
summary["error"] += 1
|
||||
except subprocess.TimeoutExpired as e:
|
||||
status = "timeout"
|
||||
summary["timeout"] += 1
|
||||
if isinstance(e.stdout, str):
|
||||
stdout = e.stdout.strip()
|
||||
if isinstance(e.stderr, str):
|
||||
stderr = e.stderr.strip()
|
||||
duration_ms = int((time.time() - started) * 1000)
|
||||
item = {
|
||||
"service": r["service"],
|
||||
"command": r["command"],
|
||||
"risk": r["risk"],
|
||||
"method": "real-backend-read; no --mock; no --dry-run",
|
||||
"status": status,
|
||||
"input": shell_join(cmd),
|
||||
"args": args,
|
||||
"stdout": stdout,
|
||||
"stderr": stderr,
|
||||
"exit_code": exit_code,
|
||||
"duration_ms": duration_ms,
|
||||
}
|
||||
item = sanitize_result(item)
|
||||
category, fixability, note = classify_failure(item)
|
||||
if category != "passed":
|
||||
item["failure_category"] = category
|
||||
item["fixability"] = fixability
|
||||
item["diagnosis"] = note
|
||||
results.append(item)
|
||||
if idx % 25 == 0 or idx == len(rows):
|
||||
print(
|
||||
f"progress {idx}/{len(rows)} ok={summary['ok']} "
|
||||
f"error={summary['error']} timeout={summary['timeout']} held={summary['held']}",
|
||||
flush=True,
|
||||
)
|
||||
if not services and not commands and not ns.failed_only:
|
||||
OUT_PATH.write_text(json.dumps({
|
||||
"generated_at": dt.datetime.now().isoformat(),
|
||||
"summary": summary,
|
||||
"failure_categories": summarize_failure_categories(results),
|
||||
"results": results,
|
||||
}, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
if OUT_PATH.exists() and (services or commands or ns.failed_only):
|
||||
existing = json.loads(OUT_PATH.read_text(encoding="utf-8"))
|
||||
replace_keys = {(r["service"], r["command"]) for r in results}
|
||||
merged = [
|
||||
r for r in existing.get("results", [])
|
||||
if (r.get("service"), r.get("command")) not in replace_keys
|
||||
]
|
||||
merged.extend(results)
|
||||
else:
|
||||
merged = results
|
||||
summary = summarize_results(merged, include_held=True)
|
||||
OUT_PATH.write_text(json.dumps({
|
||||
"generated_at": dt.datetime.now().isoformat(),
|
||||
"summary": summary,
|
||||
"failure_categories": summarize_failure_categories(merged),
|
||||
"results": merged,
|
||||
}, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
print(f"saved {OUT_PATH} batch={summarize_results(results, include_held=True)} merged={summary}")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -28,7 +28,7 @@ func TestDWSDocsCommandTreeCoverage(t *testing.T) {
|
||||
t.Skip("no command paths parsed from docs/dws (directory may not exist)")
|
||||
}
|
||||
|
||||
index := buildCommandIndex(app.NewRootCommand())
|
||||
index := buildCommandIndex(app.NewSchemaSourceRootCommand())
|
||||
missing := make([]string, 0)
|
||||
for _, path := range docPaths {
|
||||
if _, ok := index[path]; ok {
|
||||
@@ -52,7 +52,7 @@ func TestDWSDocsLocalFlagsCoverage(t *testing.T) {
|
||||
}
|
||||
docLeafSet := leafCommandSet(docPaths)
|
||||
|
||||
index := buildCommandIndex(app.NewRootCommand())
|
||||
index := buildCommandIndex(app.NewSchemaSourceRootCommand())
|
||||
missing := make([]string, 0)
|
||||
fallbackMatched := 0
|
||||
explicitMatched := 0
|
||||
|
||||
@@ -9,6 +9,12 @@ import (
|
||||
)
|
||||
|
||||
func TestMain(m *testing.M) {
|
||||
configDir, err := os.MkdirTemp("", "dws-cli-test-config-")
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
os.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
|
||||
// Set an empty catalog fixture so that EnvironmentLoader does not
|
||||
// attempt live discovery (which would hang on unreachable MCP endpoints).
|
||||
// Tests that construct app root commands must remain serial because root
|
||||
@@ -17,5 +23,6 @@ func TestMain(m *testing.M) {
|
||||
os.Setenv(cli.CatalogFixtureEnv, absFixture)
|
||||
|
||||
code := m.Run()
|
||||
_ = os.RemoveAll(configDir)
|
||||
os.Exit(code)
|
||||
}
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
@@ -45,6 +46,29 @@ var expectedReleaseAdmissionContexts = []string{
|
||||
"Mock MCP",
|
||||
}
|
||||
|
||||
func TestPackageManagerVersionVerificationReadsRawBinary(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
scriptPath, err := filepath.Abs(filepath.Join("..", "..", "scripts", "release", "verify-package-managers.sh"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(verify-package-managers.sh) error = %v", err)
|
||||
}
|
||||
data, err := os.ReadFile(scriptPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", scriptPath, err)
|
||||
}
|
||||
script := string(data)
|
||||
for _, binary := range []string{`"$vendor_bin"`, `"$prefix/bin/dws"`} {
|
||||
want := `LC_ALL=C grep -aFq "v$EXPECTED_VERSION" ` + binary
|
||||
if !strings.Contains(script, want) {
|
||||
t.Errorf("package-manager verifier is missing raw binary marker check %q", want)
|
||||
}
|
||||
}
|
||||
if strings.Contains(script, `strings "$vendor_bin"`) || strings.Contains(script, `strings "$prefix/bin/dws"`) {
|
||||
t.Fatal("package-manager verifier still requires the version marker to occupy a strings(1) line")
|
||||
}
|
||||
}
|
||||
|
||||
func seedDistArchive(t *testing.T, path string) {
|
||||
t.Helper()
|
||||
file, err := os.Create(path)
|
||||
@@ -636,6 +660,33 @@ func releaseWorkflowSection(t *testing.T, workflow, startMarker, endMarker strin
|
||||
return workflow[start : start+len(startMarker)+end]
|
||||
}
|
||||
|
||||
func releaseWorkflowRunScript(t *testing.T, workflow, stepName, nextStepName string) string {
|
||||
t.Helper()
|
||||
section := releaseWorkflowSection(
|
||||
t,
|
||||
workflow,
|
||||
" - name: "+stepName+"\n",
|
||||
"\n - name: "+nextStepName+"\n",
|
||||
)
|
||||
const runMarker = " run: |\n"
|
||||
start := strings.Index(section, runMarker)
|
||||
if start == -1 {
|
||||
t.Fatalf("release workflow step %q is missing a run block", stepName)
|
||||
}
|
||||
|
||||
lines := strings.Split(section[start+len(runMarker):], "\n")
|
||||
for i, line := range lines {
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
if !strings.HasPrefix(line, " ") {
|
||||
t.Fatalf("release workflow step %q has an unexpected run indentation: %q", stepName, line)
|
||||
}
|
||||
lines[i] = strings.TrimPrefix(line, " ")
|
||||
}
|
||||
return strings.Join(lines, "\n")
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowUsesDedicatedGovernanceIdentity(t *testing.T) {
|
||||
t.Parallel()
|
||||
workflow := readReleaseWorkflow(t)
|
||||
@@ -758,7 +809,7 @@ func TestReleaseWorkflowGovernancePreflightCannotPublish(t *testing.T) {
|
||||
"EXPECTED_REPOSITORY: DingTalk-Real-AI/dingtalk-workspace-cli",
|
||||
`DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}`,
|
||||
`test "$PREFLIGHT_COMMIT" = "$GITHUB_SHA"`,
|
||||
`ref: ${{ inputs.governance_preflight_commit }}`,
|
||||
`ref: ${{ needs.dispatch-contract.outputs.mode == 'create_release' && github.sha || inputs.governance_preflight_commit }}`,
|
||||
"persist-credentials: false",
|
||||
"governance preflight cannot be combined with npm repair",
|
||||
} {
|
||||
@@ -792,6 +843,9 @@ func TestReleaseWorkflowGovernancePreflightCannotPublish(t *testing.T) {
|
||||
for _, required := range []string{
|
||||
"needs: dispatch-contract",
|
||||
"needs.dispatch-contract.outputs.mode == 'repair_npm'",
|
||||
`ref: ` + "`tags/withdrawn/${version}`",
|
||||
"was withdrawn and cannot be repaired",
|
||||
"verify-release-workflow-delivery.sh --npm-repair",
|
||||
} {
|
||||
if !strings.Contains(repair, required) {
|
||||
t.Errorf("npm repair dispatch contract is missing %q", required)
|
||||
@@ -799,6 +853,222 @@ func TestReleaseWorkflowGovernancePreflightCannotPublish(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowPlansAndSealsCurrentMainInTheCloud(t *testing.T) {
|
||||
t.Parallel()
|
||||
workflow := readReleaseWorkflow(t)
|
||||
planStart := strings.Index(workflow, " release-plan:\n")
|
||||
sealStart := strings.Index(workflow, " seal-release:\n")
|
||||
if planStart == -1 || sealStart == -1 || planStart >= sealStart {
|
||||
t.Fatal("cloud release plan and seal jobs are missing or out of order")
|
||||
}
|
||||
plan := workflow[planStart:sealStart]
|
||||
seal := workflow[sealStart:]
|
||||
|
||||
for _, required := range []string{
|
||||
"release_operation:",
|
||||
"- none",
|
||||
"- plan",
|
||||
"- publish",
|
||||
"release_channel:",
|
||||
"release_bump:",
|
||||
"release_confirmation:",
|
||||
`release_flow + npm_repair + gitee_repair + oss_repair + governance + recovery`,
|
||||
`echo "mode=plan_release"`,
|
||||
`echo "mode=create_release"`,
|
||||
`release_confirmation must be exactly: PUBLISH $RELEASE_CHANNEL`,
|
||||
`needs.dispatch-contract.outputs.mode == 'plan_release'`,
|
||||
`needs.governance-preflight.result == 'success'`,
|
||||
"actions: read",
|
||||
"contents: read",
|
||||
`github.event.repository.default_branch`,
|
||||
`GITHUB_REPOSITORY" = "$EXPECTED_REPOSITORY`,
|
||||
`GITHUB_REF" = "refs/heads/$DEFAULT_BRANCH`,
|
||||
`ref: ${{ github.sha }}`,
|
||||
"persist-credentials: false",
|
||||
`refs/remotes/origin/main)" = "$GITHUB_SHA`,
|
||||
"next-release-version.sh",
|
||||
`'refs/tags/v*' 'refs/tags/withdrawn/v*'`,
|
||||
"release ref manifest is empty after fetching allocated tags",
|
||||
"refs_fingerprint",
|
||||
"Validate the candidate release contract before sealing",
|
||||
"release-contract.sh",
|
||||
"Require delivered previous stable baseline before sealing",
|
||||
"Require delivered beta before sealing stable",
|
||||
} {
|
||||
if !strings.Contains(workflow, required) {
|
||||
t.Errorf("cloud release contract is missing %q", required)
|
||||
}
|
||||
}
|
||||
if strings.Contains(plan, "contents: write") {
|
||||
t.Error("cloud release planning must remain read-only")
|
||||
}
|
||||
if strings.Contains(plan, "refs/tags/v refs/tags/withdrawn/v") {
|
||||
t.Error("cloud release planning must use wildcard ref patterns that match the seal API prefixes")
|
||||
}
|
||||
|
||||
for _, required := range []string{
|
||||
"name: Seal cloud release tag",
|
||||
"contents: write",
|
||||
"name: Create one immutable annotated release tag",
|
||||
`branch.data.commit.sha !== commit`,
|
||||
`actualFingerprint !== expectedFingerprint`,
|
||||
`github.rest.git.createTag`,
|
||||
`github.rest.git.createRef`,
|
||||
`ref: ` + "`refs/tags/${version}`",
|
||||
"`Release-Run: ${context.runId}`",
|
||||
"`Requested-By: ${context.actor}`",
|
||||
"`Requested-By-ID: ${context.payload.sender?.id || \"\"}`",
|
||||
"`Sealed-Commit: ${commit}`",
|
||||
"`Workflow-Commit: ${context.sha}`",
|
||||
"`Allocation-Fingerprint: ${expectedFingerprint}`",
|
||||
} {
|
||||
if !strings.Contains(seal, required) {
|
||||
t.Errorf("cloud release seal is missing %q", required)
|
||||
}
|
||||
}
|
||||
for _, forbidden := range []string{
|
||||
"actions/checkout",
|
||||
"github.rest.git.updateRef",
|
||||
"github.rest.git.deleteRef",
|
||||
"git push",
|
||||
"--force",
|
||||
} {
|
||||
if strings.Contains(seal, forbidden) {
|
||||
t.Errorf("write-capable cloud seal must not contain %q", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseFingerprintRefPatternsMatchAllAllocatedTags(t *testing.T) {
|
||||
t.Parallel()
|
||||
repo := t.TempDir()
|
||||
mustRun(t, repo, "git", "init", "-b", "main")
|
||||
mustRun(t, repo, "git", "config", "user.name", "Release Fingerprint Test")
|
||||
mustRun(t, repo, "git", "config", "user.email", "release-fingerprint@example.com")
|
||||
mustWriteFile(t, filepath.Join(repo, "tracked"), []byte("fixture\n"), 0o644)
|
||||
mustRun(t, repo, "git", "add", "tracked")
|
||||
mustRun(t, repo, "git", "commit", "-m", "fixture")
|
||||
|
||||
allocatedTags := []string{
|
||||
"v1.0.52",
|
||||
"v1.0.53-beta.5",
|
||||
"withdrawn/v1.0.51",
|
||||
}
|
||||
for _, tag := range allocatedTags {
|
||||
mustRun(t, repo, "git", "tag", "-a", tag, "-m", "Release "+tag)
|
||||
}
|
||||
mustRun(t, repo, "git", "tag", "-a", "release/v1.0.52", "-m", "unrelated namespace")
|
||||
legacy := exec.Command(
|
||||
"git", "for-each-ref", "--format=%(refname)=%(objectname)",
|
||||
"refs/tags/v", "refs/tags/withdrawn/v",
|
||||
)
|
||||
legacy.Dir = repo
|
||||
legacyOutput, err := legacy.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("legacy git for-each-ref error = %v\noutput:\n%s", err, legacyOutput)
|
||||
}
|
||||
if strings.TrimSpace(string(legacyOutput)) != "" {
|
||||
t.Fatalf("legacy component patterns unexpectedly matched flat release refs:\n%s", legacyOutput)
|
||||
}
|
||||
|
||||
cmd := exec.Command(
|
||||
"git", "for-each-ref", "--format=%(refname)=%(objectname)",
|
||||
"refs/tags/v*", "refs/tags/withdrawn/v*",
|
||||
)
|
||||
cmd.Dir = repo
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("git for-each-ref error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
got := strings.Split(strings.TrimSpace(string(output)), "\n")
|
||||
sort.Strings(got)
|
||||
|
||||
want := make([]string, 0, len(allocatedTags))
|
||||
for _, tag := range allocatedTags {
|
||||
object := strings.TrimSpace(mustOutput(t, repo, "git", "rev-parse", "refs/tags/"+tag))
|
||||
want = append(want, "refs/tags/"+tag+"="+object)
|
||||
}
|
||||
sort.Strings(want)
|
||||
if strings.Join(got, "\n") != strings.Join(want, "\n") {
|
||||
t.Fatalf("release ref set differs from the seal API set\ngot:\n%s\nwant:\n%s", strings.Join(got, "\n"), strings.Join(want, "\n"))
|
||||
}
|
||||
workflowDigest := sha256.Sum256([]byte(strings.Join(got, "\n") + "\n"))
|
||||
sealDigest := sha256.Sum256([]byte(strings.Join(want, "\n") + "\n"))
|
||||
if workflowDigest != sealDigest {
|
||||
t.Fatalf("release ref fingerprint differs from seal fingerprint: workflow=%x seal=%x", workflowDigest, sealDigest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowAcceptsGuardedLocalTagMetadata(t *testing.T) {
|
||||
t.Parallel()
|
||||
workflow := readReleaseWorkflow(t)
|
||||
|
||||
for _, required := range []string{
|
||||
`const cloudOnlyKeys = [`,
|
||||
`const cloudKeys = ["Channel", ...cloudOnlyKeys];`,
|
||||
`const hasAnyCloudMetadata = cloudOnlyKeys.some((key) => tagFields.has(key));`,
|
||||
`const isCloudSeal = cloudKeys.every((key) => tagFields.has(key));`,
|
||||
} {
|
||||
if !strings.Contains(workflow, required) {
|
||||
t.Errorf("local tag metadata compatibility is missing %q", required)
|
||||
}
|
||||
}
|
||||
if strings.Contains(workflow, `const hasAnyCloudMetadata = cloudKeys.some((key) => tagFields.has(key));`) {
|
||||
t.Error("Channel-only guarded local tags must not be classified as partial cloud seals")
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowRequiresOSSOnlyWhenMirrorIsEnabled(t *testing.T) {
|
||||
t.Parallel()
|
||||
workflow := readReleaseWorkflow(t)
|
||||
releaseContract := releaseWorkflowSection(t, workflow, " release-contract:\n", "\n release:\n")
|
||||
targetAuthority := releaseWorkflowSection(
|
||||
t,
|
||||
releaseContract,
|
||||
" - name: Resolve and verify exact release target\n",
|
||||
"\n - name: Check out repository\n",
|
||||
)
|
||||
ossStep := releaseWorkflowSection(
|
||||
t,
|
||||
workflow,
|
||||
" - name: Sync release artifacts to China OSS mirror\n",
|
||||
"\n mirror-gitee-release:\n",
|
||||
)
|
||||
|
||||
for _, required := range []string{
|
||||
`if: ${{ needs.release-contract.outputs.oss_mirror == 'enabled' }}`,
|
||||
`run: ./scripts/release/sync-to-oss.sh`,
|
||||
`DWS_REQUIRE_OSS: "1"`,
|
||||
} {
|
||||
if !strings.Contains(ossStep, required) {
|
||||
t.Errorf("opt-in OSS publication is missing %q", required)
|
||||
}
|
||||
}
|
||||
for _, required := range []string{
|
||||
`OSS_MIRROR: ${{ vars.ENABLE_OSS_MIRROR == 'true' && 'enabled' || 'deferred' }}`,
|
||||
`OSS-Mirror: ${ossMirror}`,
|
||||
`core.setOutput("oss_mirror", ossMirror);`,
|
||||
} {
|
||||
if !strings.Contains(workflow, required) {
|
||||
t.Errorf("immutable OSS release policy is missing %q", required)
|
||||
}
|
||||
}
|
||||
if strings.Contains(ossStep, "vars.ENABLE_OSS_MIRROR") {
|
||||
t.Error("channel publication must use the immutable tag policy, not the current repository variable")
|
||||
}
|
||||
for _, required := range []string{
|
||||
`const ossMirror = tagFields.has("OSS-Mirror")`,
|
||||
`? tagFields.get("OSS-Mirror")`,
|
||||
`: "enabled";`,
|
||||
`!["enabled", "deferred"].includes(ossMirror)`,
|
||||
`core.setOutput("oss_mirror", ossMirror);`,
|
||||
} {
|
||||
if !strings.Contains(targetAuthority, required) {
|
||||
t.Errorf("release target OSS policy authority is missing %q", required)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
t.Parallel()
|
||||
workflow := readReleaseWorkflow(t)
|
||||
@@ -807,7 +1077,17 @@ func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
if start == -1 {
|
||||
t.Fatal("release workflow is missing the channel repair job")
|
||||
}
|
||||
repair := workflow[start:]
|
||||
end := strings.Index(workflow[start:], "\n release-plan:\n")
|
||||
if end == -1 {
|
||||
t.Fatal("release workflow channel repair job is missing its end marker")
|
||||
}
|
||||
repair := workflow[start : start+end]
|
||||
authority := releaseWorkflowSection(
|
||||
t,
|
||||
repair,
|
||||
" - name: Verify immutable release authority\n",
|
||||
"\n - name: Require sealed OSS policy for repair\n",
|
||||
)
|
||||
tagAuthority := releaseWorkflowSection(
|
||||
t,
|
||||
repair,
|
||||
@@ -826,7 +1106,7 @@ func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
"oss_repair=0",
|
||||
`test -z "$REPAIR_GITEE_VERSION" || gitee_repair=1`,
|
||||
`test -z "$REPAIR_OSS_VERSION" || oss_repair=1`,
|
||||
"npm_repair + gitee_repair + oss_repair + governance + recovery",
|
||||
"release_flow + npm_repair + gitee_repair + oss_repair + governance + recovery",
|
||||
`echo "mode=repair_gitee" >> "$GITHUB_OUTPUT"`,
|
||||
`echo "mode=repair_oss" >> "$GITHUB_OUTPUT"`,
|
||||
} {
|
||||
@@ -847,6 +1127,8 @@ func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
"persist-credentials: false",
|
||||
"release_is_stable_version",
|
||||
"release_is_prerelease_version",
|
||||
`ref: ` + "`tags/withdrawn/${version}`",
|
||||
"was withdrawn and cannot be repaired",
|
||||
`ref: ` + "`tags/${version}`",
|
||||
`["ahead", "identical"].includes(comparison.data.status)`,
|
||||
"!release.data.immutable",
|
||||
@@ -854,6 +1136,8 @@ func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
"assetNames.length !== expectedAssets.size",
|
||||
"new Set(assetNames).size !== expectedAssets.size",
|
||||
`core.setOutput("tag_object", ref.data.object.sha)`,
|
||||
"Require sealed OSS policy for repair",
|
||||
"OSS repair is unavailable because this immutable release deferred the OSS channel.",
|
||||
`ref: ${{ steps.authority.outputs.commit_sha }}`,
|
||||
"path: release-source",
|
||||
"verify-github-tag-authority.sh",
|
||||
@@ -884,6 +1168,25 @@ func TestReleaseWorkflowChannelRepairUsesSealedReleaseAuthority(t *testing.T) {
|
||||
t.Errorf("channel repair authority is missing %q", required)
|
||||
}
|
||||
}
|
||||
for _, required := range []string{
|
||||
`const tagFields = new Map();`,
|
||||
`const ossMirror = tagFields.has("OSS-Mirror")`,
|
||||
`? tagFields.get("OSS-Mirror")`,
|
||||
`: "enabled";`,
|
||||
`!["enabled", "deferred"].includes(ossMirror)`,
|
||||
`core.setOutput("oss_mirror", ossMirror);`,
|
||||
} {
|
||||
if !strings.Contains(authority, required) {
|
||||
t.Errorf("channel repair tag policy authority is missing %q", required)
|
||||
}
|
||||
}
|
||||
npmRepair := releaseWorkflowSection(t, workflow, " repair-npm:\n", "\n release-delivery-gate:\n")
|
||||
if strings.Contains(npmRepair, `const ossMirror`) || strings.Contains(npmRepair, `core.setOutput("oss_mirror"`) {
|
||||
t.Error("npm repair must not parse or export the channel-only OSS policy")
|
||||
}
|
||||
if strings.Contains(repair, "ENABLE_OSS_MIRROR") {
|
||||
t.Error("OSS repair must use the immutable tag policy, not the current repository variable")
|
||||
}
|
||||
for _, asset := range []string{
|
||||
"dws-darwin-amd64.tar.gz",
|
||||
"dws-darwin-arm64.tar.gz",
|
||||
@@ -942,12 +1245,14 @@ func TestReleaseWorkflowRecoveryReusesGuardedJobs(t *testing.T) {
|
||||
"protected_branches !== true",
|
||||
"can_admins_bypass !== false",
|
||||
`run.path !== ".github/workflows/release.yml"`,
|
||||
`run.event !== "push"`,
|
||||
`const expectedEvent = failedByCloud ? "workflow_dispatch" : "push"`,
|
||||
`run.event !== expectedEvent`,
|
||||
`tagFields.get("Release-Run") !== failedRunId`,
|
||||
`"GET /repos/{owner}/{repo}/actions/runs/{run_id}/attempts/{attempt_number}"`,
|
||||
"attempt_number: Number(failedRunAttempt)",
|
||||
"run.run_attempt !== Number(failedRunAttempt)",
|
||||
`["failure", "cancelled", "timed_out", "startup_failure", "stale"].includes(run.conclusion)`,
|
||||
`run.head_branch !== version`,
|
||||
`run.head_branch !== expectedBranch`,
|
||||
`run.head_sha !== commit`,
|
||||
`tagObject !== expectedTagObject`,
|
||||
`["ahead", "identical"].includes(comparison.data.status)`,
|
||||
@@ -960,7 +1265,7 @@ func TestReleaseWorkflowRecoveryReusesGuardedJobs(t *testing.T) {
|
||||
`ref: sha`,
|
||||
`path: tmp/trusted-release-tooling`,
|
||||
`ref: ${{ github.sha }}`,
|
||||
`step.name === "Require immutable published GitHub Release"`,
|
||||
"verify-release-workflow-delivery.sh",
|
||||
"Require a clean sealed source before GoReleaser",
|
||||
`git status --porcelain --untracked-files=all`,
|
||||
} {
|
||||
@@ -1105,7 +1410,11 @@ func TestRecoverReleaseBindsOneFailedRunAttempt(t *testing.T) {
|
||||
"actions/runs/$find_run_id/attempts/$find_attempt",
|
||||
`select(.head_sha == \"$commit\" and .head_branch == \"$VERSION\")`,
|
||||
"Release run %s has no failed attempt",
|
||||
`[.id, .run_attempt, .repository.full_name, .path, .event, .status, .conclusion, .head_branch, .head_sha] | @tsv`,
|
||||
`[.id, .run_attempt, .repository.full_name, .path, .event, .status, .conclusion, .head_branch, .head_sha, .actor.login, .actor.id] | @tsv`,
|
||||
`Release-Run`,
|
||||
`Release-Run-Attempt`,
|
||||
`expected_attempt_event="workflow_dispatch"`,
|
||||
`is not bound by the cloud seal`,
|
||||
"actions/runs/%s/attempts/%s",
|
||||
`-f "recover_failed_run_attempt=$FAILED_RUN_ATTEMPT"`,
|
||||
} {
|
||||
@@ -1129,7 +1438,7 @@ func TestReleaseWorkflowPublicationBypassesSkippedDispatchButStopsOnCancellation
|
||||
name: "release contract",
|
||||
start: " release-contract:\n",
|
||||
end: "\n release:\n",
|
||||
condition: `if: ${{ !cancelled() && (github.event_name == 'push' || (needs.dispatch-contract.result == 'success' && needs.dispatch-contract.outputs.mode == 'recover_release' && needs.authorize-recovery.result == 'success')) }}`,
|
||||
condition: `if: ${{ !cancelled() && (github.event_name == 'push' || (needs.dispatch-contract.result == 'success' && needs.dispatch-contract.outputs.mode == 'recover_release' && needs.authorize-recovery.result == 'success') || (needs.dispatch-contract.result == 'success' && needs.dispatch-contract.outputs.mode == 'create_release' && needs.governance-preflight.result == 'success' && needs.release-plan.result == 'success' && needs.seal-release.result == 'success')) }}`,
|
||||
},
|
||||
{
|
||||
name: "build",
|
||||
@@ -1192,7 +1501,11 @@ func TestReleaseWorkflowDeliveryGateFailsClosed(t *testing.T) {
|
||||
"- mirror-gitee-release",
|
||||
"- repair-npm",
|
||||
"- repair-channel",
|
||||
"- release-plan",
|
||||
"- seal-release",
|
||||
`REPAIR_CHANNEL_RESULT: ${{ needs.repair-channel.result }}`,
|
||||
`RELEASE_PLAN_RESULT: ${{ needs.release-plan.result }}`,
|
||||
`SEAL_RELEASE_RESULT: ${{ needs.seal-release.result }}`,
|
||||
"require_publication",
|
||||
`require_result release-contract "$RELEASE_CONTRACT_RESULT" success`,
|
||||
`require_result release "$RELEASE_RESULT" success`,
|
||||
@@ -1200,6 +1513,8 @@ func TestReleaseWorkflowDeliveryGateFailsClosed(t *testing.T) {
|
||||
`require_result publish-release "$PUBLISH_RELEASE_RESULT" success`,
|
||||
`require_result publish-channels "$PUBLISH_CHANNELS_RESULT" success`,
|
||||
"workflow_dispatch:recover_release",
|
||||
"workflow_dispatch:create_release",
|
||||
"workflow_dispatch:plan_release",
|
||||
"workflow_dispatch:governance_preflight",
|
||||
"workflow_dispatch:repair_npm",
|
||||
"workflow_dispatch:repair_gitee",
|
||||
@@ -1529,6 +1844,133 @@ func TestReleaseWorkflowOpensVersionedHomebrewPRForBetaTags(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowWaitsForNPMDistTagPropagation(t *testing.T) {
|
||||
workflow := readReleaseWorkflow(t)
|
||||
script := releaseWorkflowRunScript(
|
||||
t,
|
||||
workflow,
|
||||
"Verify npm channel delivery",
|
||||
"Sync release artifacts to China OSS mirror",
|
||||
)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
sequence string
|
||||
wantSuccess bool
|
||||
wantCalls int
|
||||
wantSleeps int
|
||||
wantOutput string
|
||||
}{
|
||||
{
|
||||
name: "stale beta converges to target",
|
||||
sequence: "1.0.53-beta.6\n1.0.53-beta.6\n1.0.53-beta.7\n",
|
||||
wantSuccess: true,
|
||||
wantCalls: 3,
|
||||
wantSleeps: 2,
|
||||
},
|
||||
{
|
||||
name: "stale beta never converges",
|
||||
sequence: "1.0.53-beta.6\n",
|
||||
wantCalls: 12,
|
||||
wantSleeps: 11,
|
||||
wantOutput: "still reports older v1.0.53-beta.6 after 12 attempts",
|
||||
},
|
||||
{
|
||||
name: "permanent registry error fails immediately",
|
||||
sequence: "__NPM_ERROR__\n",
|
||||
wantCalls: 1,
|
||||
wantSleeps: 0,
|
||||
wantOutput: "permanent npm registry error",
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
root := t.TempDir()
|
||||
fakeBin := filepath.Join(root, "bin")
|
||||
if err := os.MkdirAll(fakeBin, 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll(%s) error = %v", fakeBin, err)
|
||||
}
|
||||
sequencePath := filepath.Join(root, "sequence")
|
||||
statePath := filepath.Join(root, "state")
|
||||
npmLogPath := filepath.Join(root, "npm.log")
|
||||
sleepLogPath := filepath.Join(root, "sleep.log")
|
||||
mustWriteFile(t, sequencePath, []byte(test.sequence), 0o644)
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "npm"), []byte(`#!/bin/sh
|
||||
set -eu
|
||||
printf '%s\n' "$*" >> "$NPM_CALL_LOG"
|
||||
test "$*" = "view dingtalk-workspace-cli dist-tags.beta --registry=https://registry.npmjs.org --prefer-online" || {
|
||||
echo "unexpected npm mutation: $*" >&2
|
||||
exit 97
|
||||
}
|
||||
call=0
|
||||
if test -f "$NPM_STATE"; then call="$(cat "$NPM_STATE")"; fi
|
||||
call=$((call + 1))
|
||||
printf '%s\n' "$call" > "$NPM_STATE"
|
||||
value="$(sed -n "${call}p" "$NPM_SEQUENCE")"
|
||||
if test -z "$value"; then value="$(tail -n 1 "$NPM_SEQUENCE")"; fi
|
||||
if test "$value" = "__NPM_ERROR__"; then
|
||||
echo "permanent npm registry error" >&2
|
||||
exit 42
|
||||
fi
|
||||
printf '%s\n' "$value"
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "sleep"), []byte(`#!/bin/sh
|
||||
set -eu
|
||||
printf '%s\n' "$*" >> "$SLEEP_CALL_LOG"
|
||||
`), 0o755)
|
||||
|
||||
repoRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(repository root) error = %v", err)
|
||||
}
|
||||
cmd := exec.Command("sh", "-c", script)
|
||||
cmd.Dir = repoRoot
|
||||
cmd.Env = append(os.Environ(),
|
||||
"PATH="+fakeBin+string(os.PathListSeparator)+os.Getenv("PATH"),
|
||||
"NPM_TAG=beta",
|
||||
"SEMVER=1.0.53-beta.7",
|
||||
"NPM_SEQUENCE="+sequencePath,
|
||||
"NPM_STATE="+statePath,
|
||||
"NPM_CALL_LOG="+npmLogPath,
|
||||
"SLEEP_CALL_LOG="+sleepLogPath,
|
||||
)
|
||||
output, runErr := cmd.CombinedOutput()
|
||||
if test.wantSuccess && runErr != nil {
|
||||
t.Fatalf("npm delivery verification error = %v\noutput:\n%s", runErr, output)
|
||||
}
|
||||
if !test.wantSuccess && runErr == nil {
|
||||
t.Fatalf("npm delivery verification unexpectedly succeeded\noutput:\n%s", output)
|
||||
}
|
||||
if test.wantOutput != "" && !strings.Contains(string(output), test.wantOutput) {
|
||||
t.Errorf("npm delivery verification output is missing %q:\n%s", test.wantOutput, output)
|
||||
}
|
||||
|
||||
npmLog, err := os.ReadFile(npmLogPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", npmLogPath, err)
|
||||
}
|
||||
npmCalls := strings.Split(strings.TrimSpace(string(npmLog)), "\n")
|
||||
if got := len(npmCalls); got != test.wantCalls {
|
||||
t.Errorf("npm view call count = %d, want %d; log:\n%s", got, test.wantCalls, npmLog)
|
||||
}
|
||||
if strings.Contains(string(npmLog), "dist-tag add") || strings.Contains(string(npmLog), "publish") {
|
||||
t.Errorf("delivery verification must remain read-only; log:\n%s", npmLog)
|
||||
}
|
||||
|
||||
sleepCalls := 0
|
||||
if sleepLog, err := os.ReadFile(sleepLogPath); err == nil {
|
||||
sleepCalls = len(strings.Fields(string(sleepLog)))
|
||||
} else if !os.IsNotExist(err) {
|
||||
t.Fatalf("ReadFile(%s) error = %v", sleepLogPath, err)
|
||||
}
|
||||
if sleepCalls != test.wantSleeps {
|
||||
t.Errorf("sleep call count = %d, want %d", sleepCalls, test.wantSleeps)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseStaysDraftUntilFinalizedAssetDigestsMatch(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
|
||||
@@ -141,6 +141,136 @@ func TestReleaseVersionOrdering(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestNextReleaseVersion(t *testing.T) {
|
||||
sourceRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(repo root) error = %v", err)
|
||||
}
|
||||
script := filepath.Join(sourceRoot, "scripts", "release", "next-release-version.sh")
|
||||
|
||||
newRepo := func(t *testing.T, tags ...string) string {
|
||||
t.Helper()
|
||||
root := t.TempDir()
|
||||
mustRun(t, root, "git", "init", "-b", "main")
|
||||
mustRun(t, root, "git", "config", "user.name", "Release Test")
|
||||
mustRun(t, root, "git", "config", "user.email", "release-test@example.com")
|
||||
mustWriteFile(t, filepath.Join(root, "seed.txt"), []byte("release allocator\n"), 0o644)
|
||||
mustRun(t, root, "git", "add", "seed.txt")
|
||||
mustRun(t, root, "git", "commit", "-m", "seed")
|
||||
for _, tag := range tags {
|
||||
mustRun(t, root, "git", "tag", "-a", tag, "-m", "Allocate "+tag)
|
||||
}
|
||||
return root
|
||||
}
|
||||
run := func(t *testing.T, root string, args ...string) (string, error) {
|
||||
t.Helper()
|
||||
commandArgs := append([]string{script, "--repo-root", root}, args...)
|
||||
output, err := exec.Command("sh", commandArgs...).CombinedOutput()
|
||||
return string(output), err
|
||||
}
|
||||
|
||||
t.Run("continues highest open beta core", func(t *testing.T) {
|
||||
root := newRepo(t,
|
||||
"v1.0.52",
|
||||
"v1.0.53-beta.1",
|
||||
"v1.0.53-beta.2",
|
||||
"v1.0.53-beta.3",
|
||||
"v1.0.53-beta.4",
|
||||
)
|
||||
output, err := run(t, root, "--channel", "beta")
|
||||
if err != nil {
|
||||
t.Fatalf("next beta error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
want := "release_version=v1.0.53-beta.5\nfrom_beta=\nchannel=prerelease\nbase=v1.0.52\n"
|
||||
if output != want {
|
||||
t.Fatalf("next beta output:\ngot:\n%s\nwant:\n%s", output, want)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("withdrawn beta remains allocated", func(t *testing.T) {
|
||||
root := newRepo(t,
|
||||
"v1.0.52",
|
||||
"v1.0.53-beta.4",
|
||||
"withdrawn/v1.0.53-beta.5",
|
||||
)
|
||||
output, err := run(t, root, "--channel", "prerelease")
|
||||
if err != nil {
|
||||
t.Fatalf("next beta after withdrawal error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
if !strings.Contains(output, "release_version=v1.0.53-beta.6\n") {
|
||||
t.Fatalf("withdrawn beta version was reused:\n%s", output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("starts requested bumped beta core", func(t *testing.T) {
|
||||
tests := []struct {
|
||||
bump string
|
||||
want string
|
||||
}{
|
||||
{bump: "patch", want: "v1.2.4-beta.1"},
|
||||
{bump: "minor", want: "v1.3.0-beta.1"},
|
||||
{bump: "major", want: "v2.0.0-beta.1"},
|
||||
}
|
||||
for _, test := range tests {
|
||||
t.Run(test.bump, func(t *testing.T) {
|
||||
root := newRepo(t, "v1.2.3")
|
||||
output, err := run(t, root, "--channel", "beta", "--bump", test.bump)
|
||||
if err != nil {
|
||||
t.Fatalf("new beta line error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
if !strings.Contains(output, "release_version="+test.want+"\n") {
|
||||
t.Fatalf("new beta line output:\n%s", output)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("promotes latest beta from highest open core", func(t *testing.T) {
|
||||
root := newRepo(t,
|
||||
"v1.0.0",
|
||||
"v1.1.0-beta.3",
|
||||
"v1.2.0-beta.1",
|
||||
"v1.2.0-beta.2",
|
||||
)
|
||||
output, err := run(t, root, "--channel", "stable")
|
||||
if err != nil {
|
||||
t.Fatalf("stable allocation error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
want := "release_version=v1.2.0\nfrom_beta=v1.2.0-beta.2\nchannel=stable\nbase=v1.0.0\n"
|
||||
if output != want {
|
||||
t.Fatalf("stable allocation output:\ngot:\n%s\nwant:\n%s", output, want)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("withdrawn stable closes its core", func(t *testing.T) {
|
||||
root := newRepo(t,
|
||||
"v1.0.0",
|
||||
"v1.0.1-beta.1",
|
||||
"withdrawn/v1.0.1",
|
||||
)
|
||||
output, err := run(t, root, "--channel", "beta")
|
||||
if err != nil {
|
||||
t.Fatalf("beta after withdrawn stable error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
want := "release_version=v1.0.2-beta.1\nfrom_beta=\nchannel=prerelease\nbase=v1.0.1\n"
|
||||
if output != want {
|
||||
t.Fatalf("withdrawn stable allocation output:\ngot:\n%s\nwant:\n%s", output, want)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("withdrawn latest beta cannot be promoted", func(t *testing.T) {
|
||||
root := newRepo(t,
|
||||
"v1.0.0",
|
||||
"v1.0.1-beta.1",
|
||||
"withdrawn/v1.0.1-beta.2",
|
||||
)
|
||||
output, err := run(t, root, "--channel", "stable")
|
||||
if err == nil || !strings.Contains(output, "latest beta v1.0.1-beta.2 is withdrawn") {
|
||||
t.Fatalf("withdrawn beta promotion was not rejected: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestReleaseLibHelpersPreserveCallerVariables(t *testing.T) {
|
||||
sourceRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
@@ -674,6 +804,179 @@ esac
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowDeliveryAcceptsOnlyTagBoundCloudRelease(t *testing.T) {
|
||||
sourceRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(repo root) error = %v", err)
|
||||
}
|
||||
binDir := t.TempDir()
|
||||
fakeCurl := filepath.Join(binDir, "curl")
|
||||
mustWriteFile(t, fakeCurl, []byte(`#!/bin/sh
|
||||
set -eu
|
||||
for argument in "$@"; do endpoint="$argument"; done
|
||||
case "$endpoint" in
|
||||
*event=push*)
|
||||
printf '{"workflow_runs":[]}\n'
|
||||
;;
|
||||
*/git/ref/tags/*)
|
||||
printf '{"object":{"type":"tag","sha":"%s"}}\n' "$TAG_OBJECT"
|
||||
;;
|
||||
*/git/tags/*)
|
||||
python3 - <<'PY'
|
||||
import json
|
||||
import os
|
||||
|
||||
message = "\n".join([
|
||||
f"Release {os.environ['TAG']}",
|
||||
"",
|
||||
"Channel: prerelease",
|
||||
"Release-Run: 42",
|
||||
"Release-Run-Attempt: 1",
|
||||
"Requested-By: release-user",
|
||||
"Requested-By-ID: 1234",
|
||||
f"Sealed-Commit: {os.environ['RELEASE_COMMIT']}",
|
||||
f"Workflow-Commit: {os.environ['RELEASE_COMMIT']}",
|
||||
f"Allocation-Fingerprint: {'d' * 64}",
|
||||
])
|
||||
print(json.dumps({
|
||||
"tag": os.environ["TAG"],
|
||||
"message": message,
|
||||
"object": {"type": "commit", "sha": os.environ["RELEASE_COMMIT"]},
|
||||
}))
|
||||
PY
|
||||
;;
|
||||
*/actions/runs/42/attempts/1/jobs*)
|
||||
python3 - <<'PY'
|
||||
import json
|
||||
import os
|
||||
|
||||
required = [
|
||||
"Plan next cloud release",
|
||||
"Seal cloud release tag",
|
||||
"release-contract",
|
||||
"Build signed release artifacts",
|
||||
"Verify Apple Developer ID signatures",
|
||||
"Publish immutable GitHub Release",
|
||||
"Publish npm and mirrors",
|
||||
"Release delivery gate",
|
||||
]
|
||||
if os.environ.get("MISSING_SEAL") == "1":
|
||||
required.remove("Seal cloud release tag")
|
||||
jobs = []
|
||||
for name in required:
|
||||
job = {
|
||||
"name": name,
|
||||
"status": "completed",
|
||||
"conclusion": "success",
|
||||
"head_sha": os.environ["RELEASE_COMMIT"],
|
||||
"steps": [],
|
||||
}
|
||||
if name == "Seal cloud release tag":
|
||||
job["steps"] = [{
|
||||
"name": "Create one immutable annotated release tag",
|
||||
"status": "completed",
|
||||
"conclusion": "success",
|
||||
}]
|
||||
elif name == "Publish immutable GitHub Release":
|
||||
job["steps"] = [{
|
||||
"name": "Require immutable published GitHub Release",
|
||||
"status": "completed",
|
||||
"conclusion": "success",
|
||||
}]
|
||||
if os.environ.get("PUBLISH_RELEASE_FAILURE") == "1":
|
||||
job["conclusion"] = "failure"
|
||||
elif name == "Publish npm and mirrors":
|
||||
job["conclusion"] = os.environ.get("CHANNEL_CONCLUSION", "success")
|
||||
jobs.append(job)
|
||||
print(json.dumps({"jobs": jobs}))
|
||||
PY
|
||||
;;
|
||||
*/actions/runs/42/attempts/1)
|
||||
python3 - <<'PY'
|
||||
import json
|
||||
import os
|
||||
print(json.dumps({
|
||||
"id": 42,
|
||||
"run_attempt": 1,
|
||||
"repository": {"full_name": "owner/repo"},
|
||||
"path": ".github/workflows/release.yml",
|
||||
"event": "workflow_dispatch",
|
||||
"status": "completed",
|
||||
"conclusion": os.environ.get("RUN_CONCLUSION", "success"),
|
||||
"head_branch": "main",
|
||||
"head_sha": os.environ["RELEASE_COMMIT"],
|
||||
"actor": {
|
||||
"login": os.environ.get("RUN_ACTOR", "release-user"),
|
||||
"id": int(os.environ.get("RUN_ACTOR_ID", "1234")),
|
||||
},
|
||||
}))
|
||||
PY
|
||||
;;
|
||||
*event=workflow_dispatch*)
|
||||
printf '{"workflow_runs":[]}\n'
|
||||
;;
|
||||
*) exit 1 ;;
|
||||
esac
|
||||
`), 0o755)
|
||||
script := filepath.Join(sourceRoot, "scripts", "release", "verify-release-workflow-delivery.sh")
|
||||
tag := "v1.2.3-beta.1"
|
||||
commit := strings.Repeat("a", 40)
|
||||
tagObject := strings.Repeat("b", 40)
|
||||
runWithArgs := func(arguments []string, overrides ...string) (string, error) {
|
||||
cmd := exec.Command("sh", script, tag, commit)
|
||||
cmd.Args = append([]string{"sh", script}, arguments...)
|
||||
cmd.Env = append([]string{
|
||||
"PATH=" + binDir + string(os.PathListSeparator) + os.Getenv("PATH"),
|
||||
"HOME=" + t.TempDir(),
|
||||
"DWS_RELEASE_OFFICIAL_REPOSITORY=owner/repo",
|
||||
"TAG=" + tag,
|
||||
"TAG_OBJECT=" + tagObject,
|
||||
"RELEASE_COMMIT=" + commit,
|
||||
}, overrides...)
|
||||
output, err := cmd.CombinedOutput()
|
||||
return string(output), err
|
||||
}
|
||||
run := func(overrides ...string) (string, error) {
|
||||
return runWithArgs([]string{tag, commit}, overrides...)
|
||||
}
|
||||
|
||||
if output, err := run(); err != nil || !strings.Contains(output, "cloud release run 42") {
|
||||
t.Fatalf("tag-bound cloud delivery was rejected: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := runWithArgs(
|
||||
[]string{"--channel-repair", "oss", tag, commit},
|
||||
); err != nil || !strings.Contains(output, "cloud release run 42") {
|
||||
t.Fatalf("OSS repair could not use a successful release with the mirror enabled: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := run("RUN_ACTOR=renamed-release-user"); err != nil ||
|
||||
!strings.Contains(output, "cloud release run 42") {
|
||||
t.Fatalf("cloud delivery broke after a harmless login rename: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := run("RUN_ACTOR_ID=9999"); err == nil ||
|
||||
!strings.Contains(output, "did not deliver") {
|
||||
t.Fatalf("cloud delivery with the wrong stable actor ID passed: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := run("MISSING_SEAL=1"); err == nil ||
|
||||
!strings.Contains(output, "did not deliver") {
|
||||
t.Fatalf("cloud delivery without the seal job passed: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := runWithArgs(
|
||||
[]string{"--npm-repair", tag, commit},
|
||||
"RUN_CONCLUSION=failure",
|
||||
"CHANNEL_CONCLUSION=failure",
|
||||
); err != nil || !strings.Contains(output, "npm-repair authority verified") {
|
||||
t.Fatalf("npm repair could not use a sealed immutable release after npm failure: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
if output, err := runWithArgs(
|
||||
[]string{"--npm-repair", tag, commit},
|
||||
"RUN_CONCLUSION=failure",
|
||||
"CHANNEL_CONCLUSION=failure",
|
||||
"PUBLISH_RELEASE_FAILURE=1",
|
||||
); err == nil || strings.Contains(output, "npm-repair authority verified") {
|
||||
t.Fatalf("npm repair accepted a failed immutable GitHub publication: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseWorkflowDeliveryChannelRepairRequiresLatestAttemptCoreDelivery(t *testing.T) {
|
||||
sourceRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
@@ -1031,6 +1334,109 @@ func TestReleaseContractRejectsInvalidVersionChannelPairs(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseContractTreatsWithdrawnVersionsAsPermanentlyAllocated(t *testing.T) {
|
||||
t.Run("continues after withdrawn beta", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
section := "## [1.0.1-beta.3] - 2026-07-11\n\n### Changed\n\n- Replace the withdrawn beta.\n\n"
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(section)), 0o644)
|
||||
r.commitAndPush(t, "prepare replacement beta")
|
||||
mustRun(t, r.root, "git", "tag", "-a", "v1.0.1-beta.1", "-m", "Release v1.0.1-beta.1")
|
||||
mustRun(t, r.root, "git", "tag", "-a", "withdrawn/v1.0.1-beta.2", "-m", "Withdraw v1.0.1-beta.2")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "prerelease",
|
||||
"--version", "v1.0.1-beta.3",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("replacement beta was rejected: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("never reuses exact withdrawn version", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
section := "## [1.0.1-beta.1] - 2026-07-11\n\n### Changed\n\n- Must not reuse this version.\n\n"
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(section)), 0o644)
|
||||
r.commitAndPush(t, "prepare withdrawn version")
|
||||
mustRun(t, r.root, "git", "tag", "-a", "withdrawn/v1.0.1-beta.1", "-m", "Withdraw v1.0.1-beta.1")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "prerelease",
|
||||
"--version", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil || !strings.Contains(output, "can never be reused") {
|
||||
t.Fatalf("withdrawn version was reusable: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("withdrawn beta cannot be promoted", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
r.seedBeta(t)
|
||||
mustRun(t, r.root, "git", "tag", "-a", "withdrawn/v1.0.1-beta.1", "-m", "Withdraw v1.0.1-beta.1")
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(stableSection(), betaSection())), 0o644)
|
||||
r.commitAndPush(t, "prepare blocked stable")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "stable",
|
||||
"--version", "v1.0.1",
|
||||
"--from-beta", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil || !strings.Contains(output, "beta baseline was withdrawn") {
|
||||
t.Fatalf("withdrawn beta was promoted: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("withdrawn stable is a permanent version floor", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
mustRun(t, r.root, "git", "tag", "-a", "withdrawn/v2.0.0", "-m", "Withdraw v2.0.0")
|
||||
section := "## [1.1.0-beta.1] - 2026-07-11\n\n### Changed\n\n- Must remain above every allocated stable.\n\n"
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(section)), 0o644)
|
||||
r.commitAndPush(t, "prepare invalid lower beta")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "prerelease",
|
||||
"--version", "v1.1.0-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil || !strings.Contains(output, "must be greater than latest allocated stable v2.0.0") {
|
||||
t.Fatalf("release below withdrawn stable floor passed: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("tombstone excludes stale ordinary tag from delivered baseline", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
mustRun(t, r.root, "git", "tag", "-a", "v1.0.1", "-m", "Release v1.0.1")
|
||||
mustRun(t, r.root, "git", "tag", "-a", "withdrawn/v1.0.1", "-m", "Withdraw v1.0.1")
|
||||
section := "## [1.0.2-beta.1] - 2026-07-11\n\n### Changed\n\n- Continue after the withdrawn stable.\n\n"
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(section)), 0o644)
|
||||
r.commitAndPush(t, "prepare beta after withdrawn stable")
|
||||
metadata := filepath.Join(t.TempDir(), "metadata")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "prerelease",
|
||||
"--version", "v1.0.2-beta.1",
|
||||
"--remote", "origin",
|
||||
"--metadata-output", metadata,
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("release after stale withdrawn tag failed: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
assertFileContains(t, metadata, "previous_stable=v1.0.0")
|
||||
if content, err := os.ReadFile(metadata); err != nil {
|
||||
t.Fatalf("ReadFile(metadata) error = %v", err)
|
||||
} else if strings.Contains(string(content), "previous_stable=v1.0.1") {
|
||||
t.Fatalf("withdrawn stale tag remained the delivered baseline:\n%s", content)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestReleaseContractRejectsBadChangelogSections(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -1088,12 +1494,12 @@ func TestReleaseContractRejectsDirtyOrUnsyncedMain(t *testing.T) {
|
||||
"--version", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil || !strings.Contains(output, "must exactly match origin/main") {
|
||||
if err == nil || !strings.Contains(output, "must be contained in origin/main history") {
|
||||
t.Fatalf("unsynced main was not blocked: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseContractStablePromotionAllowsOnlyChangelogDiff(t *testing.T) {
|
||||
func TestReleaseContractStablePromotionRequiresBetaAncestry(t *testing.T) {
|
||||
t.Run("sealed", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
r.seedBeta(t)
|
||||
@@ -1112,12 +1518,12 @@ func TestReleaseContractStablePromotionAllowsOnlyChangelogDiff(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("source drift", func(t *testing.T) {
|
||||
t.Run("commits after beta", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
r.seedBeta(t)
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(stableSection(), betaSection())), 0o644)
|
||||
mustWriteFile(t, filepath.Join(r.root, "drift.txt"), []byte("untested change\n"), 0o644)
|
||||
r.commitAndPush(t, "drift after beta")
|
||||
mustWriteFile(t, filepath.Join(r.root, "followup.txt"), []byte("merged after beta\n"), 0o644)
|
||||
r.commitAndPush(t, "merge follow-up after beta")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
@@ -1126,11 +1532,53 @@ func TestReleaseContractStablePromotionAllowsOnlyChangelogDiff(t *testing.T) {
|
||||
"--from-beta", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil {
|
||||
t.Fatalf("drifted stable promotion unexpectedly passed:\n%s", output)
|
||||
if err != nil {
|
||||
t.Fatalf("stable promotion with commits after beta error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
if !strings.Contains(output, "only CHANGELOG.md may differ") || !strings.Contains(output, "drift.txt") {
|
||||
t.Fatalf("drift output is not actionable:\n%s", output)
|
||||
})
|
||||
|
||||
t.Run("beta outside HEAD history", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
mustRun(t, r.root, "git", "checkout", "-b", "sidecar")
|
||||
mustWriteFile(t, filepath.Join(r.root, "sidecar.txt"), []byte("never merged\n"), 0o644)
|
||||
mustRun(t, r.root, "git", "add", ".")
|
||||
mustRun(t, r.root, "git", "commit", "-m", "sidecar beta candidate")
|
||||
mustRun(t, r.root, "git", "tag", "-a", "v1.0.1-beta.1", "-m", "Release v1.0.1-beta.1", "-m", "Channel: prerelease")
|
||||
mustRun(t, r.root, "git", "checkout", "main")
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(stableSection(), betaSection())), 0o644)
|
||||
r.commitAndPush(t, "prepare stable changelog")
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "stable",
|
||||
"--version", "v1.0.1",
|
||||
"--from-beta", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err == nil || !strings.Contains(output, "not an ancestor of HEAD") {
|
||||
t.Fatalf("beta outside HEAD history was promoted: err=%v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("older sealed commit after main advanced", func(t *testing.T) {
|
||||
r := newReleaseTestRepo(t)
|
||||
r.seedBeta(t)
|
||||
mustWriteFile(t, filepath.Join(r.root, "CHANGELOG.md"), []byte(releaseChangelog(stableSection(), betaSection())), 0o644)
|
||||
r.commitAndPush(t, "prepare stable changelog")
|
||||
sealed := strings.TrimSpace(mustOutput(t, r.root, "git", "rev-parse", "HEAD"))
|
||||
mustWriteFile(t, filepath.Join(r.root, "after.txt"), []byte("main advanced\n"), 0o644)
|
||||
r.commitAndPush(t, "advance main after stable candidate")
|
||||
mustRun(t, r.root, "git", "checkout", "--detach", sealed)
|
||||
|
||||
output, err := runReleaseScript(t, r.root, r.contract,
|
||||
"--repo-root", r.root,
|
||||
"--channel", "stable",
|
||||
"--version", "v1.0.1",
|
||||
"--from-beta", "v1.0.1-beta.1",
|
||||
"--remote", "origin",
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("older sealed commit in main history was rejected: %v\noutput:\n%s", err, output)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,943 @@
|
||||
package scripts_test
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestWithdrawReleaseWorkflowIsProtectedAndFailClosed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
workflowPath, err := filepath.Abs(filepath.Join("..", "..", ".github", "workflows", "withdraw-release.yml"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(withdraw workflow) error = %v", err)
|
||||
}
|
||||
content, err := os.ReadFile(workflowPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", workflowPath, err)
|
||||
}
|
||||
workflow := string(content)
|
||||
|
||||
for _, want := range []string{
|
||||
"workflow_dispatch:",
|
||||
"version:",
|
||||
"reason:",
|
||||
"confirmation:",
|
||||
"group: dws-release-publication",
|
||||
"cancel-in-progress: false",
|
||||
"environment: release-withdrawal",
|
||||
"prevent_self_review !== true",
|
||||
"deployment_branch_policy?.protected_branches !== true",
|
||||
"can_admins_bypass !== false",
|
||||
`const expectedRepository = "DingTalk-Real-AI/dingtalk-workspace-cli"`,
|
||||
"context.ref !== `refs/heads/${defaultBranch}`",
|
||||
"branch.data.object.sha !== context.sha",
|
||||
"contents: write",
|
||||
"persist-credentials: false",
|
||||
"NODE_AUTH_TOKEN: ${{ secrets.NPM_TOKEN }}",
|
||||
"DWS_GITEE_ENABLED:",
|
||||
"HOMEBREW_PR_TOKEN: ${{ secrets.HOMEBREW_PR_TOKEN }}",
|
||||
"./scripts/release/withdraw-release.sh",
|
||||
"already-installed clients cannot be remotely downgraded.",
|
||||
"this run remains failed until that PR is independently reviewed",
|
||||
} {
|
||||
if !strings.Contains(workflow, want) {
|
||||
t.Errorf("withdraw workflow missing %q", want)
|
||||
}
|
||||
}
|
||||
for _, forbidden := range []string{
|
||||
"\n push:",
|
||||
"\n schedule:",
|
||||
"npm unpublish",
|
||||
"cancel-in-progress: true",
|
||||
} {
|
||||
if strings.Contains(workflow, forbidden) {
|
||||
t.Errorf("withdraw workflow contains forbidden trigger/action %q", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithdrawReleaseScriptDeletesProblemReleaseLastAndUsesPermanentTombstone(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
scriptPath, err := filepath.Abs(filepath.Join("..", "..", "scripts", "release", "withdraw-release.sh"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(withdraw script) error = %v", err)
|
||||
}
|
||||
content, err := os.ReadFile(scriptPath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", scriptPath, err)
|
||||
}
|
||||
script := string(content)
|
||||
|
||||
for _, want := range []string{
|
||||
`TOMBSTONE="withdrawn/${VERSION}"`,
|
||||
`github_api --method POST "repos/${OFFICIAL_REPOSITORY}/git/tags"`,
|
||||
`github_api --method POST "repos/${OFFICIAL_REPOSITORY}/git/refs"`,
|
||||
"existing tombstone $TOMBSTONE has different immutable withdrawal metadata",
|
||||
`github_api --method PATCH`,
|
||||
`npm deprecate "${PACKAGE_NAME}@${VERSION#v}"`,
|
||||
`npm dist-tag add "${PACKAGE_NAME}@${ROLLBACK_VERSION#v}"`,
|
||||
`TARGET_OSS_MODE="$("$SCRIPT_DIR/release-tag-oss-mode.sh" "$VERSION")"`,
|
||||
`printf 'OSS-Mirror: %s\n' "$TARGET_OSS_MODE"`,
|
||||
`oss_mode = fields.get("OSS-Mirror", "enabled")`,
|
||||
`if [ "$OSS_ENABLED" = true ]; then`,
|
||||
`*) err "could not resolve immutable OSS policy for $VERSION" ;;`,
|
||||
`"$OSSUTIL" rm -rf`,
|
||||
`curl -fsS -X DELETE`,
|
||||
`git push "$GITEE_GIT_REMOTE" ":refs/tags/${VERSION}"`,
|
||||
`)" || return 1`,
|
||||
`err "could not verify Gitee tag deletion for $VERSION"`,
|
||||
`disabled_gitee_release_id="$(gitee_release_id "$VERSION")"`,
|
||||
`DWS_TAP_PR_TITLE="revert: withdraw ${VERSION} and restore ${ROLLBACK_VERSION}"`,
|
||||
"Homebrew rollback PR requires independent review and merge",
|
||||
`"repos/${OFFICIAL_REPOSITORY}/releases/${TARGET_RELEASE_ID}"`,
|
||||
`"repos/${OFFICIAL_REPOSITORY}/git/refs/tags/${VERSION}"`,
|
||||
`github_expect_404 "repos/${OFFICIAL_REPOSITORY}/releases/tags/${VERSION}"`,
|
||||
`github_expect_404 "repos/${OFFICIAL_REPOSITORY}/git/ref/tags/${VERSION}"`,
|
||||
"Already-installed clients cannot be remotely downgraded",
|
||||
} {
|
||||
if !strings.Contains(script, want) {
|
||||
t.Errorf("withdraw script missing %q", want)
|
||||
}
|
||||
}
|
||||
for _, forbidden := range []string{
|
||||
"npm unpublish",
|
||||
`"repos/${OFFICIAL_REPOSITORY}/git/refs/tags/${TOMBSTONE}"`,
|
||||
`--force "refs/tags/${TOMBSTONE}`,
|
||||
`make_latest`,
|
||||
} {
|
||||
if strings.Contains(script, forbidden) {
|
||||
t.Errorf("withdraw script contains evidence-destroying operation %q", forbidden)
|
||||
}
|
||||
}
|
||||
|
||||
tombstone := strings.LastIndex(script, "\ncreate_tombstone\n")
|
||||
githubMutation := strings.LastIndex(script, "\nupdate_github_release\n")
|
||||
npmMutation := strings.LastIndex(script, "\nupdate_npm_channel\n")
|
||||
ossMutation := strings.LastIndex(script, `update_oss_channel "$OSS_POINTER_NAME"`)
|
||||
homebrewGate := strings.LastIndex(script, "\nupdate_homebrew\n")
|
||||
githubDelete := strings.LastIndex(script, "\ndelete_github_release_and_tag\n")
|
||||
if tombstone < 0 || githubMutation < 0 || npmMutation < 0 || ossMutation < 0 ||
|
||||
homebrewGate < 0 || githubDelete < 0 {
|
||||
t.Fatalf("could not locate withdrawal mutation sequence")
|
||||
}
|
||||
if !(tombstone < homebrewGate && homebrewGate < githubMutation &&
|
||||
githubMutation < npmMutation && npmMutation < ossMutation && ossMutation < githubDelete) {
|
||||
t.Fatalf("unsafe withdrawal order: tombstone=%d github-mark=%d npm=%d oss=%d homebrew=%d github-delete=%d",
|
||||
tombstone, githubMutation, npmMutation, ossMutation, homebrewGate, githubDelete)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReleaseTagOSSModeIsImmutableAndFailClosed(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
scriptPath, err := filepath.Abs(filepath.Join("..", "..", "scripts", "release", "release-tag-oss-mode.sh"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(OSS mode script) error = %v", err)
|
||||
}
|
||||
repo := t.TempDir()
|
||||
mustRun(t, repo, "git", "init", "-b", "main")
|
||||
mustRun(t, repo, "git", "config", "user.name", "OSS Mode Test")
|
||||
mustRun(t, repo, "git", "config", "user.email", "oss-mode@example.com")
|
||||
mustWriteFile(t, filepath.Join(repo, "tracked"), []byte("fixture\n"), 0o644)
|
||||
mustRun(t, repo, "git", "add", "tracked")
|
||||
mustRun(t, repo, "git", "commit", "-m", "fixture")
|
||||
|
||||
tag := func(version, message string) {
|
||||
t.Helper()
|
||||
messagePath := filepath.Join(repo, version+".message")
|
||||
mustWriteFile(t, messagePath, []byte(message), 0o644)
|
||||
mustRun(t, repo, "git", "tag", "-a", version, "-F", messagePath)
|
||||
}
|
||||
tag("v1.0.1", "Release v1.0.1\n")
|
||||
tag("v1.0.2", "Release v1.0.2\n\nOSS-Mirror: enabled\n")
|
||||
tag("v1.0.3", "Release v1.0.3\n\nOSS-Mirror: deferred\n")
|
||||
tag("v1.0.4", "Release v1.0.4\n\nOSS-Mirror: invalid\n")
|
||||
tag("v1.0.5", "Release v1.0.5\n\nOSS-Mirror: enabled\nOSS-Mirror: deferred\n")
|
||||
mustRun(t, repo, "git", "tag", "v1.0.6")
|
||||
rawTag := func(version, message string) {
|
||||
t.Helper()
|
||||
commit := strings.TrimSpace(mustOutput(t, repo, "git", "rev-parse", "HEAD"))
|
||||
payload := strings.Join([]string{
|
||||
"object " + commit,
|
||||
"type commit",
|
||||
"tag " + version,
|
||||
"tagger OSS Mode Test <oss-mode@example.com> 1700000000 +0000",
|
||||
"",
|
||||
message,
|
||||
}, "\n")
|
||||
cmd := exec.Command("git", "mktag")
|
||||
cmd.Dir = repo
|
||||
cmd.Stdin = strings.NewReader(payload)
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
t.Fatalf("git mktag %s error = %v\noutput:\n%s", version, err, output)
|
||||
}
|
||||
mustRun(t, repo, "git", "update-ref", "refs/tags/"+version, strings.TrimSpace(string(output)))
|
||||
}
|
||||
rawTag("v1.0.7", "Release v1.0.7\n\nOSS-Mirror: \n")
|
||||
rawTag("v1.0.8", "Release v1.0.8\r\n\r\nOSS-Mirror: deferred\r\n")
|
||||
|
||||
for _, test := range []struct {
|
||||
version string
|
||||
want string
|
||||
ok bool
|
||||
}{
|
||||
{version: "v1.0.1", want: "enabled", ok: true},
|
||||
{version: "v1.0.2", want: "enabled", ok: true},
|
||||
{version: "v1.0.3", want: "deferred", ok: true},
|
||||
{version: "v1.0.4", want: "invalid OSS-Mirror metadata", ok: false},
|
||||
{version: "v1.0.5", want: "duplicate OSS-Mirror metadata", ok: false},
|
||||
{version: "v1.0.6", want: "must be annotated", ok: false},
|
||||
{version: "v1.0.7", want: "invalid OSS-Mirror metadata", ok: false},
|
||||
{version: "v1.0.8", want: "deferred", ok: true},
|
||||
} {
|
||||
t.Run(test.version, func(t *testing.T) {
|
||||
cmd := exec.Command(scriptPath, test.version)
|
||||
cmd.Dir = repo
|
||||
output, err := cmd.CombinedOutput()
|
||||
if test.ok && err != nil {
|
||||
t.Fatalf("release-tag-oss-mode.sh error = %v\noutput:\n%s", err, output)
|
||||
}
|
||||
if !test.ok && err == nil {
|
||||
t.Fatalf("release-tag-oss-mode.sh unexpectedly accepted %s", test.version)
|
||||
}
|
||||
if !strings.Contains(string(output), test.want) {
|
||||
t.Fatalf("release-tag-oss-mode.sh output missing %q:\n%s", test.want, output)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithdrawReleaseDefersOSSRequirementUntilImmutablePolicyResolution(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
scriptPath, err := filepath.Abs(filepath.Join("..", "..", "scripts", "release", "withdraw-release.sh"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(withdraw script) error = %v", err)
|
||||
}
|
||||
cmd := exec.Command("bash", scriptPath,
|
||||
"v1.2.3",
|
||||
"critical startup regression",
|
||||
"WITHDRAW v1.2.3",
|
||||
)
|
||||
cmd.Env = append(os.Environ(),
|
||||
"GITHUB_TOKEN=test-github-token",
|
||||
"NODE_AUTH_TOKEN=test-npm-token",
|
||||
"GITHUB_REPOSITORY=DingTalk-Real-AI/dingtalk-workspace-cli",
|
||||
"GITHUB_REF_NAME=main",
|
||||
"GITHUB_SHA="+strings.Repeat("a", 40),
|
||||
"GITHUB_RUN_ID=1",
|
||||
"GITHUB_ACTOR=test-operator",
|
||||
"GITHUB_EVENT_DEFAULT_BRANCH=main",
|
||||
"GITHUB_ACTIONS=false",
|
||||
)
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err == nil {
|
||||
t.Fatalf("withdraw-release.sh unexpectedly ran outside GitHub Actions:\n%s", output)
|
||||
}
|
||||
if !strings.Contains(string(output), "withdrawal may run only inside GitHub Actions") {
|
||||
t.Fatalf("withdrawal did not reach the protected execution gate before resolving OSS policy:\n%s", output)
|
||||
}
|
||||
if strings.Contains(string(output), "missing required environment variable OSS_") {
|
||||
t.Fatalf("withdrawal required OSS credentials before resolving immutable tag policy:\n%s", output)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithdrawReleaseRejectsInvalidInputsBeforeMutation(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
scriptPath, err := filepath.Abs(filepath.Join("..", "..", "scripts", "release", "withdraw-release.sh"))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(withdraw script) error = %v", err)
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
version string
|
||||
reason string
|
||||
confirmation string
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "invalid version",
|
||||
version: "1.2.3",
|
||||
reason: "critical startup regression",
|
||||
confirmation: "WITHDRAW 1.2.3",
|
||||
want: "version must be exactly",
|
||||
},
|
||||
{
|
||||
name: "wrong confirmation",
|
||||
version: "v1.2.3",
|
||||
reason: "critical startup regression",
|
||||
confirmation: "v1.2.3",
|
||||
want: "confirmation must be exactly: WITHDRAW v1.2.3",
|
||||
},
|
||||
{
|
||||
name: "multiline reason",
|
||||
version: "v1.2.3",
|
||||
reason: "critical\nregression",
|
||||
confirmation: "WITHDRAW v1.2.3",
|
||||
want: "reason must be a trimmed, printable single line",
|
||||
},
|
||||
}
|
||||
for _, test := range tests {
|
||||
test := test
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
cmd := exec.Command("bash", scriptPath, test.version, test.reason, test.confirmation)
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err == nil {
|
||||
t.Fatalf("withdraw-release.sh unexpectedly accepted invalid input:\n%s", string(output))
|
||||
}
|
||||
if !strings.Contains(string(output), test.want) {
|
||||
t.Fatalf("withdraw-release.sh output missing %q:\n%s", test.want, string(output))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithdrawReleaseRollsBackConfiguredChannelsAndStopsForHomebrewReview(t *testing.T) {
|
||||
for _, test := range []struct {
|
||||
name string
|
||||
ossDeferred bool
|
||||
}{
|
||||
{name: "legacy tag defaults OSS to enabled"},
|
||||
{name: "deferred OSS survives tombstone retry", ossDeferred: true},
|
||||
} {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
testWithdrawReleaseRollsBackConfiguredChannelsAndStopsForHomebrewReview(t, test.ossDeferred)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testWithdrawReleaseRollsBackConfiguredChannelsAndStopsForHomebrewReview(t *testing.T, ossDeferred bool) {
|
||||
root := t.TempDir()
|
||||
sourceRoot, err := filepath.Abs(filepath.Join("..", ".."))
|
||||
if err != nil {
|
||||
t.Fatalf("Abs(source root) error = %v", err)
|
||||
}
|
||||
|
||||
for _, rel := range []string{
|
||||
"scripts/release/withdraw-release.sh",
|
||||
"scripts/release/release-tag-oss-mode.sh",
|
||||
"scripts/release/release-lib.sh",
|
||||
"build/homebrew-release.rb.tmpl",
|
||||
} {
|
||||
copyTestFile(t, filepath.Join(sourceRoot, rel), filepath.Join(root, rel), 0o755)
|
||||
}
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "verify-delivery"), []byte(`#!/bin/sh
|
||||
printf 'delivery %s %s\n' "$1" "$2" >> "$CALL_LOG"
|
||||
exit 0
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "download-release-assets"), []byte(`#!/bin/sh
|
||||
set -eu
|
||||
version="$1"
|
||||
dist="$2"
|
||||
mkdir -p "$dist"
|
||||
for asset in \
|
||||
dws-darwin-amd64.tar.gz dws-darwin-arm64.tar.gz \
|
||||
dws-linux-amd64.tar.gz dws-linux-arm64.tar.gz \
|
||||
dws-windows-amd64.zip dws-windows-arm64.zip \
|
||||
dws-skills.zip; do
|
||||
printf '%s %s\n' "$version" "$asset" > "$dist/$asset"
|
||||
done
|
||||
{
|
||||
for asset in \
|
||||
dws-darwin-amd64.tar.gz dws-darwin-arm64.tar.gz \
|
||||
dws-linux-amd64.tar.gz dws-linux-arm64.tar.gz \
|
||||
dws-windows-amd64.zip dws-windows-arm64.zip \
|
||||
dws-skills.zip; do
|
||||
printf '%064d %s\n' 0 "$asset"
|
||||
done
|
||||
} > "$dist/checksums.txt"
|
||||
printf 'download-assets %s\n' "$version" >> "$CALL_LOG"
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "verify-release-assets"), []byte(`#!/bin/sh
|
||||
set -eu
|
||||
test -f "$DWS_PACKAGE_DIST_DIR/checksums.txt"
|
||||
printf 'verify-assets %s\n' "$1" >> "$CALL_LOG"
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "sync-gitee"), []byte(`#!/bin/sh
|
||||
set -eu
|
||||
git push "$GITEE_GIT_REMOTE" "refs/tags/$VERSION:refs/tags/$VERSION" >/dev/null
|
||||
printf '%s\n' "$VERSION" > "$MOCK_STATE/gitee-release-$VERSION"
|
||||
printf 'gitee-sync %s\n' "$VERSION" >> "$CALL_LOG"
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "publish-homebrew-formula.sh"), []byte(`#!/bin/sh
|
||||
cp "$DWS_FORMULA_SOURCE" "$MOCK_STATE/homebrew-formula"
|
||||
printf 'homebrew-pr %s %s\n' "$DWS_TAP_PR_TITLE" "$DWS_TAP_PR_BRANCH" >> "$CALL_LOG"
|
||||
exit 0
|
||||
`), 0o755)
|
||||
mustWriteFile(t, filepath.Join(root, "scripts", "release", "unexpected-ossutil"), []byte(`#!/bin/sh
|
||||
printf 'unexpected-oss-call\n' >> "$CALL_LOG"
|
||||
exit 99
|
||||
`), 0o755)
|
||||
|
||||
stateDir := filepath.Join(root, "state")
|
||||
fakeBin := filepath.Join(root, "bin")
|
||||
if err := os.MkdirAll(stateDir, 0o755); err != nil {
|
||||
t.Fatalf("MkdirAll(state) error = %v", err)
|
||||
}
|
||||
callLog := filepath.Join(stateDir, "calls.log")
|
||||
writeWithdrawMocks(t, fakeBin)
|
||||
|
||||
mustWriteFile(t, filepath.Join(root, "Formula", "dingtalk-workspace-cli.rb"), []byte(`class DingtalkWorkspaceCli < Formula
|
||||
version "1.0.51"
|
||||
end
|
||||
`), 0o644)
|
||||
mustRun(t, root, "git", "init", "-b", "main")
|
||||
mustRun(t, root, "git", "config", "user.name", "Withdrawal Test")
|
||||
mustRun(t, root, "git", "config", "user.email", "withdrawal@example.com")
|
||||
mustRun(t, root, "git", "add", ".")
|
||||
mustRun(t, root, "git", "commit", "-m", "candidate")
|
||||
mustRun(t, root, "git", "tag", "-a", "v1.0.51", "-m", "Release v1.0.51")
|
||||
|
||||
mustWriteFile(t, filepath.Join(root, "Formula", "dingtalk-workspace-cli.rb"), []byte(`class DingtalkWorkspaceCli < Formula
|
||||
version "1.0.52"
|
||||
end
|
||||
`), 0o644)
|
||||
mustRun(t, root, "git", "add", "Formula/dingtalk-workspace-cli.rb")
|
||||
mustRun(t, root, "git", "commit", "-m", "target")
|
||||
tagArgs := []string{"tag", "-a", "v1.0.52", "-m", "Release v1.0.52"}
|
||||
if ossDeferred {
|
||||
tagArgs = append(tagArgs, "-m", "OSS-Mirror: deferred")
|
||||
}
|
||||
mustRun(t, root, "git", tagArgs...)
|
||||
targetCommit := strings.TrimSpace(mustOutput(t, root, "git", "rev-parse", "HEAD"))
|
||||
targetTagObject := strings.TrimSpace(mustOutput(t, root, "git", "rev-parse", "refs/tags/v1.0.52"))
|
||||
|
||||
origin := filepath.Join(root, "origin.git")
|
||||
gitee := filepath.Join(root, "gitee.git")
|
||||
mustRun(t, root, "git", "init", "--bare", origin)
|
||||
mustRun(t, root, "git", "init", "--bare", gitee)
|
||||
mustRun(t, root, "git", "remote", "add", "origin", origin)
|
||||
mustRun(t, root, "git", "push", "origin", "main", "v1.0.51", "v1.0.52")
|
||||
mustRun(t, root, "git", "push", gitee, "v1.0.52")
|
||||
|
||||
writeJSONFile(t, filepath.Join(stateDir, "release-v1.0.51.json"), map[string]any{
|
||||
"id": 51, "tag_name": "v1.0.51", "name": "v1.0.51", "body": "candidate",
|
||||
"draft": false, "prerelease": false, "immutable": true,
|
||||
})
|
||||
writeJSONFile(t, filepath.Join(stateDir, "release-v1.0.52.json"), map[string]any{
|
||||
"id": 52, "tag_name": "v1.0.52", "name": "v1.0.52", "body": "target",
|
||||
"draft": false, "prerelease": false, "immutable": true,
|
||||
})
|
||||
mustWriteFile(t, filepath.Join(stateDir, "latest"), []byte("v1.0.52\n"), 0o644)
|
||||
mustWriteFile(t, filepath.Join(stateDir, "npm-latest"), []byte("1.0.52\n"), 0o644)
|
||||
mustWriteFile(t, filepath.Join(stateDir, "oss-latest"), []byte("v1.0.52\n"), 0o644)
|
||||
mustWriteFile(t, filepath.Join(stateDir, "gitee-release-v1.0.52"), []byte("752\n"), 0o644)
|
||||
mustWriteFile(t, filepath.Join(stateDir, "github-tag-v1.0.52"), []byte(targetTagObject+"\n"), 0o644)
|
||||
|
||||
scriptPath := filepath.Join(root, "scripts", "release", "withdraw-release.sh")
|
||||
cmd := exec.Command("bash", scriptPath,
|
||||
"v1.0.52",
|
||||
"critical startup regression",
|
||||
"WITHDRAW v1.0.52",
|
||||
)
|
||||
cmd.Dir = root
|
||||
withdrawEnv := append(os.Environ(),
|
||||
"PATH="+fakeBin+string(os.PathListSeparator)+os.Getenv("PATH"),
|
||||
"MOCK_STATE="+stateDir,
|
||||
"CALL_LOG="+callLog,
|
||||
"GITHUB_ACTIONS=true",
|
||||
"GITHUB_REPOSITORY=DingTalk-Real-AI/dingtalk-workspace-cli",
|
||||
"GITHUB_REF_NAME=main",
|
||||
"GITHUB_SHA="+targetCommit,
|
||||
"GITHUB_RUN_ID=12345",
|
||||
"GITHUB_ACTOR=release-operator",
|
||||
"GITHUB_EVENT_DEFAULT_BRANCH=main",
|
||||
"GITHUB_TOKEN=github-token",
|
||||
"NODE_AUTH_TOKEN=npm-token",
|
||||
"DWS_GITEE_ENABLED=true",
|
||||
"GITEE_TOKEN=gitee-token",
|
||||
"GITEE_USER=gitee-user",
|
||||
"GITEE_REPO=DingTalk-Real-AI/dingtalk-workspace-cli",
|
||||
"GITEE_GIT_REMOTE="+gitee,
|
||||
"GITEE_PUBLIC_GIT_REMOTE="+gitee,
|
||||
"HOMEBREW_PR_TOKEN=homebrew-token",
|
||||
"DWS_DELIVERY_VERIFIER="+filepath.Join(root, "scripts", "release", "verify-delivery"),
|
||||
"DWS_GITHUB_DOWNLOAD_HELPER="+filepath.Join(root, "scripts", "release", "download-release-assets"),
|
||||
"DWS_ARTIFACT_VERIFY_HELPER="+filepath.Join(root, "scripts", "release", "verify-release-assets"),
|
||||
"DWS_GITEE_SYNC_HELPER="+filepath.Join(root, "scripts", "release", "sync-gitee"),
|
||||
"ORIGIN_GIT="+origin,
|
||||
)
|
||||
if ossDeferred {
|
||||
withdrawEnv = append(withdrawEnv,
|
||||
"OSS_ACCESS_KEY_ID=",
|
||||
"OSS_ACCESS_KEY_SECRET=",
|
||||
"OSS_ENDPOINT=",
|
||||
"OSS_BUCKET=",
|
||||
"OSSUTIL="+filepath.Join(root, "scripts", "release", "unexpected-ossutil"),
|
||||
)
|
||||
} else {
|
||||
withdrawEnv = append(withdrawEnv,
|
||||
"OSS_ACCESS_KEY_ID=oss-id",
|
||||
"OSS_ACCESS_KEY_SECRET=oss-secret",
|
||||
"OSS_ENDPOINT=https://oss-cn-hangzhou.aliyuncs.com",
|
||||
"OSS_BUCKET=dws-test",
|
||||
"OSSUTIL="+filepath.Join(fakeBin, "ossutil"),
|
||||
)
|
||||
}
|
||||
cmd.Env = withdrawEnv
|
||||
output, err := cmd.CombinedOutput()
|
||||
if err == nil {
|
||||
t.Fatalf("withdrawal must remain failed while Homebrew PR is pending:\n%s", string(output))
|
||||
}
|
||||
if !strings.Contains(string(output), "Homebrew rollback PR requires independent review and merge") {
|
||||
t.Fatalf("withdrawal did not report the Homebrew manual gate:\n%s", string(output))
|
||||
}
|
||||
|
||||
assertFileEquals(t, filepath.Join(stateDir, "latest"), "v1.0.51")
|
||||
assertFileEquals(t, filepath.Join(stateDir, "npm-latest"), "1.0.51")
|
||||
assertFileContains(t, filepath.Join(stateDir, "npm-deprecated"), "WITHDRAWN v1.0.52")
|
||||
if ossDeferred {
|
||||
assertFileEquals(t, filepath.Join(stateDir, "oss-latest"), "v1.0.52")
|
||||
} else {
|
||||
assertFileEquals(t, filepath.Join(stateDir, "oss-latest"), "v1.0.51")
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "oss-removed")); err != nil {
|
||||
t.Fatalf("OSS withdrawn prefix was not removed: %v", err)
|
||||
}
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "gitee-release-v1.0.52")); !os.IsNotExist(err) {
|
||||
t.Fatalf("Gitee release still exists, stat error = %v", err)
|
||||
}
|
||||
if refs := mustOutput(t, root, "git", "ls-remote", gitee, "refs/tags/v1.0.52"); strings.TrimSpace(refs) != "" {
|
||||
t.Fatalf("Gitee still exposes withdrawn tag:\n%s", refs)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "release-v1.0.52.json")); !os.IsNotExist(err) {
|
||||
t.Fatalf("GitHub problem release still exists while Homebrew review is pending, stat error = %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "github-tag-v1.0.52")); !os.IsNotExist(err) {
|
||||
t.Fatalf("GitHub problem tag still exists while Homebrew review is pending, stat error = %v", err)
|
||||
}
|
||||
assertFileContains(t, filepath.Join(stateDir, "tombstone-message"), "Original-Commit: "+targetCommit)
|
||||
assertFileContains(t, filepath.Join(stateDir, "tombstone-message"), "Original-Tag-Object: "+targetTagObject)
|
||||
assertFileContains(t, filepath.Join(stateDir, "tombstone-message"), "Original-Release-ID: 52")
|
||||
ossMode := "enabled"
|
||||
if ossDeferred {
|
||||
ossMode = "deferred"
|
||||
}
|
||||
assertFileContains(t, filepath.Join(stateDir, "tombstone-message"), "OSS-Mirror: "+ossMode)
|
||||
assertFileContains(t, filepath.Join(stateDir, "tombstone-message"), "Reason: critical startup regression")
|
||||
assertFileContains(t, callLog, "homebrew-pr revert: withdraw v1.0.52 and restore v1.0.51")
|
||||
|
||||
calls, err := os.ReadFile(callLog)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(call log) error = %v", err)
|
||||
}
|
||||
logText := string(calls)
|
||||
tombstoneIndex := strings.Index(logText, "tombstone-ref")
|
||||
githubIndex := strings.Index(logText, "github-withdraw")
|
||||
npmIndex := strings.Index(logText, "npm-deprecate")
|
||||
ossIndex := strings.Index(logText, "oss-remove")
|
||||
if tombstoneIndex < 0 || githubIndex < 0 || npmIndex < 0 || (!ossDeferred && ossIndex < 0) {
|
||||
t.Fatalf("missing mutation audit entries:\n%s", logText)
|
||||
}
|
||||
homebrewIndex := strings.Index(logText, "homebrew-pr")
|
||||
if !(tombstoneIndex < homebrewIndex && homebrewIndex < githubIndex && githubIndex < npmIndex) ||
|
||||
(!ossDeferred && npmIndex >= ossIndex) {
|
||||
t.Fatalf("tombstone was not durable before channel mutations:\n%s", logText)
|
||||
}
|
||||
deleteReleaseIndex := strings.Index(logText, "github-delete-release")
|
||||
deleteTagIndex := strings.Index(logText, "github-delete-tag")
|
||||
channelsCompleteIndex := npmIndex
|
||||
if !ossDeferred {
|
||||
channelsCompleteIndex = ossIndex
|
||||
}
|
||||
if deleteReleaseIndex < 0 || deleteTagIndex < 0 || homebrewIndex < 0 ||
|
||||
!(homebrewIndex < channelsCompleteIndex && channelsCompleteIndex < deleteReleaseIndex && deleteReleaseIndex < deleteTagIndex) {
|
||||
t.Fatalf("GitHub problem release was not removed before the Homebrew review pause:\n%s", logText)
|
||||
}
|
||||
if ossDeferred {
|
||||
if strings.Contains(logText, "unexpected-oss-call") || ossIndex >= 0 {
|
||||
t.Fatalf("deferred withdrawal invoked OSS tooling:\n%s", logText)
|
||||
}
|
||||
for _, want := range []string{
|
||||
"OSS mirroring is disabled; no OSS rollback is required.",
|
||||
"OSS did not contain v1.0.52 because mirroring was disabled",
|
||||
} {
|
||||
if !strings.Contains(string(output), want) {
|
||||
t.Fatalf("deferred withdrawal output missing %q:\n%s", want, output)
|
||||
}
|
||||
}
|
||||
}
|
||||
if !ossDeferred {
|
||||
tombstonePath := filepath.Join(stateDir, "tombstone-message")
|
||||
tombstoneMessage, err := os.ReadFile(tombstonePath)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(legacy tombstone) error = %v", err)
|
||||
}
|
||||
legacyMessage := strings.Replace(string(tombstoneMessage), "OSS-Mirror: enabled\n", "", 1)
|
||||
if legacyMessage == string(tombstoneMessage) {
|
||||
t.Fatal("could not convert tombstone fixture to the legacy format")
|
||||
}
|
||||
mustWriteFile(t, tombstonePath, []byte(legacyMessage), 0o644)
|
||||
}
|
||||
|
||||
mergedFormula, err := os.ReadFile(filepath.Join(stateDir, "homebrew-formula"))
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(rendered Homebrew rollback) error = %v", err)
|
||||
}
|
||||
mustWriteFile(t, filepath.Join(root, "Formula", "dingtalk-workspace-cli.rb"), mergedFormula, 0o644)
|
||||
mustRun(t, root, "git", "add", "Formula/dingtalk-workspace-cli.rb")
|
||||
mustRun(t, root, "git", "commit", "-m", "merge Homebrew rollback")
|
||||
mustRun(t, root, "git", "push", "origin", "main")
|
||||
retryCommit := strings.TrimSpace(mustOutput(t, root, "git", "rev-parse", "HEAD"))
|
||||
retryEnv := replaceTestEnv(withdrawEnv,
|
||||
"GITHUB_SHA", retryCommit,
|
||||
"GITHUB_RUN_ID", "12346",
|
||||
)
|
||||
retry := exec.Command("bash", scriptPath,
|
||||
"v1.0.52",
|
||||
"critical startup regression",
|
||||
"WITHDRAW v1.0.52",
|
||||
)
|
||||
retry.Dir = root
|
||||
retry.Env = retryEnv
|
||||
retryOutput, retryErr := retry.CombinedOutput()
|
||||
if retryErr != nil {
|
||||
t.Fatalf("withdrawal retry after Homebrew merge error = %v\noutput:\n%s", retryErr, string(retryOutput))
|
||||
}
|
||||
for _, want := range []string{
|
||||
"Resuming withdrawal for v1.0.52 from exact permanent tombstone metadata.",
|
||||
"Permanent tombstone withdrawn/v1.0.52 already exists",
|
||||
"GitHub Release v1.0.52 was already absent.",
|
||||
"GitHub tag v1.0.52 was already absent.",
|
||||
"Withdrawal completed for all configured distribution channels.",
|
||||
"Already-installed clients cannot be remotely downgraded",
|
||||
} {
|
||||
if !strings.Contains(string(retryOutput), want) {
|
||||
t.Fatalf("withdrawal retry output missing %q:\n%s", want, string(retryOutput))
|
||||
}
|
||||
}
|
||||
if ossDeferred && !strings.Contains(string(retryOutput), "OSS mirroring is disabled; no OSS rollback is required.") {
|
||||
t.Fatalf("deferred tombstone retry did not restore the sealed OSS policy:\n%s", retryOutput)
|
||||
}
|
||||
assertFileEquals(t, filepath.Join(stateDir, "latest"), "v1.0.51")
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "release-v1.0.52.json")); !os.IsNotExist(err) {
|
||||
t.Fatalf("GitHub problem release still exists, stat error = %v", err)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "github-tag-v1.0.52")); !os.IsNotExist(err) {
|
||||
t.Fatalf("GitHub problem tag still exists, stat error = %v", err)
|
||||
}
|
||||
if refs := mustOutput(t, root, "git", "ls-remote", origin, "refs/tags/v1.0.52"); strings.TrimSpace(refs) != "" {
|
||||
t.Fatalf("GitHub origin still exposes withdrawn tag:\n%s", refs)
|
||||
}
|
||||
if _, err := os.Stat(filepath.Join(stateDir, "tombstone-ref")); err != nil {
|
||||
t.Fatalf("permanent withdrawal tombstone is missing: %v", err)
|
||||
}
|
||||
assertFileContains(t, callLog, "github-delete-release")
|
||||
assertFileContains(t, callLog, "github-delete-tag")
|
||||
}
|
||||
|
||||
func copyTestFile(t *testing.T, source, destination string, mode os.FileMode) {
|
||||
t.Helper()
|
||||
content, err := os.ReadFile(source)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", source, err)
|
||||
}
|
||||
mustWriteFile(t, destination, content, mode)
|
||||
}
|
||||
|
||||
func writeJSONFile(t *testing.T, path string, value any) {
|
||||
t.Helper()
|
||||
content, err := json.Marshal(value)
|
||||
if err != nil {
|
||||
t.Fatalf("Marshal(%s) error = %v", path, err)
|
||||
}
|
||||
mustWriteFile(t, path, append(content, '\n'), 0o644)
|
||||
}
|
||||
|
||||
func assertFileEquals(t *testing.T, path, want string) {
|
||||
t.Helper()
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", path, err)
|
||||
}
|
||||
if got := strings.TrimSpace(string(content)); got != want {
|
||||
t.Fatalf("%s = %q, want %q", path, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func assertFileContains(t *testing.T, path, want string) {
|
||||
t.Helper()
|
||||
content, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("ReadFile(%s) error = %v", path, err)
|
||||
}
|
||||
if !strings.Contains(string(content), want) {
|
||||
t.Fatalf("%s missing %q:\n%s", path, want, string(content))
|
||||
}
|
||||
}
|
||||
|
||||
func replaceTestEnv(env []string, replacements ...string) []string {
|
||||
result := append([]string{}, env...)
|
||||
for index := 0; index < len(replacements); index += 2 {
|
||||
key := replacements[index]
|
||||
value := replacements[index+1]
|
||||
prefix := key + "="
|
||||
replaced := false
|
||||
for envIndex, entry := range result {
|
||||
if strings.HasPrefix(entry, prefix) {
|
||||
result[envIndex] = prefix + value
|
||||
replaced = true
|
||||
}
|
||||
}
|
||||
if !replaced {
|
||||
result = append(result, prefix+value)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func writeWithdrawMocks(t *testing.T, fakeBin string) {
|
||||
t.Helper()
|
||||
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "gh"), []byte(`#!/usr/bin/env python3
|
||||
import json
|
||||
import os
|
||||
import pathlib
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
state = pathlib.Path(os.environ["MOCK_STATE"])
|
||||
log = pathlib.Path(os.environ["CALL_LOG"])
|
||||
args = sys.argv[1:]
|
||||
|
||||
def record(value):
|
||||
with log.open("a", encoding="utf-8") as handle:
|
||||
handle.write(value + "\n")
|
||||
|
||||
def value_after(flag, default=""):
|
||||
try:
|
||||
return args[args.index(flag) + 1]
|
||||
except (ValueError, IndexError):
|
||||
return default
|
||||
|
||||
def not_found():
|
||||
print("gh: Not Found (HTTP 404)", file=sys.stderr)
|
||||
raise SystemExit(1)
|
||||
|
||||
if args[:2] == ["release", "download"]:
|
||||
output = pathlib.Path(value_after("--dir"))
|
||||
output.mkdir(parents=True, exist_ok=True)
|
||||
assets = [
|
||||
"dws-darwin-amd64.tar.gz",
|
||||
"dws-darwin-arm64.tar.gz",
|
||||
"dws-linux-amd64.tar.gz",
|
||||
"dws-linux-arm64.tar.gz",
|
||||
"dws-skills.zip",
|
||||
]
|
||||
(output / "checksums.txt").write_text(
|
||||
"".join(("a" * 64) + " " + asset + "\n" for asset in assets),
|
||||
encoding="utf-8",
|
||||
)
|
||||
record("github-download-checksums")
|
||||
raise SystemExit(0)
|
||||
|
||||
if not args or args[0] != "api":
|
||||
raise SystemExit("unsupported gh invocation: " + " ".join(args))
|
||||
|
||||
method = value_after("--method", "GET")
|
||||
skip = {"-H", "--method", "-f", "-F", "--input", "--jq"}
|
||||
endpoint = ""
|
||||
i = 1
|
||||
while i < len(args):
|
||||
if args[i] in skip:
|
||||
i += 2
|
||||
continue
|
||||
if args[i].startswith("repos/"):
|
||||
endpoint = args[i]
|
||||
break
|
||||
i += 1
|
||||
jq = value_after("--jq")
|
||||
|
||||
if endpoint.endswith("/git/ref/heads/main"):
|
||||
payload = {"object": {"sha": os.environ["GITHUB_SHA"]}}
|
||||
elif "/releases/tags/" in endpoint:
|
||||
version = endpoint.rsplit("/", 1)[1]
|
||||
path = state / f"release-{version}.json"
|
||||
if not path.exists():
|
||||
not_found()
|
||||
payload = json.loads(path.read_text(encoding="utf-8"))
|
||||
elif endpoint.endswith("/releases/latest"):
|
||||
payload = {"tag_name": (state / "latest").read_text(encoding="utf-8").strip()}
|
||||
elif "/git/ref/tags/withdrawn/" in endpoint:
|
||||
ref_path = state / "tombstone-ref"
|
||||
if not ref_path.exists():
|
||||
not_found()
|
||||
payload = {"object": {"sha": ref_path.read_text(encoding="utf-8").strip()}}
|
||||
elif endpoint.endswith("/git/ref/tags/v1.0.52") and method == "GET":
|
||||
tag_path = state / "github-tag-v1.0.52"
|
||||
if not tag_path.exists():
|
||||
not_found()
|
||||
payload = {"object": {"sha": tag_path.read_text(encoding="utf-8").strip()}}
|
||||
elif "/git/tags/" in endpoint and method == "GET":
|
||||
payload = {
|
||||
"tag": "withdrawn/v1.0.52",
|
||||
"message": (state / "tombstone-message").read_text(encoding="utf-8"),
|
||||
"object": {
|
||||
"type": "commit",
|
||||
"sha": (state / "tombstone-target").read_text(encoding="utf-8").strip(),
|
||||
},
|
||||
}
|
||||
elif endpoint.endswith("/git/tags") and method == "POST":
|
||||
fields = {}
|
||||
for index, arg in enumerate(args):
|
||||
if arg == "-f":
|
||||
key, value = args[index + 1].split("=", 1)
|
||||
fields[key] = value
|
||||
(state / "tombstone-message").write_text(fields["message"], encoding="utf-8")
|
||||
(state / "tombstone-target").write_text(fields["object"] + "\n", encoding="utf-8")
|
||||
payload = {"sha": "a" * 40}
|
||||
record("tombstone-object")
|
||||
elif endpoint.endswith("/git/refs") and method == "POST":
|
||||
(state / "tombstone-ref").write_text("a" * 40 + "\n", encoding="utf-8")
|
||||
payload = {"ref": "refs/tags/withdrawn/v1.0.52", "object": {"sha": "a" * 40}}
|
||||
record("tombstone-ref")
|
||||
elif endpoint.endswith("/releases/52") and method == "DELETE":
|
||||
release_path = state / "release-v1.0.52.json"
|
||||
if not release_path.exists():
|
||||
not_found()
|
||||
release_path.unlink()
|
||||
(state / "latest").write_text("v1.0.51\n", encoding="utf-8")
|
||||
payload = {}
|
||||
record("github-delete-release")
|
||||
elif endpoint.endswith("/git/refs/tags/v1.0.52") and method == "DELETE":
|
||||
tag_path = state / "github-tag-v1.0.52"
|
||||
if not tag_path.exists():
|
||||
not_found()
|
||||
tag_path.unlink()
|
||||
subprocess.run(
|
||||
["git", f"--git-dir={os.environ['ORIGIN_GIT']}", "update-ref", "-d", "refs/tags/v1.0.52"],
|
||||
check=True,
|
||||
)
|
||||
payload = {}
|
||||
record("github-delete-tag")
|
||||
elif "/releases/" in endpoint and method == "PATCH":
|
||||
release_id = endpoint.rsplit("/", 1)[1]
|
||||
if release_id == "52":
|
||||
path = state / "release-v1.0.52.json"
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
input_path = value_after("--input")
|
||||
patch = json.loads(pathlib.Path(input_path).read_text(encoding="utf-8"))
|
||||
data.update({"name": patch["name"], "body": patch["body"]})
|
||||
path.write_text(json.dumps(data), encoding="utf-8")
|
||||
record("github-withdraw")
|
||||
payload = {}
|
||||
else:
|
||||
raise SystemExit("unsupported gh api endpoint: " + endpoint + " " + method)
|
||||
|
||||
if jq == ".object.sha":
|
||||
print(payload["object"]["sha"])
|
||||
elif jq == ".tag_name":
|
||||
print(payload["tag_name"])
|
||||
elif jq == ".sha":
|
||||
print(payload["sha"])
|
||||
else:
|
||||
print(json.dumps(payload))
|
||||
`), 0o755)
|
||||
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "npm"), []byte(`#!/usr/bin/env python3
|
||||
import os
|
||||
import pathlib
|
||||
import sys
|
||||
|
||||
state = pathlib.Path(os.environ["MOCK_STATE"])
|
||||
log = pathlib.Path(os.environ["CALL_LOG"])
|
||||
args = sys.argv[1:]
|
||||
|
||||
def record(value):
|
||||
with log.open("a", encoding="utf-8") as handle:
|
||||
handle.write(value + "\n")
|
||||
|
||||
if args[0] == "view":
|
||||
spec = args[1]
|
||||
field = args[2]
|
||||
if field == "version":
|
||||
print(spec.rsplit("@", 1)[1])
|
||||
elif field == "deprecated":
|
||||
path = state / "npm-deprecated"
|
||||
if path.exists() and spec.endswith("@1.0.52"):
|
||||
print(path.read_text(encoding="utf-8").strip())
|
||||
elif field == "dist-tags.latest":
|
||||
print((state / "npm-latest").read_text(encoding="utf-8").strip())
|
||||
else:
|
||||
raise SystemExit("unsupported npm view field: " + field)
|
||||
elif args[0] == "deprecate":
|
||||
(state / "npm-deprecated").write_text(args[2] + "\n", encoding="utf-8")
|
||||
record("npm-deprecate")
|
||||
elif args[:2] == ["dist-tag", "add"]:
|
||||
version = args[2].rsplit("@", 1)[1]
|
||||
(state / "npm-latest").write_text(version + "\n", encoding="utf-8")
|
||||
record("npm-pointer")
|
||||
else:
|
||||
raise SystemExit("unsupported npm invocation: " + " ".join(args))
|
||||
`), 0o755)
|
||||
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "ossutil"), []byte(`#!/usr/bin/env python3
|
||||
import os
|
||||
import pathlib
|
||||
import sys
|
||||
|
||||
state = pathlib.Path(os.environ["MOCK_STATE"])
|
||||
log = pathlib.Path(os.environ["CALL_LOG"])
|
||||
args = sys.argv[1:]
|
||||
|
||||
if args[0] == "cp":
|
||||
source, target = args[-2:]
|
||||
if source.startswith("oss://"):
|
||||
if source.endswith("/latest.txt"):
|
||||
stored = state / "oss-latest"
|
||||
else:
|
||||
stored = state / ("oss-object-" + source.rsplit("/", 1)[1])
|
||||
pathlib.Path(target).write_bytes(stored.read_bytes())
|
||||
else:
|
||||
if target.endswith("/latest.txt"):
|
||||
stored = state / "oss-latest"
|
||||
else:
|
||||
stored = state / ("oss-object-" + target.rsplit("/", 1)[1])
|
||||
stored.write_bytes(pathlib.Path(source).read_bytes())
|
||||
elif args[0] == "rm":
|
||||
(state / "oss-removed").write_text("yes\n", encoding="utf-8")
|
||||
with log.open("a", encoding="utf-8") as handle:
|
||||
handle.write("oss-remove\n")
|
||||
elif args[0] == "ls":
|
||||
raise SystemExit(0)
|
||||
else:
|
||||
raise SystemExit("unsupported ossutil invocation: " + " ".join(args))
|
||||
`), 0o755)
|
||||
|
||||
mustWriteFile(t, filepath.Join(fakeBin, "curl"), []byte(`#!/usr/bin/env python3
|
||||
import json
|
||||
import os
|
||||
import pathlib
|
||||
import sys
|
||||
|
||||
state = pathlib.Path(os.environ["MOCK_STATE"])
|
||||
log = pathlib.Path(os.environ["CALL_LOG"])
|
||||
args = sys.argv[1:]
|
||||
url = args[-1]
|
||||
version = "v1.0.51" if "v1.0.51" in url else "v1.0.52"
|
||||
release = state / ("gitee-release-" + version)
|
||||
if "-X" in args and args[args.index("-X") + 1] == "DELETE":
|
||||
(state / "gitee-release-v1.0.52").unlink(missing_ok=True)
|
||||
with log.open("a", encoding="utf-8") as handle:
|
||||
handle.write("gitee-delete-release\n")
|
||||
elif "/releases/tags/" in url:
|
||||
output = pathlib.Path(args[args.index("-o") + 1])
|
||||
if release.exists():
|
||||
release_id = 751 if version == "v1.0.51" else 752
|
||||
output.write_text(json.dumps({"id": release_id}), encoding="utf-8")
|
||||
print("200", end="")
|
||||
else:
|
||||
output.write_text('{"message":"Not Found"}', encoding="utf-8")
|
||||
print("404", end="")
|
||||
else:
|
||||
raise SystemExit(22)
|
||||
`), 0o755)
|
||||
}
|
||||
Reference in New Issue
Block a user