Compare commits

...
Author SHA1 Message Date
修雨 6e070a7e24 ci: tighten PR coverage gate 2026-07-20 17:14:35 +08:00
修雨 e867abd03c Merge pull request #687 from shangguanxuan633-lab/codex/auth-token-manager-complete
fix(auth): unify token resolution and recover rejected tokens
2026-07-20 15:07:07 +08:00
修雨 9d89965de9 Merge branch 'main' into codex/auth-token-manager-complete 2026-07-20 14:53:06 +08:00
修雨 41088bb965 fix(release): derive OSS_REGION for ossutil v2 V4 signing (#692)
ossutil 2.x signs requests with V4 and refuses to run without an
explicit region, so the OSS mirror sync would fail in CI even with
valid credentials. Derive OSS_REGION from the endpoint host
(including -internal variants) and fail fast when it cannot be
derived.
2026-07-20 14:51:41 +08:00
修雨 8259116f15 test(auth): isolate Windows keychain packages 2026-07-20 14:33:17 +08:00
上官玄 9ec1fa0638 test(auth): use synthetic log redaction sentinel 2026-07-20 14:22:18 +08:00
修雨 9afd3be79b Merge branch 'main' into codex/auth-token-manager-complete 2026-07-20 14:15:48 +08:00
修雨 67da5019e3 Merge pull request #689 from DingTalk-Real-AI/codex/fix-beta4-channel-repair
fix(release): recover immutable mirror channels safely
2026-07-20 14:07:54 +08:00
shangguanxuan.sgx b0ded7deb8 fix(auth): retry rejected access tokens safely 2026-07-20 13:37:18 +08:00
shangguanxuan.sgx 22905fc41e fix(auth): unify access token resolution 2026-07-20 12:17:04 +08:00
49 changed files with 2708 additions and 373 deletions
+4 -3
View File
@@ -1,4 +1,4 @@
name: Code Admission — PR 合入门禁
name: CI
on:
push:
@@ -714,9 +714,10 @@ jobs:
- name: Enforce coverage gate
if: needs.lint.outputs.changelog_only != 'true'
env:
COVERAGE_TARGET: "80"
COVERAGE_TARGET: "100"
COVERAGE_ENFORCE_OVERALL: "false"
run: COVERAGE_ADDITIONAL_PROFILE=coverage-shortcut.txt make coverage-gate BASE_REF="$COVERAGE_BASE_REF"
COVERAGE_OVERALL_TOLERANCE: "0"
run: COVERAGE_ADDITIONAL_DIFF_PROFILE=coverage-shortcut.txt make coverage-gate BASE_REF="$COVERAGE_BASE_REF"
- name: Generate coverage report
if: needs.lint.outputs.changelog_only != 'true'
+1 -1
View File
@@ -3,7 +3,7 @@ name: Main Integration — Wukong Overlay
on:
workflow_run:
workflows:
- Code Admission — PR 合入门禁
- CI
types:
- completed
+1
View File
@@ -13,6 +13,7 @@ The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and th
### Fixed
- **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.
+2 -2
View File
@@ -19,8 +19,8 @@ help:
@printf " make policy - Check the built dws plus open-source and Schema policies\n"
@printf " make interface-integrity - Check historical commands and help contracts still work\n"
@printf " make authoritative-interface-integrity BASE_REF=<ref> - Check the Git-owned PR merge-base\n"
@printf " make coverage-gate BASE_REF=<ref> - Enforce overall non-regression and changed-code coverage\n"
@printf " make coverage-gate-platform BASE_REF=<ref> PROFILE=<file> - Enforce native-platform changed-code coverage\n"
@printf " make coverage-gate BASE_REF=<ref> - Enforce overall non-regression and 100%% changed-code coverage\n"
@printf " make coverage-gate-platform BASE_REF=<ref> PROFILE=<file> - Enforce 100%% native changed-code coverage\n"
@printf " make update-interface-baseline - Add new CLI contracts without removing history\n"
@printf " make reset-interface-baseline - DANGEROUS: replace all CLI compatibility history\n"
@printf " make schema-compatibility BASE_REF=<ref> - Check the complete Schema contract against the PR merge-base\n"
+1 -1
View File
@@ -52,7 +52,7 @@ run.
```mermaid
flowchart TB
PR["Pull request"] --> CA["Code Admission — PR 合入门禁"]
PR["Pull request"] --> CA["CI"]
subgraph CA_CHECKS["Nine required contexts"]
L["Lint"]
T["Test"]
+11 -5
View File
@@ -1,4 +1,4 @@
# Code Admission — PR 合入门禁
# CI — PR 合入门禁
The pull-request admission layer has exactly nine required external contexts:
@@ -6,7 +6,7 @@ The pull-request admission layer has exactly nine required external contexts:
|---|---|
| `Lint` | Stable PR revision classification, formatting, `go vet`, and Actionlint |
| `Test` | Race/unit/release-script tests plus fast cross-platform compilation |
| `Coverage` | Overall non-regression and changed-code coverage |
| `Coverage` | Overall non-regression and 100% changed-code coverage |
| `Policy` | Repository policy and the fail-closed CHANGELOG contract |
| `Edition` | Edition contract tests |
| `Interface Integrity` | CLI, Schema, Skill, and stable-release compatibility |
@@ -14,7 +14,7 @@ The pull-request admission layer has exactly nine required external contexts:
| `CLI Smoke` | Offline help for every public top-level command |
| `Mock MCP` | HTTP and stdio MCP lifecycle smoke tests |
The workflow display name is `Code Admission — PR 合入门禁`. Parallel helper
The workflow display name is `CI`. Parallel helper
jobs may implement `Test` and `Coverage`, but they are not ruleset contexts.
Do not require an aggregate alias or a downstream integration check in place of
the nine contracts above.
@@ -81,7 +81,7 @@ successful PR check.
```mermaid
flowchart TB
PR["Pull request"] --> ADMISSION["Code Admission — PR 合入门禁"]
PR["Pull request"] --> ADMISSION["CI"]
ADMISSION --> L["Lint"]
ADMISSION --> T["Test"]
ADMISSION --> C["Coverage"]
@@ -122,7 +122,13 @@ base_ref=$(git merge-base HEAD origin/main)
`make coverage-gate` is an enforcement step, not a profile generator. CI
generates the candidate, supporting, merge-base, and (when risk-selected)
native profiles before the aggregate `Coverage` context evaluates them.
native profiles before the aggregate `Coverage` context evaluates them. The
aggregate and native gates require 100% coverage for changed executable Go
statements. Overall coverage remains an unrounded, zero-tolerance merge-base
non-regression check. Candidate and baseline profiles are evaluated by the
same block-deduplicating checker; supporting policy and shortcut profiles
contribute to changed-code coverage only. The checked-in badge is presentation
only and is never read as a gate input.
Compatibility checks derive authoritative Interface snapshots from the PR
merge-base and the latest reachable stable release. The candidate cannot bless
+1 -1
View File
@@ -2,7 +2,7 @@
发布只走一条链路:本地脚本负责封板、验证并推送 annotated tag;GitHub Actions 负责构建和发布最终产物。不要直接运行 `goreleaser release`,也不要手工补打或移动 tag。
发布前必须完成平台治理:目标 GitHub 仓库已启用 immutable releases,`main` 精确要求 `Code Admission — PR 合入门禁` 的九个 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`,操作机已安装并登录 `gh`。本地脚本会在封 tag 前通过 API 检查 immutable releases、当前 SHA 的全部九个 context 和在途 Release;`v*` tag ruleset 仍需仓库管理员预先配置并由操作人确认。
## 日常只用一个入口
+223
View File
@@ -0,0 +1,223 @@
package app
import (
"context"
"errors"
"os"
"path/filepath"
"sync"
"sync/atomic"
"testing"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
type tokenManagerSnapshotProvider struct {
load func() (*authpkg.TokenData, error)
}
func (p tokenManagerSnapshotProvider) GetAccessToken(context.Context) (string, error) {
data, err := p.load()
if err != nil || data == nil {
return "", err
}
return data.AccessToken, nil
}
func (p tokenManagerSnapshotProvider) GetTokenSnapshot(context.Context) (*authpkg.TokenData, error) {
return p.load()
}
type tokenManagerLegacyGetter struct {
token string
err error
}
func (g tokenManagerLegacyGetter) GetToken() (string, string, error) {
return g.token, "file", g.err
}
func installTokenManagerFakes(t *testing.T, load func() (*authpkg.TokenData, error)) {
t.Helper()
oldProvider, oldLegacy := newAccessTokenProvider, newLegacyTokenManager
oldEdition := edition.Get()
edition.Override(&edition.Hooks{})
newAccessTokenProvider = func(string) accessTokenGetter {
return tokenManagerSnapshotProvider{load: load}
}
newLegacyTokenManager = func(string) legacyTokenGetter {
return tokenManagerLegacyGetter{err: authpkg.ErrTokenDataNotFound}
}
t.Cleanup(func() {
newAccessTokenProvider, newLegacyTokenManager = oldProvider, oldLegacy
edition.Override(oldEdition)
})
}
func TestCrossPlatformCoverageTokenManagerCachesUntilMarkerRevisionChanges(t *testing.T) {
configDir := t.TempDir()
if err := authpkg.WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
var calls atomic.Int32
token := "token-a"
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
calls.Add(1)
return &authpkg.TokenData{AccessToken: token, ExpiresAt: time.Now().Add(time.Hour)}, nil
})
manager := NewTokenManager()
first, err := manager.Get(context.Background(), configDir, "")
if err != nil || first.AccessToken != "token-a" {
t.Fatalf("first token = %#v, %v", first, err)
}
second, err := manager.Get(context.Background(), configDir, "")
if err != nil || second.AccessToken != "token-a" || calls.Load() != 1 {
t.Fatalf("cached token = %#v, %v, calls=%d", second, err, calls.Load())
}
token = "token-b"
if err := authpkg.WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
rotated, err := manager.Get(context.Background(), configDir, "")
if err != nil || rotated.AccessToken != "token-b" || calls.Load() != 2 {
t.Fatalf("rotated token = %#v, %v, calls=%d", rotated, err, calls.Load())
}
}
func TestCrossPlatformCoverageTokenManagerDoesNotCacheWithoutExpiryOrRevision(t *testing.T) {
configDir := t.TempDir()
var calls atomic.Int32
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
calls.Add(1)
return &authpkg.TokenData{AccessToken: "token"}, nil
})
manager := NewTokenManager()
for range 2 {
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
t.Fatal(err)
}
}
if calls.Load() != 2 {
t.Fatalf("provider calls = %d, want 2", calls.Load())
}
}
func TestCrossPlatformCoverageTokenManagerTreatsMalformedMarkerAsUncacheable(t *testing.T) {
configDir := t.TempDir()
if err := os.WriteFile(filepath.Join(configDir, "token.json"), []byte("{"), 0o600); err != nil {
t.Fatal(err)
}
var calls atomic.Int32
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
calls.Add(1)
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
})
manager := NewTokenManager()
for range 2 {
if snapshot, err := manager.Get(context.Background(), configDir, ""); err != nil || snapshot.AccessToken != "token" {
t.Fatalf("snapshot = %#v, error = %v", snapshot, err)
}
}
if calls.Load() != 2 {
t.Fatalf("provider calls = %d, want 2", calls.Load())
}
}
func TestCrossPlatformCoverageTokenManagerDoesNotCacheOpaqueEditionStorageWithProviderFallback(t *testing.T) {
configDir := t.TempDir()
if err := authpkg.WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
var calls atomic.Int32
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
calls.Add(1)
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
})
edition.Override(&edition.Hooks{
LoadToken: func(string) ([]byte, error) { return nil, nil },
TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
return fallback()
},
})
manager := NewTokenManager()
for range 2 {
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
t.Fatal(err)
}
}
if calls.Load() != 2 {
t.Fatalf("provider calls = %d, want 2", calls.Load())
}
}
func TestCrossPlatformCoverageTokenManagerCoalescesConcurrentLoads(t *testing.T) {
configDir := t.TempDir()
if err := authpkg.WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
var calls atomic.Int32
release := make(chan struct{})
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
calls.Add(1)
<-release
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
})
manager := NewTokenManager()
const workers = 8
var wg sync.WaitGroup
wg.Add(workers)
errs := make(chan error, workers)
for range workers {
go func() {
defer wg.Done()
_, err := manager.Get(context.Background(), configDir, "")
errs <- err
}()
}
for calls.Load() == 0 {
time.Sleep(time.Millisecond)
}
close(release)
wg.Wait()
close(errs)
for err := range errs {
if err != nil {
t.Fatal(err)
}
}
if calls.Load() != 1 {
t.Fatalf("provider calls = %d, want 1", calls.Load())
}
}
func TestCrossPlatformCoverageTokenManagerPreservesProviderFailure(t *testing.T) {
configDir := t.TempDir()
want := errors.New("keychain permission denied")
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) { return nil, want })
_, err := NewTokenManager().Get(context.Background(), configDir, "")
if !errors.Is(err, want) {
t.Fatalf("error = %v, want cause %v", err, want)
}
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
t.Fatalf("provider failure was misclassified as missing credentials: %v", err)
}
}
func TestCrossPlatformCoverageTokenResolutionErrorOnlyClassifiesTrueMissingCredential(t *testing.T) {
missing := tokenResolutionError(authpkg.ErrTokenDataNotFound)
var typed interface{ Unwrap() error }
if !errors.As(missing, &typed) || !errors.Is(missing, authpkg.ErrTokenDataNotFound) {
t.Fatalf("missing error = %v", missing)
}
want := errors.New("decrypt failed")
if got := tokenResolutionError(want); !errors.Is(got, want) || errors.Is(got, authpkg.ErrTokenDataNotFound) {
t.Fatalf("storage error = %v", got)
}
if got := tokenResolutionError(context.Canceled); !errors.Is(got, context.Canceled) {
t.Fatalf("cancellation = %v", got)
}
}
+244 -48
View File
@@ -21,19 +21,61 @@ import (
"log/slog"
"path/filepath"
"strings"
"sync"
"time"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
const accessTokenRefreshWindow = 5 * time.Minute
type legacyTokenGetter interface {
GetToken() (string, string, error)
}
type accessTokenSnapshotGetter interface {
GetTokenSnapshot(context.Context) (*authpkg.TokenData, error)
}
// AccessTokenSnapshot is the minimal bearer view needed by the process cache.
// Refresh-token material never leaves the auth package.
type AccessTokenSnapshot struct {
AccessToken string
ExpiresAt time.Time
Source string
}
type tokenManagerKey struct {
configDir string
profile string
}
type tokenManagerEntry struct {
mu sync.Mutex
snapshot AccessTokenSnapshot
revision string
}
// TokenManager is the only process cache for user access tokens. Cache entries
// are isolated by config directory and profile, expiry-aware, and invalidated
// by the credential publication marker written by auth storage.
type TokenManager struct {
mu sync.Mutex
entries map[tokenManagerKey]*tokenManagerEntry
now func() time.Time
}
func NewTokenManager() *TokenManager {
return &TokenManager{entries: make(map[tokenManagerKey]*tokenManagerEntry), now: time.Now}
}
var runtimeTokenManager = NewTokenManager()
var (
newAccessTokenProvider = func(configDir string) accessTokenGetter {
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, disc)
discard := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, discard)
configureOAuthProviderCompatibility(provider, configDir)
return provider
}
@@ -44,64 +86,218 @@ var (
}
)
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
// the same host compatibility hooks as MCP. It mirrors the former body of
// getCachedRuntimeToken (excluding process-level cache and timing).
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
provider := newAccessTokenProvider(configDir)
token, tokenErr := provider.GetAccessToken(ctx)
if tokenErr == nil && strings.TrimSpace(token) != "" {
return strings.TrimSpace(token), nil
}
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
return "", tokenErr
}
if strings.TrimSpace(authpkg.RuntimeProfile()) != "" {
if tokenErr != nil {
return "", tokenErr
}
return "", nil
}
manager := newLegacyTokenManager(configDir)
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
return strings.TrimSpace(leg), nil
}
if tokenErr != nil {
return "", tokenErr
}
return "", nil
}
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
// the active edition config directory, the same process-cached path as MCP is used.
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
if t := strings.TrimSpace(explicitToken); t != "" {
return t, nil
// Get resolves an access token for the active runtime profile.
func (m *TokenManager) Get(ctx context.Context, configDir, explicitToken string) (AccessTokenSnapshot, error) {
if token := strings.TrimSpace(explicitToken); token != "" {
return AccessTokenSnapshot{AccessToken: token, Source: "explicit"}, nil
}
if strings.TrimSpace(configDir) == "" {
return "", fmt.Errorf("config directory is empty")
return AccessTokenSnapshot{}, fmt.Errorf("config directory is empty")
}
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
return tok, nil
key := tokenManagerKey{
configDir: canonicalTokenConfigDir(configDir),
profile: strings.TrimSpace(authpkg.RuntimeProfile()),
}
entry := m.entry(key)
entry.mu.Lock()
defer entry.mu.Unlock()
now := time.Now()
if m != nil && m.now != nil {
now = m.now()
}
revision, present, err := authpkg.ReadTokenMarkerRevision(configDir)
if err != nil {
return AccessTokenSnapshot{}, err
}
if tokenSnapshotUsable(entry.snapshot, now) && present && revision != "" && revision == entry.revision {
return entry.snapshot, nil
}
// Treat the marker and credential as one optimistic snapshot. A concurrent
// login/refresh between the reads causes a retry instead of caching stale A
// under the publication marker for B.
for attempt := 0; attempt < 4; attempt++ {
beforeRevision, beforePresent, err := authpkg.ReadTokenMarkerRevision(configDir)
if err != nil {
return AccessTokenSnapshot{}, err
}
return "", noCredentialsError()
snapshot, err := resolveTokenSnapshotWithEdition(ctx, configDir, key.profile)
if err != nil {
return AccessTokenSnapshot{}, err
}
afterRevision, afterPresent, err := authpkg.ReadTokenMarkerRevision(configDir)
if err != nil {
return AccessTokenSnapshot{}, err
}
if beforePresent != afterPresent || beforeRevision != afterRevision {
continue
}
if strings.TrimSpace(snapshot.AccessToken) == "" {
return AccessTokenSnapshot{}, noCredentialsError()
}
if tokenSnapshotUsable(snapshot, now) && afterPresent && afterRevision != "" {
entry.snapshot = snapshot
entry.revision = afterRevision
} else {
entry.snapshot = AccessTokenSnapshot{}
entry.revision = ""
}
return snapshot, nil
}
tok, err := resolveAccessTokenFromDir(ctx, configDir)
return AccessTokenSnapshot{}, fmt.Errorf("token publication changed repeatedly while resolving credentials")
}
func (m *TokenManager) entry(key tokenManagerKey) *tokenManagerEntry {
m.mu.Lock()
defer m.mu.Unlock()
if m.entries == nil {
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
}
entry := m.entries[key]
if entry == nil {
entry = &tokenManagerEntry{}
m.entries[key] = entry
}
return entry
}
func (m *TokenManager) Invalidate() {
if m == nil {
return
}
m.mu.Lock()
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
m.mu.Unlock()
}
func resolveTokenSnapshotWithEdition(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
hooks := edition.Get()
opaqueStorage := hooks.LoadToken != nil || hooks.SaveToken != nil || hooks.DeleteToken != nil
provider := hooks.TokenProvider
if provider == nil {
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
if err != nil {
return AccessTokenSnapshot{}, err
}
// Opaque edition storage hooks have no publication-revision contract.
// Resolve them on every logical request instead of caching a token that
// may be replaced outside the default auth store.
if opaqueStorage {
snapshot.ExpiresAt = time.Time{}
}
return snapshot, nil
}
var fallbackSnapshot AccessTokenSnapshot
var fallbackCalled bool
token, err := provider(ctx, func() (string, error) {
fallbackCalled = true
var fallbackErr error
fallbackSnapshot, fallbackErr = resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
if fallbackErr != nil {
return "", fallbackErr
}
return fallbackSnapshot.AccessToken, nil
})
if err != nil {
return AccessTokenSnapshot{}, fmt.Errorf("edition token provider: %w", err)
}
token = strings.TrimSpace(token)
if token == "" {
return AccessTokenSnapshot{}, noCredentialsError()
}
if fallbackCalled && token == fallbackSnapshot.AccessToken {
if opaqueStorage {
fallbackSnapshot.ExpiresAt = time.Time{}
}
return fallbackSnapshot, nil
}
// Edition providers expose no lifetime metadata, so resolve them on every
// logical request instead of recreating a process-lifetime string cache.
return AccessTokenSnapshot{AccessToken: token, Source: "edition"}, nil
}
func resolveAccessTokenSnapshotFromDir(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
provider := newAccessTokenProvider(configDir)
if snapshotProvider, ok := provider.(accessTokenSnapshotGetter); ok {
data, err := snapshotProvider.GetTokenSnapshot(ctx)
if err == nil && data != nil && strings.TrimSpace(data.AccessToken) != "" {
return AccessTokenSnapshot{
AccessToken: strings.TrimSpace(data.AccessToken),
ExpiresAt: data.ExpiresAt,
Source: "oauth",
}, nil
}
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
return AccessTokenSnapshot{}, err
}
if strings.TrimSpace(profile) != "" {
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
}
return resolveLegacyToken(configDir, err)
}
token, err := provider.GetAccessToken(ctx)
if err == nil && strings.TrimSpace(token) != "" {
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: "oauth_compat"}, nil
}
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
return AccessTokenSnapshot{}, err
}
if strings.TrimSpace(profile) != "" {
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
}
return resolveLegacyToken(configDir, err)
}
func resolveLegacyToken(configDir string, oauthErr error) (AccessTokenSnapshot, error) {
token, source, err := newLegacyTokenManager(configDir).GetToken()
if err == nil && strings.TrimSpace(token) != "" {
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: source}, nil
}
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
return AccessTokenSnapshot{}, err
}
if oauthErr != nil {
return AccessTokenSnapshot{}, oauthErr
}
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
}
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, authpkg.RuntimeProfile())
if err != nil {
return "", err
}
if tok != "" {
return tok, nil
return snapshot.AccessToken, nil
}
// ResolveAuxiliaryAccessToken resolves every non-runner bearer token through
// the same TokenManager used by MCP tool calls.
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
snapshot, err := runtimeTokenManager.Get(ctx, configDir, explicitToken)
if err != nil {
return "", err
}
return "", noCredentialsError()
return snapshot.AccessToken, nil
}
func tokenSnapshotUsable(snapshot AccessTokenSnapshot, now time.Time) bool {
return strings.TrimSpace(snapshot.AccessToken) != "" &&
!snapshot.ExpiresAt.IsZero() &&
now.Before(snapshot.ExpiresAt.Add(-accessTokenRefreshWindow))
}
func canonicalTokenConfigDir(configDir string) string {
if absolute, err := filepath.Abs(configDir); err == nil {
return filepath.Clean(absolute)
}
return filepath.Clean(configDir)
}
func noCredentialsError() error {
if edition.Get().IsEmbedded {
return fmt.Errorf("认证信息已失效,请重新认证")
return fmt.Errorf("认证信息已失效,请重新认证: %w", authpkg.ErrTokenDataNotFound)
}
return fmt.Errorf("no credentials found, run: dws auth login")
return fmt.Errorf("no credentials found, run: dws auth login: %w", authpkg.ErrTokenDataNotFound)
}
+16 -8
View File
@@ -34,6 +34,10 @@ func (g fakeAccessTokenGetter) GetAccessToken(context.Context) (string, error) {
return g.token, g.err
}
func (g fakeAccessTokenGetter) ForceRefreshRejectedToken(context.Context, string) (string, error) {
return g.token, g.err
}
type fakeLegacyTokenGetter struct {
token string
err error
@@ -212,14 +216,16 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
if _, err := resolveAccessTokenFromDir(context.Background(), "unused"); !errors.Is(err, authpkg.ErrTokenDecryption) {
t.Fatalf("decryption error = %v", err)
}
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{err: errors.New("missing")} }
newAccessTokenProvider = func(string) accessTokenGetter {
return fakeAccessTokenGetter{err: authpkg.ErrTokenDataNotFound}
}
newLegacyTokenManager = func(string) legacyTokenGetter { return fakeLegacyTokenGetter{token: " legacy "} }
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); err != nil || got != "legacy" {
t.Fatalf("legacy token = %q, %v", got, err)
}
authpkg.SetRuntimeProfile("corp:user")
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || err == nil || err.Error() != "missing" {
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
t.Fatalf("explicit profile fallback = token %q error %v, want profile error", got, err)
}
authpkg.SetRuntimeProfile("")
@@ -246,10 +252,10 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
}
func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T) {
oldMark, oldFactory := markAccessTokenStale, newRefreshProvider
oldLoad, oldFactory := loadRefreshTokenData, newRefreshProvider
oldStop := stopStdio
t.Cleanup(func() {
markAccessTokenStale, newRefreshProvider = oldMark, oldFactory
loadRefreshTokenData, newRefreshProvider = oldLoad, oldFactory
stopStdio = oldStop
stdioMu.Lock()
stdioClients = make(map[string]*transport.StdioClient)
@@ -257,11 +263,13 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
})
fail := errors.New("failure")
_ = oldFactory(t.TempDir())
markAccessTokenStale = func(string) error { return fail }
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) { return nil, fail }
if _, err := ForceRefreshAccessToken(context.Background(), "config"); !errors.Is(err, fail) {
t.Fatalf("mark stale error = %v", err)
t.Fatalf("load rejected token error = %v", err)
}
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) {
return &authpkg.TokenData{AccessToken: "rejected"}, nil
}
markAccessTokenStale = func(string) error { return nil }
for _, tc := range []struct {
getter fakeAccessTokenGetter
want string
@@ -270,7 +278,7 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
{getter: fakeAccessTokenGetter{token: " "}, want: "empty"},
{getter: fakeAccessTokenGetter{token: " refreshed "}},
} {
newRefreshProvider = func(string) accessTokenGetter { return tc.getter }
newRefreshProvider = func(string) rejectedAccessTokenRefresher { return tc.getter }
got, err := ForceRefreshAccessToken(context.Background(), "config")
if tc.want != "" && (err == nil || !strings.Contains(err.Error(), tc.want)) {
t.Fatalf("refresh error = %v, want %q", err, tc.want)
+136
View File
@@ -15,6 +15,15 @@ package app
import (
"context"
"errors"
"fmt"
"log/slog"
"strings"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/logging"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/authretry"
)
// authRetryingKey marks a context that has already attempted one
@@ -23,8 +32,35 @@ import (
// to the user instead.
type authRetryingKeyType struct{}
type authRefreshFailureError struct {
rejection error
refresh error
}
func (e *authRefreshFailureError) Error() string {
return "automatic access token refresh failed"
}
func (e *authRefreshFailureError) Unwrap() []error {
if e == nil {
return nil
}
return []error{e.rejection, e.refresh}
}
var authRetryingKey = authRetryingKeyType{}
var (
runnerForceRefreshRejectedAccessToken = forceRefreshRejectedAccessToken
runnerExecuteAuthRetry func(*runtimeRunner, context.Context, string, executor.Invocation) (executor.Result, error)
)
func init() {
runnerExecuteAuthRetry = func(r *runtimeRunner, ctx context.Context, endpoint string, invocation executor.Invocation) (executor.Result, error) {
return r.executeInvocation(ctx, endpoint, invocation)
}
}
// IsAuthRetrying reports whether the current context is already inside an
// AuthRefreshRequired retry. Mirrors IsPatRetrying.
func IsAuthRetrying(ctx context.Context) bool {
@@ -34,3 +70,103 @@ func IsAuthRetrying(ctx context.Context) bool {
v, _ := ctx.Value(authRetryingKey).(bool)
return v
}
func withAuthRetrying(ctx context.Context) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, authRetryingKey, true)
}
func authRefreshLogger() *slog.Logger {
if logger := FileLoggerInstance(); logger != nil {
return logger
}
return slog.Default()
}
func (r *runtimeRunner) managesRuntimeOAuth(hasPluginAuth bool) bool {
if r == nil || hasPluginAuth {
return false
}
return r.globalFlags == nil || strings.TrimSpace(r.globalFlags.Token) == ""
}
// retryAuthRefreshRequired consumes only the explicit edition marker. It does
// not infer retryability from free text, generic auth categories, HTTP 403, or
// ordinary business errors.
func (r *runtimeRunner) retryAuthRefreshRequired(
ctx context.Context,
endpoint string,
invocation executor.Invocation,
rejectedAccessToken string,
markerErr error,
hasPluginAuth bool,
) (executor.Result, error, bool) {
marker, marked := authretry.As(markerErr)
if !marked {
return executor.Result{}, nil, false
}
cause := marker.Cause
if cause == nil {
cause = markerErr
}
// Explicit --token and plugin credentials are not backed by the default
// OAuth refresh store. Preserve the overlay cause without mutating an
// unrelated persisted login.
if !r.managesRuntimeOAuth(hasPluginAuth) {
return executor.Result{}, cause, true
}
if IsAuthRetrying(ctx) {
authRefreshLogger().Warn("auth.runtime.refresh.retry_exhausted",
"product", invocation.CanonicalProduct,
"tool", invocation.Tool,
)
return executor.Result{}, cause, true
}
if _, err := runnerForceRefreshRejectedAccessToken(ctx, defaultConfigDir(), rejectedAccessToken); err != nil {
// Keep every log credential-safe. The returned error chain retains the
// complete cause for in-process diagnosis; even DWS_DEBUG_AUTH must not
// serialize an OAuth response body or other attacker-controlled text.
authRefreshLogger().Warn("auth.runtime.refresh.failed",
"product", invocation.CanonicalProduct,
"tool", invocation.Tool,
"stage", "force_refresh_rejected_token",
"error_type", fmt.Sprintf("%T", err),
)
logging.AuthDebug("auth.runtime.refresh.failed.detail",
"product", invocation.CanonicalProduct,
"tool", invocation.Tool,
"stage", "force_refresh_rejected_token",
"error_type", fmt.Sprintf("%T", err),
)
combined := &authRefreshFailureError{rejection: cause, refresh: err}
return executor.Result{}, apperrors.NewAuth(
"automatic access token refresh failed",
apperrors.WithOperation("auth/token/refresh"),
apperrors.WithReason("auth_refresh_failed"),
apperrors.WithHint("本地凭证已保留;可稍后重试,若持续失败请查看认证诊断日志。"),
apperrors.WithCause(combined),
), true
}
logging.AuthDebug("auth.runtime.refresh.succeeded",
"product", invocation.CanonicalProduct,
"tool", invocation.Tool,
)
result, err := runnerExecuteAuthRetry(r, withAuthRetrying(ctx), endpoint, invocation)
return result, err, true
}
// isRefreshableTransportAuthError deliberately excludes HTTP/RPC 403 and
// generic CategoryAuth values. OnAuthError may request a refresh only for an
// exact transport-level unauthorized signal.
func isRefreshableTransportAuthError(err error) bool {
var typed *apperrors.Error
if !errors.As(err, &typed) || typed.Category != apperrors.CategoryAuth {
return false
}
return typed.Reason == "http_401" || typed.RPCCode == 401
}
@@ -0,0 +1,354 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package app
import (
"bytes"
"context"
"errors"
"log/slog"
"strings"
"testing"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/audit"
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
apperrors "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/errors"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/executor"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/transport"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/authretry"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
func installAuthRefreshRunnerSeams(t *testing.T) {
t.Helper()
previousHooks := edition.Get()
previousCall := runnerCallTool
previousPreflight := runnerPreflightDocDownload
previousRefresh := runnerForceRefreshRejectedAccessToken
previousRetry := runnerExecuteAuthRetry
previousCapture := runnerCaptureRuntimeFailure
previousProfile := authpkg.RuntimeProfile()
pluginAuthMu.Lock()
previousPlugins := pluginAuthRegistry
pluginAuthRegistry = make(map[string]*PluginAuth)
pluginAuthMu.Unlock()
runnerPreflightDocDownload = func(*runtimeRunner, context.Context, *transport.Client, string, executor.Invocation) error {
return nil
}
runnerCaptureRuntimeFailure = func(executor.Invocation, error, error) {}
authpkg.SetRuntimeProfile("")
runtimeTokenManager.Invalidate()
t.Setenv("DWS_CONFIG_DIR", "")
t.Setenv("DWS_DEBUG_AUTH", "0")
t.Cleanup(func() {
edition.Override(previousHooks)
runnerCallTool = previousCall
runnerPreflightDocDownload = previousPreflight
runnerForceRefreshRejectedAccessToken = previousRefresh
runnerExecuteAuthRetry = previousRetry
runnerCaptureRuntimeFailure = previousCapture
authpkg.SetRuntimeProfile(previousProfile)
runtimeTokenManager.Invalidate()
pluginAuthMu.Lock()
pluginAuthRegistry = previousPlugins
pluginAuthMu.Unlock()
})
}
func authRefreshTestRunner(flags *GlobalFlags) *runtimeRunner {
return &runtimeRunner{
transport: transport.NewClient(nil),
globalFlags: flags,
auditSink: audit.NopSink{},
}
}
func authRefreshTestInvocation() executor.Invocation {
return executor.Invocation{
CanonicalProduct: "auth-retry-test-product",
Tool: "test_tool",
Params: map[string]any{"value": "safe"},
}
}
func authRefreshTokenHooks(configDir string, token *string, classify func(map[string]any) error) *edition.Hooks {
return &edition.Hooks{
ConfigDir: func() string { return configDir },
TokenProvider: func(context.Context, func() (string, error)) (string, error) {
return *token, nil
},
ClassifyToolResult: classify,
}
}
func TestCrossPlatformCoverageRunnerRetriesEditionAuthMarkerOnce(t *testing.T) {
installAuthRefreshRunnerSeams(t)
configDir := t.TempDir()
token := "old-access"
rejection := apperrors.NewAuth("server rejected access token", apperrors.WithReason("access_token_rejected"))
edition.Override(authRefreshTokenHooks(configDir, &token, func(content map[string]any) error {
if expired, _ := content["expired"].(bool); expired {
return &authretry.AuthRefreshRequired{Cause: rejection}
}
return nil
}))
var callTokens []string
runnerCallTool = func(client *transport.Client, _ context.Context, _, _ string, _ map[string]any) (transport.ToolCallResult, error) {
callTokens = append(callTokens, client.AuthToken)
if len(callTokens) == 1 {
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
}
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(_ context.Context, gotDir, rejected string) (string, error) {
refreshCalls++
if gotDir != configDir || rejected != "old-access" {
t.Fatalf("refresh input = dir %q token %q", gotDir, rejected)
}
token = "new-access"
return token, nil
}
result, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
if err != nil {
t.Fatal(err)
}
if refreshCalls != 1 || len(callTokens) != 2 || callTokens[0] != "old-access" || callTokens[1] != "new-access" {
t.Fatalf("refreshes=%d call tokens=%v", refreshCalls, callTokens)
}
content, _ := result.Response["content"].(map[string]any)
if content["value"] != "ok" || content["success"] != true {
t.Fatalf("result content = %#v", content)
}
}
func TestCrossPlatformCoverageRunnerRefreshFailurePreservesBothCausesAndSafeLog(t *testing.T) {
installAuthRefreshRunnerSeams(t)
t.Setenv("DWS_DEBUG_AUTH", "1")
configDir := t.TempDir()
token := "old-access"
rejection := apperrors.NewAuth("server rejected access token", apperrors.WithReason("access_token_rejected"))
edition.Override(authRefreshTokenHooks(configDir, &token, func(map[string]any) error {
return &authretry.AuthRefreshRequired{Cause: rejection}
}))
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
}
refreshErr := errors.New(`oauth refresh response parse failed: body={"access_token":"access-token-secret","refresh_token":"refresh-token-secret","uid":"uid-secret-value"}`)
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
return "", refreshErr
}
var logs bytes.Buffer
previousLogger := slog.Default()
slog.SetDefault(slog.New(slog.NewJSONHandler(&logs, &slog.HandlerOptions{Level: slog.LevelDebug})))
t.Cleanup(func() { slog.SetDefault(previousLogger) })
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
if !errors.Is(err, rejection) || !errors.Is(err, refreshErr) {
t.Fatalf("error = %v, want rejection and refresh causes", err)
}
var typed *apperrors.Error
if !errors.As(err, &typed) || typed.Category != apperrors.CategoryAuth || typed.Reason != "auth_refresh_failed" || typed.Operation != "auth/token/refresh" {
t.Fatalf("refresh envelope = %#v", typed)
}
var rendered bytes.Buffer
if printErr := apperrors.PrintJSON(&rendered, err); printErr != nil {
t.Fatal(printErr)
}
for _, want := range []string{`"category": "auth"`, `"reason": "auth_refresh_failed"`, `"operation": "auth/token/refresh"`} {
if !strings.Contains(rendered.String(), want) {
t.Fatalf("structured stderr missing %s: %s", want, rendered.String())
}
}
for _, secret := range []string{"access-token-secret", "refresh-token-secret", "uid-secret-value"} {
if strings.Contains(err.Error(), secret) || strings.Contains(logs.String(), secret) || strings.Contains(rendered.String(), secret) {
t.Fatalf("auth output leaked %q: error=%q logs=%s stderr=%s", secret, err, logs.String(), rendered.String())
}
}
for _, want := range []string{"auth.runtime.refresh.failed", "auth.runtime.refresh.failed.detail", "force_refresh_rejected_token", "error_type"} {
if !strings.Contains(logs.String(), want) {
t.Fatalf("safe refresh log missing %q: %s", want, logs.String())
}
}
}
func TestCrossPlatformCoverageRunnerSecondEditionMarkerReturnsSecondCause(t *testing.T) {
installAuthRefreshRunnerSeams(t)
configDir := t.TempDir()
token := "old-access"
firstCause := errors.New("first rejection")
secondCause := errors.New("second rejection")
edition.Override(authRefreshTokenHooks(configDir, &token, func(content map[string]any) error {
attempt, _ := content["attempt"].(int)
if attempt == 1 {
return &authretry.AuthRefreshRequired{Cause: firstCause}
}
return &authretry.AuthRefreshRequired{Cause: secondCause}
}))
calls := 0
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
calls++
return transport.ToolCallResult{Content: map[string]any{"attempt": calls}}, nil
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
refreshCalls++
token = "new-access"
return token, nil
}
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
if !errors.Is(err, secondCause) || errors.Is(err, firstCause) {
t.Fatalf("error = %v, want only second rejection cause", err)
}
if calls != 2 || refreshCalls != 1 {
t.Fatalf("calls=%d refreshes=%d", calls, refreshCalls)
}
}
func TestCrossPlatformCoverageRunnerOnAuthErrorOnlyRetriesExactUnauthorized(t *testing.T) {
t.Run("http 401 marker retries once", func(t *testing.T) {
installAuthRefreshRunnerSeams(t)
configDir := t.TempDir()
token := "old-access"
rejection := errors.New("transport rejected token")
hookCalls := 0
hooks := authRefreshTokenHooks(configDir, &token, nil)
hooks.OnAuthError = func(string, error) error {
hookCalls++
return &authretry.AuthRefreshRequired{Cause: rejection}
}
edition.Override(hooks)
calls := 0
var callTokens []string
runnerCallTool = func(client *transport.Client, _ context.Context, _, _ string, _ map[string]any) (transport.ToolCallResult, error) {
calls++
callTokens = append(callTokens, client.AuthToken)
if calls == 1 {
return transport.ToolCallResult{}, apperrors.NewAuth("unauthorized", apperrors.WithReason("http_401"))
}
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
refreshCalls++
token = "new-access"
return token, nil
}
if _, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation()); err != nil {
t.Fatal(err)
}
if hookCalls != 1 || refreshCalls != 1 || calls != 2 || strings.Join(callTokens, ",") != "old-access,new-access" {
t.Fatalf("hook=%d refresh=%d calls=%d tokens=%v", hookCalls, refreshCalls, calls, callTokens)
}
})
for _, tc := range []struct {
name string
err error
}{
{name: "http 403", err: apperrors.NewAuth("forbidden", apperrors.WithReason("http_403"))},
{name: "ordinary auth", err: apperrors.NewAuth("load failed", apperrors.WithReason("auth_load_failed"))},
} {
t.Run(tc.name+" does not enter hook", func(t *testing.T) {
installAuthRefreshRunnerSeams(t)
configDir := t.TempDir()
token := "old-access"
hookCalls := 0
hooks := authRefreshTokenHooks(configDir, &token, nil)
hooks.OnAuthError = func(string, error) error {
hookCalls++
return &authretry.AuthRefreshRequired{Cause: errors.New("must not run")}
}
edition.Override(hooks)
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
return transport.ToolCallResult{}, tc.err
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
refreshCalls++
return "", nil
}
_, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
if !errors.Is(err, tc.err) || hookCalls != 0 || refreshCalls != 0 {
t.Fatalf("error=%v hook=%d refresh=%d", err, hookCalls, refreshCalls)
}
})
}
}
func TestCrossPlatformCoverageRunnerDoesNotRefreshExplicitTokenMarker(t *testing.T) {
installAuthRefreshRunnerSeams(t)
rejection := errors.New("explicit token rejected")
edition.Override(&edition.Hooks{ClassifyToolResult: func(map[string]any) error {
return &authretry.AuthRefreshRequired{Cause: rejection}
}})
calls := 0
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
calls++
return transport.ToolCallResult{Content: map[string]any{"expired": true}}, nil
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
refreshCalls++
return "", nil
}
_, err := authRefreshTestRunner(&GlobalFlags{Token: "explicit-token"}).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation())
if !errors.Is(err, rejection) || calls != 1 || refreshCalls != 0 {
t.Fatalf("error=%v calls=%d refresh=%d", err, calls, refreshCalls)
}
}
func TestCrossPlatformCoverageRunnerRetriesPreflightEditionMarkerOnce(t *testing.T) {
installAuthRefreshRunnerSeams(t)
configDir := t.TempDir()
token := "old-access"
rejection := errors.New("preflight token rejected")
edition.Override(authRefreshTokenHooks(configDir, &token, nil))
preflightCalls := 0
runnerPreflightDocDownload = func(*runtimeRunner, context.Context, *transport.Client, string, executor.Invocation) error {
preflightCalls++
if preflightCalls == 1 {
return &authretry.AuthRefreshRequired{Cause: rejection}
}
return nil
}
toolCalls := 0
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
toolCalls++
return transport.ToolCallResult{Content: map[string]any{"value": "ok"}}, nil
}
refreshCalls := 0
runnerForceRefreshRejectedAccessToken = func(context.Context, string, string) (string, error) {
refreshCalls++
token = "new-access"
return token, nil
}
if _, err := authRefreshTestRunner(nil).executeInvocation(context.Background(), "https://example.test", authRefreshTestInvocation()); err != nil {
t.Fatal(err)
}
if preflightCalls != 2 || toolCalls != 1 || refreshCalls != 1 {
t.Fatalf("preflights=%d tools=%d refreshes=%d", preflightCalls, toolCalls, refreshCalls)
}
}
+3 -1
View File
@@ -554,7 +554,7 @@ func TestCrossPlatformCoverageRecoveryRuntimeHTTP(t *testing.T) {
defer server.Close()
SetDynamicServers([]mcptypes.ServerDescriptor{{Endpoint: server.URL, CLI: mcptypes.CLIOverlay{ID: "devdoc", Tools: []mcptypes.CLITool{{Name: "search_open_platform_docs_rag"}}}}})
t.Cleanup(func() { SetDynamicServers(nil) })
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client())}
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client()), flags: &GlobalFlags{Token: "token"}}
got, err := runtime.Search(context.Background(), "query", recovery.RecoveryContext{ToolName: "search"})
if err != nil || got.DocSearch.Status != "success" || len(got.KBHits) == 0 {
t.Fatalf("recovery search = %#v %v", got, err)
@@ -1650,6 +1650,8 @@ func TestCrossPlatformCoveragePersonalSubscriptionAndSourceCoverage(t *testing.T
}
func TestCrossPlatformCoveragePersonalEventCommandRuntimeCoverage(t *testing.T) {
authpkg.SetRuntimeProfile("")
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
AccessToken: "access", RefreshToken: "refresh", ExpiresAt: time.Now().Add(time.Hour),
CorpID: "corp", UserID: "user", ClientID: "client",
+2 -1
View File
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
@@ -164,7 +165,7 @@ func TestCrossPlatformCoverageRawAPIAndTokenCoverage(t *testing.T) {
}
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{} }
missing := t.TempDir()
if got, err := resolveAccessTokenFromDir(context.Background(), missing); err != nil || got != "" {
if got, err := resolveAccessTokenFromDir(context.Background(), missing); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
t.Fatalf("missing access token = %q, %v", got, err)
}
if _, err := ResolveAuxiliaryAccessToken(context.Background(), missing, ""); err == nil {
+5 -11
View File
@@ -413,7 +413,7 @@ func eventStreamBusID(streamOpts eventStreamTicketOptions) string {
return "portal-ticket-normal:" + sourceID
}
func newEventSource(ctx context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
func newEventSource(_ context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
if !streamOpts.enabled() {
return eventNewDingtalkSource(source.Config{
ClientID: clientID,
@@ -421,14 +421,6 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
})
}
token, err := eventResolveAccessToken(ctx, configDir, "")
if err != nil {
return nil, fmt.Errorf("event stream ticket: resolve user token: %w", err)
}
if strings.TrimSpace(token) == "" {
return nil, errors.New("event stream ticket: empty user token")
}
portalClientID := clientID
portalClientSecret := clientSecret
if streamOpts.usesPortalNormalMode() {
@@ -440,8 +432,10 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
ClientID: portalClientID,
ClientSecret: portalClientSecret,
PortalTicket: &source.PortalTicketConfig{
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
AccessToken: token,
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
AccessTokenProvider: func(ctx context.Context) (string, error) {
return eventResolveAccessToken(ctx, configDir, "")
},
SourceID: eventStreamSourceID(streamOpts.SourceID),
Mode: streamOpts.Mode,
ClientID: portalClientID,
@@ -132,14 +132,18 @@ func TestCrossPlatformCoverageEventSourcesAndForegroundCoverage(t *testing.T) {
if _, err := newEventSource(context.Background(), "config", "client", "secret", eventStreamTicketOptions{}); err != nil {
t.Fatal(err)
}
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
stream := eventStreamTicketOptions{Mode: "custom"}
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); !errors.Is(err, fail) {
t.Fatalf("stream token error = %v", err)
var captured source.Config
eventNewDingtalkSource = func(cfg source.Config, _ ...source.SourceOption) (*source.DingtalkSource, error) {
captured = cfg
return &source.DingtalkSource{}, nil
}
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return " ", nil }
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err == nil {
t.Fatal("empty stream token succeeded")
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err != nil {
t.Fatalf("stream source construction = %v", err)
}
if _, err := captured.PortalTicket.AccessTokenProvider(context.Background()); !errors.Is(err, fail) {
t.Fatalf("stream token provider error = %v", err)
}
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
for _, mode := range []string{"custom", "normal"} {
+19 -5
View File
@@ -259,7 +259,7 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
return personalConsumeRun(ctx, cfg)
}
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
sub, eventKey, ruleType, err := personalEnsureSubscription(ctx, client, identity, opts)
if err != nil {
return fmt.Errorf("event consume --as user: %w", err)
@@ -498,7 +498,7 @@ func runPersonalEventStatus(c *cobra.Command, opts personalStatusOptions) error
if status == "" || status == "all" {
status = ""
}
subs, err := personalListSubscriptions(personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
subs, err := personalListSubscriptions(newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
Status: status,
EventKey: opts.EventKey,
SubscribeID: opts.SubscribeID,
@@ -613,7 +613,7 @@ func runPersonalEventStop(c *cobra.Command, opts personalStopOptions) error {
if err != nil {
return fmt.Errorf("event stop --as user: %w", err)
}
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
for _, id := range subscribeIDs {
if err := personalDeleteSubscription(client, ctx, id); err != nil {
return fmt.Errorf("event stop --as user: cancel subscription %s: %w", id, err)
@@ -723,7 +723,10 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
if err != nil {
return personal.Identity{}, err
}
tokenData, _ := personalLoadTokenData(configDir)
tokenData, err := personalLoadTokenData(configDir)
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
return personal.Identity{}, fmt.Errorf("load OAuth identity metadata: %w", err)
}
var corpID, userID, clientID, refreshToken string
if tokenData != nil {
corpID = tokenData.CorpID
@@ -769,6 +772,15 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
}, nil
}
func newPersonalEventControlClient(configDir, baseURL string, identity personal.Identity) *personal.Client {
identity.AccessToken = ""
client := personal.NewClient(baseURL, identity)
client.AccessTokenProvider = func(ctx context.Context) (string, error) {
return personalResolveAuxiliaryAccessToken(ctx, configDir, "")
}
return client
}
func personalTokenSubject(kind, token string) string {
token = strings.TrimSpace(token)
if token == "" {
@@ -817,7 +829,9 @@ func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptio
}
_ = ctx
return source.NewPersonal(source.PersonalConfig{
AccessToken: opts.Identity.AccessToken,
AccessTokenProvider: func(ctx context.Context) (string, error) {
return personalResolveAuxiliaryAccessToken(ctx, opts.ConfigDir, "")
},
ClientID: clientID,
ClientSecret: clientSecret,
SourceID: opts.Identity.SourceID,
+27 -16
View File
@@ -27,9 +27,13 @@ type accessTokenGetter interface {
GetAccessToken(context.Context) (string, error)
}
type rejectedAccessTokenRefresher interface {
ForceRefreshRejectedToken(context.Context, string) (string, error)
}
var (
markAccessTokenStale = authpkg.MarkAccessTokenStale
newRefreshProvider = func(configDir string) accessTokenGetter {
loadRefreshTokenData = authpkg.LoadTokenData
newRefreshProvider = func(configDir string) rejectedAccessTokenRefresher {
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, disc)
configureOAuthProviderCompatibility(provider, configDir)
@@ -42,26 +46,33 @@ var (
// server-side rejection (HTTP 401 or business code such as
// TOKEN_VERIFIED_FAILED) on what locally appeared to be a still-valid token.
//
// Steps:
// 1. MarkAccessTokenStale rewrites ExpiresAt to a past instant so
// OAuthProvider.GetAccessToken's fast-path will miss.
// 2. NewOAuthProvider + GetAccessToken triggers lockedRefresh, which uses the
// existing dual-layer lock (process + file) to serialize concurrent
// refresh attempts across goroutines and processes.
// 3. ResetRuntimeTokenCache clears the per-process sync.Once cache so the
// next resolveAuthToken call re-reads from disk.
//
// Existing OAuthProvider.GetAccessToken behaviour is unchanged; this helper
// is the only entry point that orchestrates "force refresh" semantics.
// It snapshots the current access token, then delegates to the OAuth
// provider's dual-locked compare-and-refresh operation. If another caller has
// already rotated the token, that newer token is reused without another
// refresh request.
func ForceRefreshAccessToken(ctx context.Context, configDir string) (string, error) {
if strings.TrimSpace(configDir) == "" {
return "", fmt.Errorf("config directory is empty")
}
if err := markAccessTokenStale(configDir); err != nil {
return "", fmt.Errorf("mark access token stale: %w", err)
data, err := loadRefreshTokenData(configDir)
if err != nil {
return "", err
}
if data == nil || strings.TrimSpace(data.AccessToken) == "" {
return "", fmt.Errorf("stored access token is empty")
}
return forceRefreshRejectedAccessToken(ctx, configDir, data.AccessToken)
}
func forceRefreshRejectedAccessToken(ctx context.Context, configDir, rejectedAccessToken string) (string, error) {
if strings.TrimSpace(configDir) == "" {
return "", fmt.Errorf("config directory is empty")
}
if strings.TrimSpace(rejectedAccessToken) == "" {
return "", fmt.Errorf("rejected access token is empty")
}
provider := newRefreshProvider(configDir)
tok, err := provider.GetAccessToken(ctx)
tok, err := provider.ForceRefreshRejectedToken(ctx, rejectedAccessToken)
if err != nil {
return "", err
}
+20 -20
View File
@@ -56,7 +56,7 @@ var openBrowserFunc = tryOpenBrowser
var (
patAuthorizationTimeout = PatAuthRetryTimeout
patAuthorizationPollInterval = PatAuthPollInterval
patLoadTokenData = authpkg.LoadTokenData
patResolveAccessToken = ResolveAuxiliaryAccessToken
patWaitForAuthorization = WaitForPatAuthorization
patPollDeviceFlowWithInterval = pollPatDeviceFlowWithInterval
patSaveAppConfig = authpkg.SaveAppConfig
@@ -272,7 +272,7 @@ func patAuthorizationURIFromData(data map[string]any) string {
// WaitForPatAuthorization polls until the user completes authorization or timeout.
// It returns true if authorization was completed, false if timed out or cancelled.
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) (bool, error) {
timeout := patAuthorizationTimeout
deadline := time.Now().Add(timeout)
pollTicker := time.NewTicker(patAuthorizationPollInterval)
@@ -290,27 +290,26 @@ func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Wr
select {
case <-ctx.Done():
fmt.Fprintf(output, "%s 操作已取消\n", tui.StateMark("error"))
return false
return false, ctx.Err()
case <-time.After(time.Until(deadline)):
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", tui.StateMark("error"), timeout)
fmt.Fprintf(output, " %s 请重新执行命令\n", tui.Dim("ℹ"))
return false
return false, nil
case <-pollTicker.C:
pollCount++
elapsed := time.Since(start).Truncate(time.Second)
remaining := time.Until(deadline).Truncate(time.Second)
// Check if token is now valid
tokenData, err := patLoadTokenData(configDir)
if err == nil && tokenData != nil {
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
fmt.Fprintln(output)
return true
}
// Check the same resolver used by every outbound bearer request.
if _, err := patResolveAccessToken(ctx, configDir, ""); err == nil {
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
fmt.Fprintln(output)
return true, nil
} else if !stderrors.Is(err, authpkg.ErrTokenDataNotFound) {
return false, fmt.Errorf("check authorization token: %w", err)
}
// Show polling status
@@ -340,7 +339,10 @@ func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocati
PrintPatAuthError(output, scopeErr)
// Wait for user to complete authorization
authorized := patWaitForAuthorization(ctx, configDir, output)
authorized, waitErr := patWaitForAuthorization(ctx, configDir, output)
if waitErr != nil {
return executor.Result{}, waitErr
}
if !authorized {
return executor.Result{}, apperrors.NewAuth(
"等待用户授权超时",
@@ -794,12 +796,6 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
pollURL := fmt.Sprintf("%s%s?flowId=%s",
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
// Load user access token for the poll request header.
var accessToken string
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
accessToken = tokenData.AccessToken
}
// Use a client that does NOT follow redirects, so we can detect SSO 302.
noRedirectClient := &http.Client{
CheckRedirect: func(req *http.Request, via []*http.Request) error {
@@ -828,6 +824,10 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
slog.Debug("PAT poll: failed to create request", "error", err)
continue
}
accessToken, tokenErr := patResolveAccessToken(ctx, configDir, "")
if tokenErr != nil && !stderrors.Is(tokenErr, authpkg.ErrTokenDataNotFound) {
return "", "", fmt.Errorf("resolve PAT poll access token: %w", tokenErr)
}
if accessToken != "" {
req.Header.Set("x-user-access-token", accessToken)
}
@@ -52,39 +52,41 @@ func TestCrossPlatformCoveragePATRetryRemainingPureAndWaitCoverage(t *testing.T)
oldTimeout := patAuthorizationTimeout
oldInterval := patAuthorizationPollInterval
oldLoad := patLoadTokenData
oldResolve := patResolveAccessToken
t.Cleanup(func() {
patAuthorizationTimeout = oldTimeout
patAuthorizationPollInterval = oldInterval
patLoadTokenData = oldLoad
patResolveAccessToken = oldResolve
})
patAuthorizationTimeout = 50 * time.Millisecond
patAuthorizationPollInterval = time.Millisecond
patLoadTokenData = func(string) (*authpkg.TokenData, error) {
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
patResolveAccessToken = func(context.Context, string, string) (string, error) {
return "token", nil
}
out.Reset()
if !WaitForPatAuthorization(context.Background(), "", &out) {
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || !ok {
t.Fatal("valid token did not authorize")
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
out.Reset()
if WaitForPatAuthorization(ctx, "", &out) {
t.Fatal("cancelled authorization succeeded")
if ok, err := WaitForPatAuthorization(ctx, "", &out); ok || !errors.Is(err, context.Canceled) {
t.Fatalf("cancelled authorization = %v, %v", ok, err)
}
patAuthorizationTimeout = time.Millisecond
patAuthorizationPollInterval = time.Hour
out.Reset()
if WaitForPatAuthorization(context.Background(), "", &out) {
t.Fatal("timed out authorization succeeded")
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok {
t.Fatalf("timed out authorization = %v, %v", ok, err)
}
patAuthorizationTimeout = 5 * time.Millisecond
patAuthorizationPollInterval = time.Millisecond
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
patResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
out.Reset()
if WaitForPatAuthorization(context.Background(), "", &out) || !strings.Contains(out.String(), "等待授权中") {
t.Fatalf("invalid-token polling output = %q", out.String())
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok || !strings.Contains(out.String(), "等待授权中") {
t.Fatalf("invalid-token polling = %v, %v, output %q", ok, err, out.String())
}
}
@@ -109,12 +111,12 @@ func TestCrossPlatformCoveragePATRetryRemainingOrchestrationCoverage(t *testing.
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
scope := &PatScopeError{OriginalError: "missing", Identity: "user", ErrorType: "missing_scope", Message: "missing", Hint: "login", MissingScope: "calendar:read"}
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return false }
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return false, nil }
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{}, executor.Invocation{}, scope, t.TempDir(), io.Discard); err == nil {
t.Fatal("PAT retry timeout succeeded")
}
wantErr := errors.New("runner failed")
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return true }
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return true, nil }
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{err: wantErr}, executor.Invocation{}, scope, t.TempDir(), io.Discard); !errors.Is(err, wantErr) {
t.Fatalf("authorized retry = %v", err)
}
@@ -215,15 +217,15 @@ func patRaw(flowID, clientID, secret string) string {
func TestCrossPlatformCoveragePATRetryRemainingPollAndBrowserCoverage(t *testing.T) {
oldDo := patPollHTTPDo
oldRequest := patPollNewRequest
oldLoad := patLoadTokenData
oldResolve := patResolveAccessToken
oldBrowser := patBrowserOpenCommand
t.Cleanup(func() {
patPollHTTPDo = oldDo
patPollNewRequest = oldRequest
patLoadTokenData = oldLoad
patResolveAccessToken = oldResolve
patBrowserOpenCommand = oldBrowser
})
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return &authpkg.TokenData{AccessToken: "token"}, nil }
patResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
cancelled, cancelNow := context.WithCancel(context.Background())
cancelNow()
if status, _, err := pollPatDeviceFlowWithInterval(cancelled, "flow", t.TempDir(), io.Discard, 0); err != nil || status != authpkg.StatusCancelled {
+5 -1
View File
@@ -333,7 +333,11 @@ func (r *recoveryRuntime) CallToolDirect(ctx context.Context, serverID, toolName
if err != nil {
return nil, err
}
tc := r.transport.WithAuth(resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags)), resolveIdentityHeaders())
authToken, err := resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags))
if err != nil {
return nil, tokenResolutionError(err)
}
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
result, err := tc.CallTool(ctx, endpoint, toolName, args)
if err != nil {
return nil, err
+64 -54
View File
@@ -235,7 +235,9 @@ func (r *runtimeRunner) runSingle(ctx context.Context, invocation executor.Invoc
// ~70ms on macOS; starting it here lets the load overlap with endpoint
// resolution and catalog loading below.
if prefetchToken {
go runnerGetCachedRuntimeToken(ctx)
go func() {
_, _ = runnerGetCachedRuntimeToken(ctx)
}()
}
if shouldUseDirectRuntime(invocation) {
@@ -534,8 +536,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
authToken := ""
if hasPluginAuth {
authToken = pluginAuth.Token
} else {
authToken = r.resolveAuthToken(ctx)
} else if !invocation.DryRun && (r.globalFlags == nil || !r.globalFlags.Mock) {
var tokenErr error
authToken, tokenErr = r.resolveAuthToken(ctx)
if tokenErr != nil {
return executor.Result{}, tokenResolutionError(tokenErr)
}
}
var timeoutSec int
@@ -617,6 +623,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
}
return runnerHandlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
}
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, err, hasPluginAuth); handled {
if retryErr != nil {
runnerCaptureRuntimeFailure(invocation, err, retryErr)
}
return result, retryErr
}
runnerCaptureRuntimeFailure(invocation, err, err)
return executor.Result{}, err
}
@@ -625,9 +637,15 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
callResult, err := runnerCallTool(tc, callCtx, endpoint, invocation.Tool, invocation.Params)
RecordTiming(ctx, "mcp_call", time.Since(callStart))
if err != nil {
if isAuthError(err) {
if isRefreshableTransportAuthError(err) {
if fn := edition.Get().OnAuthError; fn != nil {
if overrideErr := fn(defaultConfigDir(), err); overrideErr != nil {
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, overrideErr, hasPluginAuth); handled {
if retryErr != nil {
runnerCaptureRuntimeFailure(invocation, err, retryErr)
}
return result, retryErr
}
runnerCaptureRuntimeFailure(invocation, err, overrideErr)
return executor.Result{}, overrideErr
}
@@ -652,6 +670,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
}
return runnerHandlePatAuthCheck(ctx, r, invocation, patCheck, defaultConfigDir(), os.Stderr)
}
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, editionErr, hasPluginAuth); handled {
if retryErr != nil {
runnerCaptureRuntimeFailure(invocation, editionErr, retryErr)
}
return result, retryErr
}
return executor.Result{}, editionErr
}
}
@@ -672,6 +696,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
// patterns (PAT permission, gateway-auth) before generic handling.
if classify := edition.Get().ClassifyToolResult; classify != nil {
if hookErr := classify(callResult.Content); hookErr != nil {
if result, retryErr, handled := r.retryAuthRefreshRequired(ctx, endpoint, invocation, authToken, hookErr, hasPluginAuth); handled {
if retryErr != nil {
runnerCaptureRuntimeFailure(invocation, hookErr, retryErr)
}
return result, retryErr
}
runnerCaptureRuntimeFailure(invocation, hookErr, hookErr)
return executor.Result{}, hookErr
}
@@ -796,67 +826,49 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
}, nil
}
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) (string, error) {
explicitToken := ""
if r != nil && r.globalFlags != nil {
explicitToken = r.globalFlags.Token
}
if token := strings.TrimSpace(explicitToken); token != "" {
return token
}
if tp := edition.Get().TokenProvider; tp != nil {
token, _ := tp(ctx, func() (string, error) {
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
})
return token
}
return getCachedRuntimeToken(ctx)
return resolveRuntimeAuthToken(ctx, explicitToken)
}
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
if token := strings.TrimSpace(explicitToken); token != "" {
return token
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) (string, error) {
snapshot, err := runtimeTokenManager.Get(ctx, defaultConfigDir(), explicitToken)
if err != nil {
return "", err
}
// Use cached token to avoid repeated Keychain access (~70ms per call)
return getCachedRuntimeToken(ctx)
return snapshot.AccessToken, nil
}
// Cached token state for process lifetime
var (
cachedRuntimeTokenMu sync.Mutex
cachedRuntimeTokens = map[string]string{}
)
// getCachedRuntimeToken returns a cached access token, loading it only once per process.
// This avoids repeated Keychain access which takes ~70ms each time.
func getCachedRuntimeToken(ctx context.Context) string {
cacheKey := strings.TrimSpace(authpkg.RuntimeProfile())
if cacheKey == "" {
cacheKey = "__default__"
}
cachedRuntimeTokenMu.Lock()
if token := cachedRuntimeTokens[cacheKey]; token != "" {
cachedRuntimeTokenMu.Unlock()
return token
}
cachedRuntimeTokenMu.Unlock()
// getCachedRuntimeToken is kept as the prefetch seam used by runner tests. The
// cache itself lives exclusively in TokenManager.
func getCachedRuntimeToken(ctx context.Context) (string, error) {
loadStart := time.Now()
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
return resolveRuntimeAuthToken(ctx, "")
}
configDir := defaultConfigDir()
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
slog.Error(tokenErr.Error())
return ""
func tokenResolutionError(err error) error {
if err == nil {
return nil
}
if token == "" {
return ""
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return err
}
cachedRuntimeTokenMu.Lock()
cachedRuntimeTokens[cacheKey] = token
cachedRuntimeTokenMu.Unlock()
return token
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
return apperrors.NewAuth(
"未登录,请先执行 dws auth login",
apperrors.WithReason("not_authenticated"),
apperrors.WithHint("运行 'dws auth login' 完成登录后重试"),
apperrors.WithActions("dws auth login"),
apperrors.WithCause(err),
)
}
// Keychain, parse, permission, lock, and refresh failures are real local or
// network errors. Preserve their cause instead of disguising them as logout.
return fmt.Errorf("resolve access token: %w", err)
}
// generateExecutionID returns a random 16-char hex string used to correlate
@@ -871,9 +883,7 @@ func generateExecutionID() string {
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
// This should be called after login/logout operations.
func ResetRuntimeTokenCache() {
cachedRuntimeTokenMu.Lock()
defer cachedRuntimeTokenMu.Unlock()
cachedRuntimeTokens = map[string]string{}
runtimeTokenManager.Invalidate()
}
func newRuntimeContentScanner() safety.Scanner {
+9 -9
View File
@@ -50,9 +50,9 @@ func TestCrossPlatformCoverageRunnerRemainingRoutingCoverage(t *testing.T) {
inv := executor.Invocation{CanonicalProduct: "product", Tool: "tool"}
prefetched := make(chan struct{}, 1)
runnerGetCachedRuntimeToken = func(context.Context) string {
runnerGetCachedRuntimeToken = func(context.Context) (string, error) {
prefetched <- struct{}{}
return ""
return "", nil
}
r := &runtimeRunner{
loader: cli.CatalogLoaderFrom(cli.Catalog{}, wantErr),
@@ -197,7 +197,7 @@ func TestCrossPlatformCoverageRunnerRemainingExecutionCoverage(t *testing.T) {
return nil
}
authErr := apperrors.NewAuth("expired")
authErr := apperrors.NewAuth("expired", apperrors.WithReason("http_401"))
runnerCallTool = func(*transport.Client, context.Context, string, string, map[string]any) (transport.ToolCallResult, error) {
return transport.ToolCallResult{}, authErr
}
@@ -329,19 +329,19 @@ func TestCrossPlatformCoverageRunnerRemainingStdioAuthAndHeadersCoverage(t *test
}
r.globalFlags.Token = " explicit "
if got := r.resolveAuthToken(context.Background()); got != "explicit" {
t.Fatalf("explicit auth token = %q", got)
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "explicit" {
t.Fatalf("explicit auth token = %q, %v", got, err)
}
edition.Override(&edition.Hooks{TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
_, _ = fallback()
return "provided", nil
}})
r.globalFlags.Token = ""
if got := r.resolveAuthToken(context.Background()); got != "provided" {
t.Fatalf("provided auth token = %q", got)
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "provided" {
t.Fatalf("provided auth token = %q, %v", got, err)
}
if got := resolveRuntimeAuthToken(context.Background(), " runtime "); got != "runtime" {
t.Fatalf("runtime explicit token = %q", got)
if got, err := resolveRuntimeAuthToken(context.Background(), " runtime "); err != nil || got != "runtime" {
t.Fatalf("runtime explicit token = %q, %v", got, err)
}
t.Setenv(envDWSChannel, "channel")
+30 -26
View File
@@ -17,6 +17,7 @@ import (
"archive/zip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"mime"
@@ -36,25 +37,25 @@ import (
)
var (
skillLoadAccessToken = loadSkillAccessToken
skillDownloadToTmp = downloadSkillToTmpDir
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
skillNewRequest = http.NewRequestWithContext
skillLoadTokenData = authpkg.LoadTokenData
skillResolveTargetPath = resolveSkillTargetPath
skillFetchDownloadInfo = fetchSkillDownloadInfo
skillDownloadFile = downloadSkillFile
skillExtractZip = extractSkillZip
skillUserHomeDir = os.UserHomeDir
skillMkdirTemp = os.MkdirTemp
skillCreate = os.Create
skillCreateTemp = os.CreateTemp
skillRemoveAll = os.RemoveAll
skillRemove = os.Remove
skillMkdirAll = os.MkdirAll
skillOpenFile = os.OpenFile
skillCopy = io.Copy
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
skillLoadAccessToken = loadSkillAccessToken
skillDownloadToTmp = downloadSkillToTmpDir
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
skillNewRequest = http.NewRequestWithContext
skillResolveAccessToken = ResolveAuxiliaryAccessToken
skillResolveTargetPath = resolveSkillTargetPath
skillFetchDownloadInfo = fetchSkillDownloadInfo
skillDownloadFile = downloadSkillFile
skillExtractZip = extractSkillZip
skillUserHomeDir = os.UserHomeDir
skillMkdirTemp = os.MkdirTemp
skillCreate = os.Create
skillCreateTemp = os.CreateTemp
skillRemoveAll = os.RemoveAll
skillRemove = os.Remove
skillMkdirAll = os.MkdirAll
skillOpenFile = os.OpenFile
skillCopy = io.Copy
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
)
func init() {
@@ -296,7 +297,7 @@ func newSkillAddHintCommand() *cobra.Command {
func runSkillGet(cmd *cobra.Command, args []string) error {
skillID, _ := cmd.Flags().GetString("skill-id")
accessToken, err := skillLoadAccessToken()
accessToken, err := skillLoadAccessToken(cmd.Context())
if err != nil {
return err
}
@@ -319,7 +320,7 @@ func runSkillFind(cmd *cobra.Command, args []string) error {
if source == "" {
source, _ = cmd.Flags().GetString("scopes")
}
accessToken, err := skillLoadAccessToken()
accessToken, err := skillLoadAccessToken(cmd.Context())
if err != nil {
return err
}
@@ -388,7 +389,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
}
accessToken, err := skillLoadAccessToken()
accessToken, err := skillLoadAccessToken(cmd.Context())
if err != nil {
return err
}
@@ -441,13 +442,16 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
return nil
}
func loadSkillAccessToken() (string, error) {
func loadSkillAccessToken(ctx context.Context) (string, error) {
configDir := defaultConfigDir()
tokenData, err := skillLoadTokenData(configDir)
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
token, err := skillResolveAccessToken(ctx, configDir, "")
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
return "", skillAuthError()
}
return tokenData.AccessToken, nil
if err != nil {
return "", fmt.Errorf("resolve skill access token: %w", err)
}
return token, nil
}
func skillAuthError() error {
@@ -57,11 +57,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
})
fail := errors.New("failure")
cmd := skillCoverageCommand()
skillLoadAccessToken = func() (string, error) { return "", fail }
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
if err := runSkillGet(cmd, nil); !errors.Is(err, fail) {
t.Fatalf("skill get auth error = %v", err)
}
skillLoadAccessToken = func() (string, error) { return "token", nil }
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
if err := runSkillFind(cmd, nil); err == nil {
t.Fatal("skill find request failure should propagate")
@@ -72,11 +72,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
t.Fatalf("skill get download error = %v", err)
}
skillLoadAccessToken = func() (string, error) { return "", fail }
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
if err := runSkillFind(cmd, nil); !errors.Is(err, fail) {
t.Fatalf("skill find auth error = %v", err)
}
skillLoadAccessToken = func() (string, error) { return "token", nil }
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
skillHTTPDo = func(*http.Client, *http.Request) (*http.Response, error) { return nil, fail }
if err := runSkillFind(cmd, nil); err == nil {
t.Fatal("skill find network failure should propagate")
@@ -111,11 +111,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
t.Fatal("invalid skill target should fail")
}
skillResolveTargetPath = func(string) (string, error) { return "dest", nil }
skillLoadAccessToken = func() (string, error) { return "", fail }
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
t.Fatalf("skill add auth error = %v", err)
}
skillLoadAccessToken = func() (string, error) { return "token", nil }
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
skillFetchDownloadInfo = func(context.Context, string, string) (*downloadSkillResponse, error) { return nil, fail }
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
t.Fatalf("skill info error = %v", err)
@@ -152,25 +152,35 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
func TestCrossPlatformCoverageSkillCommandLowLevelRemainingCoverage(t *testing.T) {
oldHTTP := skillHTTPDo
oldNewRequest, oldLoadToken := skillNewRequest, skillLoadTokenData
oldNewRequest, oldResolveToken := skillNewRequest, skillResolveAccessToken
oldHome := skillUserHomeDir
oldMkdirTemp, oldCreate, oldCreateTemp := skillMkdirTemp, skillCreate, skillCreateTemp
oldRemoveAll, oldRemove, oldMkdir := skillRemoveAll, skillRemove, skillMkdirAll
oldOpen, oldCopy, oldZipOpen := skillOpenFile, skillCopy, skillOpenZipFile
t.Cleanup(func() {
skillHTTPDo = oldHTTP
skillNewRequest, skillLoadTokenData = oldNewRequest, oldLoadToken
skillNewRequest, skillResolveAccessToken = oldNewRequest, oldResolveToken
skillUserHomeDir = oldHome
skillMkdirTemp, skillCreate, skillCreateTemp = oldMkdirTemp, oldCreate, oldCreateTemp
skillRemoveAll, skillRemove, skillMkdirAll = oldRemoveAll, oldRemove, oldMkdir
skillOpenFile, skillCopy, skillOpenZipFile = oldOpen, oldCopy, oldZipOpen
})
fail := errors.New("failure")
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
if _, err := loadSkillAccessToken(); err == nil {
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
if _, err := loadSkillAccessToken(context.Background()); err == nil {
t.Fatal("invalid skill access token succeeded")
}
skillLoadTokenData = oldLoadToken
canceled, cancel := context.WithCancel(context.Background())
cancel()
skillResolveAccessToken = func(ctx context.Context, _, _ string) (string, error) {
return "", ctx.Err()
}
if _, err := loadSkillAccessToken(canceled); !errors.Is(err, context.Canceled) {
t.Fatalf("skill token cancellation = %v", err)
}
skillResolveAccessToken = oldResolveToken
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
if _, err := fetchSkillDownloadInfo(context.Background(), "token", "id"); err == nil {
t.Fatal("download-info request failure should propagate")
+9 -13
View File
@@ -18,7 +18,6 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
@@ -397,9 +396,11 @@ func TestSkillInstallRequiresAuth(t *testing.T) {
configDir := filepath.Join(tempDir, "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
t.Cleanup(CloseFileLogger)
originalLoadToken := skillLoadTokenData
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, errors.New("missing") }
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
originalResolveToken := skillResolveAccessToken
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
// Ensure the config directory exists but has no token
if err := os.MkdirAll(configDir, 0755); err != nil {
@@ -679,16 +680,11 @@ func TestSkillSearchUsesSourceQueryAndKeepsScopesCompat(t *testing.T) {
configDir := filepath.Join(t.TempDir(), "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
t.Cleanup(CloseFileLogger)
originalLoadToken := skillLoadTokenData
skillLoadTokenData = func(string) (*authpkg.TokenData, error) {
return &authpkg.TokenData{
AccessToken: "test-token",
RefreshToken: "refresh-token",
ExpiresAt: time.Now().Add(time.Hour),
RefreshExpAt: time.Now().Add(24 * time.Hour),
}, nil
originalResolveToken := skillResolveAccessToken
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "test-token", nil
}
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
var gotSources []string
var gotScopes []string
+10
View File
@@ -52,6 +52,10 @@ func TestMain(m *testing.M) {
_ = os.RemoveAll(tmpDir)
panic("set " + keychain.StorageDirEnv + ": " + err.Error())
}
if err := os.Setenv(keychain.TestNamespaceEnv, tmpDir); err != nil {
_ = os.RemoveAll(tmpDir)
panic("set " + keychain.TestNamespaceEnv + ": " + err.Error())
}
if err := os.Setenv("DWS_CONFIG_DIR", filepath.Join(tmpDir, "config")); err != nil {
_ = os.RemoveAll(tmpDir)
panic("set DWS_CONFIG_DIR: " + err.Error())
@@ -79,6 +83,12 @@ func TestMain(m *testing.M) {
StopAllStdioClients()
CloseAuditSink()
CloseFileLogger()
if err := keychain.RemoveAuthTokenEntries(keychain.Service); err != nil {
fmt.Fprintf(os.Stderr, "internal/app keychain cleanup: %v\n", err)
if code == 0 {
code = 1
}
}
if err := os.RemoveAll(tmpDir); err != nil {
fmt.Fprintf(os.Stderr, "internal/app test cleanup %s: %v\n", tmpDir, err)
if code == 0 {
@@ -0,0 +1,70 @@
package auth
import (
"context"
"errors"
"testing"
"time"
)
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotPreservesLoadFailure(t *testing.T) {
oldLoad := oauthLoadToken
want := errors.New("keychain permission denied")
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
t.Cleanup(func() { oauthLoadToken = oldLoad })
_, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
if !errors.Is(err, want) {
t.Fatalf("error = %v, want cause %v", err, want)
}
if errors.Is(err, ErrTokenDataNotFound) {
t.Fatalf("load failure was misclassified as missing credentials: %v", err)
}
}
func TestCrossPlatformCoverageOAuthProviderLoginPreservesLoadFailure(t *testing.T) {
oldLoad := oauthLoadToken
want := errors.New("keychain permission denied")
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
t.Cleanup(func() { oauthLoadToken = oldLoad })
_, err := NewOAuthProvider(t.TempDir(), nil).Login(context.Background(), false)
if !errors.Is(err, want) {
t.Fatalf("error = %v, want cause %v", err, want)
}
}
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotReturnsExpiryMetadata(t *testing.T) {
oldLoad := oauthLoadToken
expiresAt := time.Now().Add(time.Hour)
oauthLoadToken = func(string) (*TokenData, error) {
return &TokenData{AccessToken: "token", ExpiresAt: expiresAt}, nil
}
t.Cleanup(func() { oauthLoadToken = oldLoad })
snapshot, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
if err != nil {
t.Fatal(err)
}
if snapshot.AccessToken != "token" || !snapshot.ExpiresAt.Equal(expiresAt) {
t.Fatalf("snapshot = %#v", snapshot)
}
}
func TestCrossPlatformCoverageTokenMarkerRevisionChangesOnEveryPublication(t *testing.T) {
configDir := t.TempDir()
if err := WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
first, present, err := ReadTokenMarkerRevision(configDir)
if err != nil || !present || first == "" {
t.Fatalf("first marker = %q, %v, %v", first, present, err)
}
if err := WriteTokenMarker(configDir); err != nil {
t.Fatal(err)
}
second, present, err := ReadTokenMarkerRevision(configDir)
if err != nil || !present || second == "" || second == first {
t.Fatalf("second marker = %q, %v, %v; first=%q", second, present, err, first)
}
}
+233 -1
View File
@@ -13,7 +13,50 @@
package auth
import "time"
import (
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"path/filepath"
"strings"
"sync"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
const rejectedTokenRefreshFailureCooldown = 3 * time.Second
type rejectedTokenRefreshKey struct {
configDir string
profile string
tokenDigest [sha256.Size]byte
}
type rejectedTokenRefreshCall struct {
done chan struct{}
participants int
token string
err error
}
type rejectedTokenRefreshFailure struct {
at time.Time
err error
}
var rejectedTokenRefreshCoordinator = struct {
sync.Mutex
inFlight map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall
failures map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure
now func() time.Time
}{
inFlight: make(map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall),
failures: make(map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure),
now: time.Now,
}
// MarkAccessTokenStale loads the persisted TokenData, sets ExpiresAt to a past
// instant (preserving access_token and refresh_token), and writes it back. The
@@ -39,3 +82,192 @@ func MarkAccessTokenStale(configDir string) error {
data.ExpiresAt = time.Now().Add(-1 * time.Minute)
return SaveTokenData(configDir, data)
}
// ForceRefreshRejectedToken refreshes rejectedAccessToken only while it is
// still the credential stored for the active profile. The compare and refresh
// run under the same process + file lock used by ordinary expiry refresh, so a
// late rejection cannot invalidate or refresh over a token another caller has
// already rotated.
//
// When the stored token no longer matches, the newer token is returned without
// calling the refresh endpoint. Refresh failures leave the stored credential in
// place; login/logout remain the only owners of credential deletion.
func (p *OAuthProvider) ForceRefreshRejectedToken(ctx context.Context, rejectedAccessToken string) (string, error) {
if p == nil || strings.TrimSpace(p.configDir) == "" {
return "", fmt.Errorf("config directory is empty")
}
rejectedAccessToken = strings.TrimSpace(rejectedAccessToken)
if rejectedAccessToken == "" {
return "", fmt.Errorf("rejected access token is empty")
}
if ctx == nil {
ctx = context.Background()
}
profile := strings.TrimSpace(RuntimeProfile())
key := newRejectedTokenRefreshKey(p.configDir, profile, rejectedAccessToken)
call, leader := beginRejectedTokenRefresh(key)
if !leader {
select {
case <-call.done:
return call.token, call.err
case <-ctx.Done():
return "", ctx.Err()
}
}
token, err, recordFailure := p.forceRefreshRejectedTokenOnce(ctx, profile, rejectedAccessToken, key)
finishRejectedTokenRefresh(key, call, token, err, recordFailure)
return token, err
}
func (p *OAuthProvider) forceRefreshRejectedTokenOnce(
ctx context.Context,
profile string,
rejectedAccessToken string,
key rejectedTokenRefreshKey,
) (string, error, bool) {
lock, err := oauthAcquireLock(ctx, p.configDir)
if err != nil {
return "", fmt.Errorf("acquiring dual lock: %w", err), false
}
defer lock.Release()
data, err := loadOAuthTokenUnderHeldLock(p.configDir, profile)
if err != nil {
return "", fmt.Errorf("reload rejected token: %w", err), false
}
current := strings.TrimSpace(data.AccessToken)
if current == "" {
clearRejectedTokenRefreshFailure(key)
return "", fmt.Errorf("stored access token is empty"), false
}
if current != rejectedAccessToken {
clearRejectedTokenRefreshFailure(key)
return current, nil, false
}
if cachedErr := recentRejectedTokenRefreshFailure(key); cachedErr != nil {
return "", cachedErr, false
}
if !data.IsRefreshTokenValid() {
return "", fmt.Errorf("refresh_token 已过期"), true
}
if err := preflightTokenRefreshPersistence(p.configDir, data); err != nil {
return "", fmt.Errorf("本地登录态无法安全更新: %w", err), true
}
refreshed, err := oauthRefreshToken(p, ctx, data)
if err != nil {
return "", err, true
}
if refreshed == nil || strings.TrimSpace(refreshed.AccessToken) == "" {
return "", fmt.Errorf("force refresh returned empty access token"), true
}
return strings.TrimSpace(refreshed.AccessToken), nil, false
}
func newRejectedTokenRefreshKey(configDir, profile, rejectedAccessToken string) rejectedTokenRefreshKey {
canonicalDir := filepath.Clean(configDir)
if absolute, err := filepath.Abs(configDir); err == nil {
canonicalDir = filepath.Clean(absolute)
}
return rejectedTokenRefreshKey{
configDir: canonicalDir,
profile: strings.TrimSpace(profile),
tokenDigest: sha256.Sum256([]byte(strings.TrimSpace(rejectedAccessToken))),
}
}
func beginRejectedTokenRefresh(key rejectedTokenRefreshKey) (*rejectedTokenRefreshCall, bool) {
rejectedTokenRefreshCoordinator.Lock()
defer rejectedTokenRefreshCoordinator.Unlock()
if call := rejectedTokenRefreshCoordinator.inFlight[key]; call != nil {
call.participants++
return call, false
}
call := &rejectedTokenRefreshCall{done: make(chan struct{}), participants: 1}
rejectedTokenRefreshCoordinator.inFlight[key] = call
return call, true
}
func finishRejectedTokenRefresh(
key rejectedTokenRefreshKey,
call *rejectedTokenRefreshCall,
token string,
err error,
recordFailure bool,
) {
rejectedTokenRefreshCoordinator.Lock()
defer rejectedTokenRefreshCoordinator.Unlock()
call.token = token
call.err = err
delete(rejectedTokenRefreshCoordinator.inFlight, key)
if err == nil {
delete(rejectedTokenRefreshCoordinator.failures, key)
} else if recordFailure && !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
now := time.Now()
if rejectedTokenRefreshCoordinator.now != nil {
now = rejectedTokenRefreshCoordinator.now()
}
rejectedTokenRefreshCoordinator.failures[key] = rejectedTokenRefreshFailure{at: now, err: err}
}
close(call.done)
}
func recentRejectedTokenRefreshFailure(key rejectedTokenRefreshKey) error {
rejectedTokenRefreshCoordinator.Lock()
defer rejectedTokenRefreshCoordinator.Unlock()
now := time.Now()
if rejectedTokenRefreshCoordinator.now != nil {
now = rejectedTokenRefreshCoordinator.now()
}
for failureKey, failure := range rejectedTokenRefreshCoordinator.failures {
age := now.Sub(failure.at)
if age < 0 || age >= rejectedTokenRefreshFailureCooldown {
delete(rejectedTokenRefreshCoordinator.failures, failureKey)
}
}
if failure, ok := rejectedTokenRefreshCoordinator.failures[key]; ok {
return failure.err
}
return nil
}
func clearRejectedTokenRefreshFailure(key rejectedTokenRefreshKey) {
rejectedTokenRefreshCoordinator.Lock()
delete(rejectedTokenRefreshCoordinator.failures, key)
rejectedTokenRefreshCoordinator.Unlock()
}
// loadOAuthTokenUnderHeldLock mirrors LoadTokenDataForProfile without taking a
// second, non-reentrant auth lock. Opaque edition storage hooks (for example
// Wukong's encrypted .data file) are read inside the caller's dual lock so the
// compare-and-refresh decision covers both Core and embedded storage.
func loadOAuthTokenUnderHeldLock(configDir, profile string) (*TokenData, error) {
hooks := edition.Get()
if hooks.LoadToken == nil {
data, err := oauthLoadTokenLocked(configDir, profile)
if err != nil {
return nil, err
}
if data == nil {
return nil, fmt.Errorf("stored token data is empty")
}
return data, nil
}
if strings.TrimSpace(profile) != "" {
return nil, fmt.Errorf("profile selection is not supported by the current auth backend")
}
blob, err := hooks.LoadToken(configDir)
if err != nil {
return nil, err
}
var data TokenData
if err := json.Unmarshal(blob, &data); err != nil {
return nil, fmt.Errorf("parsing token data from hook: %w", err)
}
return &data, nil
}
+4 -2
View File
@@ -47,8 +47,10 @@ func (m *Manager) GetToken() (string, string, error) {
}
return token, "file", nil
}
return "", "", fmt.Errorf("%s", i18n.T("未找到认证信息,请运行 dws auth login"))
if err != nil && !os.IsNotExist(err) {
return "", "", fmt.Errorf("load legacy token: %w", err)
}
return "", "", fmt.Errorf("%s: %w", i18n.T("未找到认证信息,请运行 dws auth login"), ErrTokenDataNotFound)
}
func (m *Manager) GetMCPURL() (string, error) {
+30 -10
View File
@@ -115,6 +115,12 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
// Smart degradation: try silent refresh before opening browser.
if !force {
data, err := oauthLoadToken(p.configDir)
if err != nil && !errors.Is(err, ErrTokenDataNotFound) && !os.IsNotExist(err) {
if preflightErr := preflightTokenPersistence(p.configDir); preflightErr != nil {
return nil, fmt.Errorf("%s: %w", i18n.T("本地登录态无法安全更新"), preflightErr)
}
return nil, fmt.Errorf("load existing access token: %w", err)
}
if err == nil {
// Case 1: access_token still valid — no action needed.
if data.IsAccessTokenValid() {
@@ -636,36 +642,50 @@ continueLogin:
return tokenData, nil
}
// GetAccessToken returns a valid access token, auto-refreshing if needed.
// Uses a file lock with double-check pattern to prevent concurrent refresh
// from multiple CLI processes.
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
// GetTokenSnapshot returns a valid token together with its expiry metadata.
// Storage and refresh failures retain their original cause; only a confirmed
// missing credential is reported as ErrTokenDataNotFound.
func (p *OAuthProvider) GetTokenSnapshot(ctx context.Context) (*TokenData, error) {
data, err := oauthLoadToken(p.configDir)
if err != nil {
return "", errors.New(i18n.T("未登录,请运行 dws auth login"))
if errors.Is(err, ErrTokenDataNotFound) || os.IsNotExist(err) {
return nil, fmt.Errorf("%s: %w", i18n.T("未登录,请运行 dws auth login"), ErrTokenDataNotFound)
}
return nil, fmt.Errorf("load access token: %w", err)
}
// Fast path: access_token still valid — no lock needed.
if data.IsAccessTokenValid() {
return data.AccessToken, nil
return data, nil
}
// Slow path: token expired — try locked refresh.
if data.IsRefreshTokenValid() {
refreshed, rErr := p.lockedRefresh(ctx)
if rErr == nil {
return refreshed.AccessToken, nil
return refreshed, nil
}
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
if p.logger != nil {
p.logger.Warn(i18n.T("refresh_token 刷新失败"), "error", rErr)
}
return "", fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
return nil, fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
} else {
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
}
return "", errors.New(i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"))
return nil, fmt.Errorf("%s: %w", i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"), ErrTokenDataNotFound)
}
// GetAccessToken returns a valid access token, auto-refreshing if needed.
// Uses a file lock with double-check pattern to prevent concurrent refresh
// from multiple CLI processes.
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
data, err := p.GetTokenSnapshot(ctx)
if err != nil {
return "", err
}
return strings.TrimSpace(data.AccessToken), nil
}
// lockedRefresh attempts to refresh the token while holding dual-layer locks.
@@ -697,7 +717,7 @@ func (p *OAuthProvider) lockedRefresh(ctx context.Context) (*TokenData, error) {
// Double-check: re-load from disk — another goroutine/process may have refreshed
// while we were waiting for the lock.
data, err := oauthLoadTokenLocked(p.configDir, RuntimeProfile())
data, err := loadOAuthTokenUnderHeldLock(p.configDir, RuntimeProfile())
if err != nil {
return nil, err
}
@@ -0,0 +1,510 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"context"
"encoding/json"
"errors"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
)
type rejectedTokenHookStore struct {
mu sync.Mutex
data TokenData
deletes int
}
func (s *rejectedTokenHookStore) load(string) ([]byte, error) {
s.mu.Lock()
defer s.mu.Unlock()
return json.Marshal(s.data)
}
func (s *rejectedTokenHookStore) save(_ string, blob []byte) error {
var data TokenData
if err := json.Unmarshal(blob, &data); err != nil {
return err
}
s.mu.Lock()
s.data = data
s.mu.Unlock()
return nil
}
func (s *rejectedTokenHookStore) delete(string) error {
s.mu.Lock()
s.data = TokenData{}
s.deletes++
s.mu.Unlock()
return nil
}
func (s *rejectedTokenHookStore) snapshot() (TokenData, int) {
s.mu.Lock()
defer s.mu.Unlock()
return s.data, s.deletes
}
func installRejectedTokenHookStore(t *testing.T, data TokenData) *rejectedTokenHookStore {
t.Helper()
store := &rejectedTokenHookStore{data: data}
previousHooks := edition.Get()
edition.Override(&edition.Hooks{
LoadToken: store.load,
SaveToken: store.save,
DeleteToken: store.delete,
})
t.Cleanup(func() { edition.Override(previousHooks) })
return store
}
func installOAuthRefreshStub(t *testing.T, fn func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error)) {
t.Helper()
resetRejectedTokenRefreshCoordinator(t)
previous := oauthRefreshToken
oauthRefreshToken = fn
t.Cleanup(func() { oauthRefreshToken = previous })
}
func resetRejectedTokenRefreshCoordinator(t *testing.T) {
t.Helper()
reset := func() {
rejectedTokenRefreshCoordinator.Lock()
rejectedTokenRefreshCoordinator.inFlight = make(map[rejectedTokenRefreshKey]*rejectedTokenRefreshCall)
rejectedTokenRefreshCoordinator.failures = make(map[rejectedTokenRefreshKey]rejectedTokenRefreshFailure)
rejectedTokenRefreshCoordinator.now = time.Now
rejectedTokenRefreshCoordinator.Unlock()
}
reset()
t.Cleanup(reset)
}
func waitForRejectedTokenRefreshParticipants(t *testing.T, key rejectedTokenRefreshKey, want int) {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for {
rejectedTokenRefreshCoordinator.Lock()
call := rejectedTokenRefreshCoordinator.inFlight[key]
got := 0
if call != nil {
got = call.participants
}
rejectedTokenRefreshCoordinator.Unlock()
if got >= want {
return
}
if time.Now().After(deadline) {
t.Fatalf("refresh participants = %d, want %d", got, want)
}
time.Sleep(time.Millisecond)
}
}
func installProfilesAcquireProbe(t *testing.T) <-chan struct{} {
t.Helper()
previous := profilesAcquireDualLock
attempted := make(chan struct{}, 1)
profilesAcquireDualLock = func(ctx context.Context, configDir string) (*DualLock, error) {
attempted <- struct{}{}
return previous(ctx, configDir)
}
t.Cleanup(func() { profilesAcquireDualLock = previous })
return attempted
}
func waitForProfilesAcquire(t *testing.T, attempted <-chan struct{}) {
t.Helper()
select {
case <-attempted:
case <-time.After(2 * time.Second):
t.Fatal("public opaque token mutation did not enter the Core dual lock")
}
}
func validRejectedTokenData(accessToken string) TokenData {
return TokenData{
AccessToken: accessToken,
RefreshToken: "refresh-token",
ExpiresAt: time.Now().Add(time.Hour),
RefreshExpAt: time.Now().Add(24 * time.Hour),
Source: "mcp",
ClientID: "client-id",
}
}
func TestCrossPlatformCoverageForceRefreshRejectedTokenConcurrentCallersExchangeOnce(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
var refreshCalls atomic.Int32
started := make(chan struct{})
release := make(chan struct{})
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
if refreshCalls.Add(1) == 1 {
close(started)
}
<-release
updated := *data
updated.AccessToken = "new-access"
updated.ExpiresAt = time.Now().Add(time.Hour)
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
return nil, err
}
return &updated, nil
})
provider := NewOAuthProvider(t.TempDir(), nil)
const workers = 8
results := make(chan string, workers)
errs := make(chan error, workers)
var wg sync.WaitGroup
wg.Add(workers)
for range workers {
go func() {
defer wg.Done()
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
results <- token
errs <- err
}()
}
<-started
close(release)
wg.Wait()
close(results)
close(errs)
for err := range errs {
if err != nil {
t.Fatal(err)
}
}
for token := range results {
if token != "new-access" {
t.Fatalf("token = %q, want new-access", token)
}
}
if got := refreshCalls.Load(); got != 1 {
t.Fatalf("refresh calls = %d, want 1", got)
}
stored, deletes := store.snapshot()
if stored.AccessToken != "new-access" || deletes != 0 {
t.Fatalf("stored token = %q, deletes = %d", stored.AccessToken, deletes)
}
}
func TestCrossPlatformCoverageForceRefreshRejectedTokenFailurePreservesCredential(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
refreshErr := errors.New("temporary refresh failure")
installOAuthRefreshStub(t, func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
return nil, refreshErr
})
_, err := NewOAuthProvider(t.TempDir(), nil).ForceRefreshRejectedToken(context.Background(), "old-access")
if !errors.Is(err, refreshErr) {
t.Fatalf("error = %v, want refresh cause", err)
}
stored, deletes := store.snapshot()
if stored.AccessToken != "old-access" || stored.RefreshToken != "refresh-token" || deletes != 0 {
t.Fatalf("credential changed after transient failure: %#v, deletes=%d", stored, deletes)
}
}
func TestCrossPlatformCoverageForceRefreshRejectedTokenFailureIsSingleflightAndCooledDown(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
previousProfile := RuntimeProfile()
SetRuntimeProfile("")
t.Cleanup(func() { SetRuntimeProfile(previousProfile) })
refreshErr := errors.New("temporary refresh failure")
var refreshCalls atomic.Int32
started := make(chan struct{})
release := make(chan struct{})
baseNow := time.Now()
var nowNanos atomic.Int64
nowNanos.Store(baseNow.UnixNano())
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
call := refreshCalls.Add(1)
if call == 1 {
close(started)
<-release
return nil, refreshErr
}
updated := *data
updated.AccessToken = "recovered-access"
updated.ExpiresAt = time.Now().Add(time.Hour)
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
return nil, err
}
return &updated, nil
})
rejectedTokenRefreshCoordinator.Lock()
rejectedTokenRefreshCoordinator.now = func() time.Time {
return time.Unix(0, nowNanos.Load())
}
rejectedTokenRefreshCoordinator.Unlock()
configDir := t.TempDir()
provider := NewOAuthProvider(configDir, nil)
const workers = 8
start := make(chan struct{})
ready := make(chan struct{}, workers)
errs := make(chan error, workers)
var wg sync.WaitGroup
wg.Add(workers)
for range workers {
go func() {
defer wg.Done()
ready <- struct{}{}
<-start
_, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
errs <- err
}()
}
for range workers {
<-ready
}
close(start)
<-started
key := newRejectedTokenRefreshKey(configDir, "", "old-access")
waitForRejectedTokenRefreshParticipants(t, key, workers)
close(release)
wg.Wait()
close(errs)
for err := range errs {
if !errors.Is(err, refreshErr) {
t.Fatalf("shared refresh error = %v, want %v", err, refreshErr)
}
}
if got := refreshCalls.Load(); got != 1 {
t.Fatalf("refresh calls after concurrent failure = %d, want 1", got)
}
stored, deletes := store.snapshot()
if stored.AccessToken != "old-access" || stored.RefreshToken != "refresh-token" || deletes != 0 {
t.Fatalf("credential changed after shared failure: %#v, deletes=%d", stored, deletes)
}
if _, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access"); !errors.Is(err, refreshErr) {
t.Fatalf("cooldown error = %v, want %v", err, refreshErr)
}
if got := refreshCalls.Load(); got != 1 {
t.Fatalf("refresh calls inside cooldown = %d, want 1", got)
}
nowNanos.Store(baseNow.Add(rejectedTokenRefreshFailureCooldown + time.Nanosecond).UnixNano())
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
if err != nil || token != "recovered-access" {
t.Fatalf("refresh after cooldown = %q, %v", token, err)
}
if got := refreshCalls.Load(); got != 2 {
t.Fatalf("refresh calls after cooldown = %d, want 2", got)
}
stored, deletes = store.snapshot()
if stored.AccessToken != "recovered-access" || deletes != 0 {
t.Fatalf("stored token after cooldown recovery = %q, deletes=%d", stored.AccessToken, deletes)
}
}
func TestCrossPlatformCoverageForceRefreshRejectedTokenChangedDuringCooldownUsesNewToken(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
previousProfile := RuntimeProfile()
SetRuntimeProfile("")
t.Cleanup(func() { SetRuntimeProfile(previousProfile) })
refreshErr := errors.New("temporary refresh failure")
var refreshCalls atomic.Int32
baseNow := time.Now()
installOAuthRefreshStub(t, func(*OAuthProvider, context.Context, *TokenData) (*TokenData, error) {
refreshCalls.Add(1)
return nil, refreshErr
})
rejectedTokenRefreshCoordinator.Lock()
rejectedTokenRefreshCoordinator.now = func() time.Time { return baseNow }
rejectedTokenRefreshCoordinator.Unlock()
configDir := t.TempDir()
provider := NewOAuthProvider(configDir, nil)
if _, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access"); !errors.Is(err, refreshErr) {
t.Fatalf("initial refresh error = %v, want %v", err, refreshErr)
}
if got := refreshCalls.Load(); got != 1 {
t.Fatalf("initial refresh calls = %d, want 1", got)
}
store.mu.Lock()
store.data = validRejectedTokenData("externally-refreshed")
store.mu.Unlock()
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
if err != nil || token != "externally-refreshed" {
t.Fatalf("refresh after external publication = %q, %v", token, err)
}
if got := refreshCalls.Load(); got != 1 {
t.Fatalf("external publication triggered another exchange: calls=%d", got)
}
key := newRejectedTokenRefreshKey(configDir, "", "old-access")
rejectedTokenRefreshCoordinator.Lock()
_, failurePresent := rejectedTokenRefreshCoordinator.failures[key]
rejectedTokenRefreshCoordinator.Unlock()
if failurePresent {
t.Fatal("old-token failure cache was not cleared after external publication")
}
}
func TestCrossPlatformCoverageOpaquePublisherWaitsForRejectedTokenRefresh(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
started := make(chan struct{})
release := make(chan struct{})
var releaseOnce sync.Once
t.Cleanup(func() { releaseOnce.Do(func() { close(release) }) })
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
close(started)
<-release
updated := *data
updated.AccessToken = "refreshed-from-old"
updated.ExpiresAt = time.Now().Add(time.Hour)
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
return nil, err
}
return &updated, nil
})
configDir := t.TempDir()
provider := NewOAuthProvider(configDir, nil)
refreshResult := make(chan struct {
token string
err error
}, 1)
go func() {
token, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
refreshResult <- struct {
token string
err error
}{token: token, err: err}
}()
<-started
acquireAttempted := installProfilesAcquireProbe(t)
publishResult := make(chan error, 1)
go func() {
publishResult <- SaveTokenData(configDir, ptrTokenData(validRejectedTokenData("login-published")))
}()
waitForProfilesAcquire(t, acquireAttempted)
releaseOnce.Do(func() { close(release) })
refresh := <-refreshResult
if refresh.err != nil || refresh.token != "refreshed-from-old" {
t.Fatalf("refresh result = %q, %v", refresh.token, refresh.err)
}
if err := <-publishResult; err != nil {
t.Fatalf("publish token: %v", err)
}
stored, deletes := store.snapshot()
if stored.AccessToken != "login-published" || deletes != 0 {
t.Fatalf("older refresh overwrote login publication: token=%q deletes=%d", stored.AccessToken, deletes)
}
}
func TestCrossPlatformCoverageOpaqueLogoutWaitsForRejectedTokenRefresh(t *testing.T) {
for _, tc := range []struct {
name string
logout func(string) error
}{
{name: "current profile", logout: func(configDir string) error {
return DeleteTokenDataForProfile(configDir, "")
}},
{name: "all profiles", logout: DeleteAllTokenData},
} {
t.Run(tc.name, func(t *testing.T) {
store := installRejectedTokenHookStore(t, validRejectedTokenData("old-access"))
started := make(chan struct{})
release := make(chan struct{})
var releaseOnce sync.Once
t.Cleanup(func() { releaseOnce.Do(func() { close(release) }) })
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, data *TokenData) (*TokenData, error) {
close(started)
<-release
updated := *data
updated.AccessToken = "refreshed-before-logout"
updated.ExpiresAt = time.Now().Add(time.Hour)
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
return nil, err
}
return &updated, nil
})
configDir := t.TempDir()
provider := NewOAuthProvider(configDir, nil)
refreshResult := make(chan error, 1)
go func() {
_, err := provider.ForceRefreshRejectedToken(context.Background(), "old-access")
refreshResult <- err
}()
<-started
acquireAttempted := installProfilesAcquireProbe(t)
logoutResult := make(chan error, 1)
go func() { logoutResult <- tc.logout(configDir) }()
waitForProfilesAcquire(t, acquireAttempted)
releaseOnce.Do(func() { close(release) })
if err := <-refreshResult; err != nil {
t.Fatalf("refresh: %v", err)
}
if err := <-logoutResult; err != nil {
t.Fatalf("logout: %v", err)
}
stored, deletes := store.snapshot()
if stored.AccessToken != "" || stored.RefreshToken != "" || deletes != 1 {
t.Fatalf("refresh resurrected logged-out credential: %#v deletes=%d", stored, deletes)
}
})
}
}
func ptrTokenData(data TokenData) *TokenData {
return &data
}
func TestCrossPlatformCoverageOAuthLockedRefreshReadsOpaqueEditionStore(t *testing.T) {
data := validRejectedTokenData("expired-access")
data.ExpiresAt = time.Now().Add(-time.Hour)
store := installRejectedTokenHookStore(t, data)
var refreshCalls atomic.Int32
installOAuthRefreshStub(t, func(p *OAuthProvider, _ context.Context, current *TokenData) (*TokenData, error) {
refreshCalls.Add(1)
updated := *current
updated.AccessToken = "proactively-refreshed"
updated.ExpiresAt = time.Now().Add(time.Hour)
if err := saveTokenDataLocked(p.configDir, &updated); err != nil {
return nil, err
}
return &updated, nil
})
token, err := NewOAuthProvider(t.TempDir(), nil).GetAccessToken(context.Background())
if err != nil || token != "proactively-refreshed" {
t.Fatalf("GetAccessToken() = %q, %v", token, err)
}
if refreshCalls.Load() != 1 {
t.Fatalf("refresh calls = %d, want 1", refreshCalls.Load())
}
stored, deletes := store.snapshot()
if stored.AccessToken != "proactively-refreshed" || deletes != 0 {
t.Fatalf("stored token = %q, deletes = %d", stored.AccessToken, deletes)
}
}
+11
View File
@@ -14,6 +14,7 @@
package auth
import (
"fmt"
"os"
"testing"
@@ -34,7 +35,17 @@ func TestMain(m *testing.M) {
_ = os.RemoveAll(tmpDir)
panic("set " + keychain.StorageDirEnv + ": " + err.Error())
}
if err := os.Setenv(keychain.TestNamespaceEnv, tmpDir); err != nil {
_ = os.RemoveAll(tmpDir)
panic("set " + keychain.TestNamespaceEnv + ": " + err.Error())
}
code := m.Run()
if err := keychain.RemoveAuthTokenEntries(keychain.Service); err != nil {
fmt.Fprintf(os.Stderr, "internal/auth keychain cleanup: %v\n", err)
if code == 0 {
code = 1
}
}
_ = os.RemoveAll(tmpDir)
os.Exit(code)
}
+37 -11
View File
@@ -127,6 +127,10 @@ const tokenJSONFile = "token.json"
type TokenMarker struct {
UpdatedAt string `json:"updated_at"`
ManualToken bool `json:"manual_token,omitempty"`
// Revision changes on every credential publication. Runtime token caches
// use it as a cheap cross-process invalidation signal without reading the
// platform keychain on every request.
Revision string `json:"revision,omitempty"`
}
// WriteTokenMarker writes a token.json marker containing only an updated_at
@@ -147,6 +151,7 @@ func writeTokenMarker(configDir string, manual bool) error {
marker := TokenMarker{
UpdatedAt: time.Now().Format(time.RFC3339),
ManualToken: manual,
Revision: uuid.NewString(),
}
data, _ := tokenJSONMarshalIndent(marker, "", " ")
if err := tokenMkdirAll(configDir, 0o700); err != nil {
@@ -159,6 +164,27 @@ func writeTokenMarker(configDir string, manual bool) error {
return tokenRename(tmp, filepath.Join(configDir, tokenJSONFile))
}
// ReadTokenMarkerRevision returns the current credential publication revision.
// Existing markers without a revision remain readable, but callers must avoid
// caching them because they cannot prove that the credential is unchanged.
func ReadTokenMarkerRevision(configDir string) (revision string, present bool, err error) {
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
if err != nil {
if os.IsNotExist(err) {
return "", false, nil
}
return "", false, fmt.Errorf("read token marker: %w", err)
}
var marker TokenMarker
if err := json.Unmarshal(data, &marker); err != nil {
// The marker is only a cache-coherency hint. A malformed historical or
// externally modified marker must disable caching, not make an otherwise
// valid credential unusable.
return "", true, nil
}
return strings.TrimSpace(marker.Revision), true, nil
}
func manualTokenMarkerActive(configDir string) (bool, error) {
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
if err != nil {
@@ -184,13 +210,10 @@ func DeleteTokenMarker(configDir string) error {
return nil
}
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to the default keychain-based storage.
// SaveTokenData persists TokenData under the auth dual lock. When an edition
// hook (SaveToken) is registered, the locked write delegates to that hook;
// otherwise it falls back to the default keychain-based storage.
func SaveTokenData(configDir string, data *TokenData) error {
if h := edition.Get(); h.SaveToken != nil {
return saveTokenViaHook(h, configDir, data)
}
return withProfilesLock(configDir, func() error {
return saveTokenDataLocked(configDir, data)
})
@@ -464,9 +487,8 @@ func tokenLoadProfileIdentity(profile Profile) (*TokenData, error) {
return orgData, nil
}
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to keychain + legacy cleanup.
// DeleteTokenData removes token data. Edition hooks and the default keychain
// path are both serialized with refresh through the auth dual lock.
func DeleteTokenData(configDir string) error {
return DeleteTokenDataForProfile(configDir, RuntimeProfile())
}
@@ -478,7 +500,9 @@ func DeleteTokenDataForProfile(configDir, profile string) error {
if strings.TrimSpace(profile) != "" {
return fmt.Errorf("profile selection is not supported by the current auth backend")
}
return h.DeleteToken(configDir)
return withProfilesLock(configDir, func() error {
return h.DeleteToken(configDir)
})
}
return withProfilesLock(configDir, func() error {
return deleteTokenDataForProfileLocked(configDir, profile)
@@ -913,7 +937,9 @@ func restoreTokenMarker(configDir string, marker tokenMarkerSnapshot) error {
// DeleteAllTokenData removes all profile-scoped and legacy token data.
func DeleteAllTokenData(configDir string) error {
if h := edition.Get(); h.DeleteToken != nil {
return h.DeleteToken(configDir)
return withProfilesLock(configDir, func() error {
return h.DeleteToken(configDir)
})
}
return withProfilesLock(configDir, func() error {
var firstErr error
+28 -9
View File
@@ -63,9 +63,10 @@ func (i Identity) Key() string {
}
type Client struct {
BaseURL string
HTTPClient *http.Client
Identity Identity
BaseURL string
HTTPClient *http.Client
Identity Identity
AccessTokenProvider func(context.Context) (string, error)
}
type CreateSubscriptionRequest struct {
@@ -337,8 +338,9 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
if c == nil {
return errors.New("personal event: nil client")
}
if c.Identity.AccessToken == "" {
return errors.New("personal event: access token is required")
accessToken, err := c.resolveAccessToken(ctx)
if err != nil {
return err
}
u := strings.TrimRight(c.BaseURL, "/") + path
if len(q) > 0 {
@@ -358,7 +360,7 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
if err != nil {
return fmt.Errorf("personal event: create request: %w", err)
}
c.decorate(req)
c.decorate(req, accessToken)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
@@ -424,9 +426,26 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
return json.Unmarshal(data, out)
}
func (c *Client) decorate(req *http.Request) {
req.Header.Set("Authorization", "Bearer "+c.Identity.AccessToken)
req.Header.Set("x-user-access-token", c.Identity.AccessToken)
func (c *Client) resolveAccessToken(ctx context.Context) (string, error) {
if c.AccessTokenProvider != nil {
token, err := c.AccessTokenProvider(ctx)
if err != nil {
return "", fmt.Errorf("personal event: resolve access token: %w", err)
}
if token = strings.TrimSpace(token); token != "" {
return token, nil
}
return "", errors.New("personal event: access token provider returned empty token")
}
if token := strings.TrimSpace(c.Identity.AccessToken); token != "" {
return token, nil
}
return "", errors.New("personal event: access token is required")
}
func (c *Client) decorate(req *http.Request, accessToken string) {
req.Header.Set("Authorization", "Bearer "+accessToken)
req.Header.Set("x-user-access-token", accessToken)
req.Header.Set("X-DWS-Client-Id", c.Identity.ClientID)
req.Header.Set("X-DWS-Source-Id", c.Identity.SourceID)
if c.Identity.CorpID != "" {
@@ -0,0 +1,56 @@
package personal
import (
"context"
"errors"
"io"
"net/http"
"strings"
"testing"
)
type accessTokenRoundTripper func(*http.Request) (*http.Response, error)
func (f accessTokenRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestCrossPlatformCoverageClientResolvesAccessTokenPerRequest(t *testing.T) {
tokens := []string{"token-a", "token-b"}
calls := 0
client := NewClient("https://control.test", Identity{AccessToken: "stale", ClientID: "client", SourceID: "source"})
client.AccessTokenProvider = func(context.Context) (string, error) {
token := tokens[calls]
calls++
return token, nil
}
client.HTTPClient = &http.Client{Transport: accessTokenRoundTripper(func(req *http.Request) (*http.Response, error) {
want := tokens[calls-1]
if got := req.Header.Get("Authorization"); got != "Bearer "+want {
t.Fatalf("Authorization = %q, want Bearer %s", got, want)
}
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"success":true,"result":{"items":[]}}`)), Header: make(http.Header)}, nil
})}
for range 2 {
if _, err := client.ListSubscriptions(context.Background(), ListOptions{}); err != nil {
t.Fatal(err)
}
}
if calls != 2 {
t.Fatalf("provider calls = %d, want 2", calls)
}
}
func TestCrossPlatformCoverageClientDoesNotFallBackAfterProviderFailure(t *testing.T) {
want := errors.New("keychain failed")
client := NewClient("https://control.test", Identity{AccessToken: "stale", ClientID: "client", SourceID: "source"})
client.AccessTokenProvider = func(context.Context) (string, error) { return "", want }
client.HTTPClient = &http.Client{Transport: accessTokenRoundTripper(func(*http.Request) (*http.Response, error) {
t.Fatal("HTTP must not run after token provider failure")
return nil, nil
})}
_, err := client.ListSubscriptions(context.Background(), ListOptions{})
if !errors.Is(err, want) {
t.Fatalf("error = %v, want %v", err, want)
}
}
+39 -15
View File
@@ -41,19 +41,22 @@ const (
)
type PersonalConfig struct {
AccessToken string
ClientID string
ClientSecret string
SourceID string
TicketURL string
TicketMode string
HTTPClient *http.Client
WebSocketDialer *websocket.Dialer
Now func() time.Time
ReconnectMin time.Duration
ReconnectMax time.Duration
AccessToken string
AccessTokenProvider AccessTokenProvider
ClientID string
ClientSecret string
SourceID string
TicketURL string
TicketMode string
HTTPClient *http.Client
WebSocketDialer *websocket.Dialer
Now func() time.Time
ReconnectMin time.Duration
ReconnectMax time.Duration
}
type AccessTokenProvider func(context.Context) (string, error)
type PersonalSource struct {
cfg PersonalConfig
machine *Machine
@@ -73,8 +76,8 @@ type ticketResponse struct {
}
func NewPersonal(cfg PersonalConfig) (*PersonalSource, error) {
if strings.TrimSpace(cfg.AccessToken) == "" {
return nil, errors.New("personal source: AccessToken is required")
if cfg.AccessTokenProvider == nil && strings.TrimSpace(cfg.AccessToken) == "" {
return nil, errors.New("personal source: AccessToken or AccessTokenProvider is required")
}
if strings.TrimSpace(cfg.ClientID) == "" {
return nil, errors.New("personal source: ClientID is required")
@@ -192,6 +195,10 @@ 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 {
return nil, err
}
body := map[string]any{
"sourceId": s.cfg.SourceID,
"mode": s.cfg.TicketMode,
@@ -207,8 +214,8 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Accept", "application/json")
req.Header.Set("x-user-access-token", s.cfg.AccessToken)
req.Header.Set("Authorization", "Bearer "+s.cfg.AccessToken)
req.Header.Set("x-user-access-token", accessToken)
req.Header.Set("Authorization", "Bearer "+accessToken)
req.Header.Set("X-DWS-Client-Id", s.cfg.ClientID)
req.Header.Set("X-DWS-Source-Id", s.cfg.SourceID)
@@ -238,6 +245,23 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
return ticket, nil
}
func resolveSourceAccessToken(ctx context.Context, provider AccessTokenProvider, fallback, component string) (string, error) {
if provider != nil {
token, err := provider(ctx)
if err != nil {
return "", fmt.Errorf("%s: resolve access token: %w", component, err)
}
if token = strings.TrimSpace(token); token != "" {
return token, nil
}
return "", fmt.Errorf("%s: access token provider returned empty token", component)
}
if token := strings.TrimSpace(fallback); token != "" {
return token, nil
}
return "", fmt.Errorf("%s: access token is required", component)
}
func (s *PersonalSource) handleFrame(conn *websocket.Conn, data []byte, emit dwsevent.EmitFn) error {
df, err := payload.DecodeDataFrame(data)
if err != nil {
+16 -11
View File
@@ -39,14 +39,15 @@ const (
// normal mode uses portal-side managed credentials; custom mode asks portal to
// open the user connection with the caller-provided clientId/clientSecret.
type PortalTicketConfig struct {
TicketURL string
AccessToken string
SourceID string
Mode string
ClientID string
ClientSecret string
UserAgent string
HTTPClient *http.Client
TicketURL string
AccessToken string
AccessTokenProvider AccessTokenProvider
SourceID string
Mode string
ClientID string
ClientSecret string
UserAgent string
HTTPClient *http.Client
}
var portalWriteMessage = func(conn *websocket.Conn, messageType int, data []byte) error {
@@ -60,8 +61,8 @@ func (c *PortalTicketConfig) Valid() error {
if strings.TrimSpace(c.TicketURL) == "" {
return errors.New("source: portal ticket URL is required")
}
if strings.TrimSpace(c.AccessToken) == "" {
return errors.New("source: portal access token is required")
if c.AccessTokenProvider == nil && strings.TrimSpace(c.AccessToken) == "" {
return errors.New("source: portal access token or provider is required")
}
if strings.TrimSpace(c.SourceID) == "" {
return errors.New("source: portal sourceId is required")
@@ -161,6 +162,10 @@ 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 {
return portalStreamTicket{}, err
}
httpClient := cfg.HTTPClient
if httpClient == nil {
httpClient = &http.Client{Timeout: 20 * time.Second}
@@ -184,7 +189,7 @@ func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalSt
if ua := strings.TrimSpace(cfg.UserAgent); ua != "" {
req.Header.Set("User-Agent", ua)
}
req.Header.Set("x-user-access-token", cfg.AccessToken)
req.Header.Set("x-user-access-token", accessToken)
resp, err := httpClient.Do(req)
if err != nil {
@@ -0,0 +1,66 @@
package source
import (
"context"
"errors"
"io"
"net/http"
"strings"
"testing"
)
type tokenProviderRoundTripper func(*http.Request) (*http.Response, error)
func (f tokenProviderRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}
func TestCrossPlatformCoveragePersonalSourceResolvesTokenForEveryTicketRequest(t *testing.T) {
tokens := []string{"token-a", "token-b"}
calls := 0
source, err := NewPersonal(PersonalConfig{
AccessTokenProvider: func(context.Context) (string, error) {
token := tokens[calls]
calls++
return token, nil
},
ClientID: "client",
SourceID: "source",
TicketURL: "https://ticket.test",
HTTPClient: &http.Client{Transport: tokenProviderRoundTripper(func(req *http.Request) (*http.Response, error) {
want := tokens[calls-1]
if got := req.Header.Get("x-user-access-token"); got != want {
t.Fatalf("token header = %q, want %q", got, want)
}
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"endpoint":"wss://stream.test","ticket":"ticket"}`)), Header: make(http.Header)}, nil
})},
})
if err != nil {
t.Fatal(err)
}
for range 2 {
if _, err := source.fetchTicket(context.Background()); err != nil {
t.Fatal(err)
}
}
if calls != 2 {
t.Fatalf("provider calls = %d, want 2", calls)
}
}
func TestCrossPlatformCoveragePortalTicketProviderFailureStopsBeforeHTTP(t *testing.T) {
want := errors.New("token store failed")
httpCalled := false
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
TicketURL: "https://ticket.test",
AccessTokenProvider: func(context.Context) (string, error) { return "", want },
SourceID: "source",
HTTPClient: &http.Client{Transport: tokenProviderRoundTripper(func(*http.Request) (*http.Response, error) {
httpCalled = true
return nil, errors.New("unexpected HTTP")
})},
})
if !errors.Is(err, want) || httpCalled {
t.Fatalf("request error = %v, httpCalled=%v", err, httpCalled)
}
}
+5
View File
@@ -36,6 +36,11 @@ const (
// parallel. When empty, the platform default applies.
StorageDirEnv = "DWS_KEYCHAIN_DIR"
// TestNamespaceEnv isolates the Windows HKCU registry backend for tests.
// Production code must not set it. Other platforms already isolate secure
// storage through StorageDirEnv and ignore this value.
TestNamespaceEnv = "DWS_KEYCHAIN_TEST_NAMESPACE"
// DisableKeychainEnv opts the macOS implementation out of system
// Keychain access for the DEK, falling back to a file-based DEK
// (same scheme as Linux). Intended for sandboxed runtimes where
+8
View File
@@ -14,6 +14,7 @@
package keychain
import (
"fmt"
"os"
"path/filepath"
"testing"
@@ -25,8 +26,15 @@ func TestMain(m *testing.M) {
panic(err)
}
_ = os.Setenv(StorageDirEnv, dir)
_ = os.Setenv(TestNamespaceEnv, dir)
_ = os.Setenv(DisableKeychainEnv, "1")
code := m.Run()
if err := RemoveAuthTokenEntries(Service); err != nil {
fmt.Fprintf(os.Stderr, "internal/keychain test cleanup: %v\n", err)
if code == 0 {
code = 1
}
}
_ = os.RemoveAll(dir)
os.Exit(code)
}
+19 -4
View File
@@ -16,6 +16,7 @@
package keychain
import (
"crypto/sha256"
"encoding/base64"
"errors"
"fmt"
@@ -47,9 +48,8 @@ const regRootPath = `Software\DwsCli\keychain`
// The Windows keychain backend keeps secrets in DPAPI-protected HKCU registry
// values rather than on disk, so this path is used only by the portable
// auth-bundle export/import (internal/auth) to colocate config. When the
// DWS_KEYCHAIN_DIR environment variable is set (used by tests for isolation),
// the storage root is taken from that env var instead; otherwise it defaults
// to %LocalAppData%\<service>.
// DWS_KEYCHAIN_DIR environment variable is set, the storage root is taken from
// that env var; otherwise it defaults to %LocalAppData%\<service>.
func StorageDir(service string) string {
if override := os.Getenv(StorageDirEnv); override != "" {
return filepath.Join(override, service)
@@ -66,7 +66,22 @@ func StorageDir(service string) string {
}
func registryPathForService(service string) string {
return regRootPath + `\` + safeRegistryComponent(service)
path := regRootPath + `\` + safeRegistryComponent(service)
namespace := strings.TrimSpace(os.Getenv(TestNamespaceEnv))
if namespace == "" {
return path
}
// Windows stores credentials in HKCU instead of DWS_KEYCHAIN_DIR. Tests set
// an explicit process namespace so concurrent package binaries cannot
// delete each other's credentials. Hash it to avoid leaking temp paths or
// introducing registry separators.
namespace = filepath.Clean(namespace)
if absolute, err := filepath.Abs(namespace); err == nil {
namespace = absolute
}
sum := sha256.Sum256([]byte(strings.ToLower(namespace)))
return fmt.Sprintf(`%s\test-%x`, path, sum[:16])
}
var safeRegRe = regexp.MustCompile(`[^a-zA-Z0-9._-]`)
@@ -17,12 +17,37 @@ package keychain
import (
"errors"
"strings"
"testing"
"golang.org/x/sys/windows"
"golang.org/x/sys/windows/registry"
)
func TestCrossPlatformCoverageRegistryPathForServiceHonorsTestNamespace(t *testing.T) {
t.Setenv(TestNamespaceEnv, "")
defaultPath := registryPathForService("service")
if defaultPath != regRootPath+`\service` {
t.Fatalf("default registry path = %q, want historical path %q", defaultPath, regRootPath+`\service`)
}
t.Setenv(TestNamespaceEnv, t.TempDir())
firstPath := registryPathForService("service")
t.Setenv(TestNamespaceEnv, t.TempDir())
secondPath := registryPathForService("service")
if firstPath == defaultPath || secondPath == defaultPath {
t.Fatalf("isolated registry paths = %q, %q; want paths distinct from %q", firstPath, secondPath, defaultPath)
}
if firstPath == secondPath {
t.Fatalf("isolated registry paths collide: %q", firstPath)
}
if !strings.HasPrefix(firstPath, defaultPath+`\test-`) {
t.Fatalf("isolated registry path = %q, want prefix %q", firstPath, defaultPath+`\test-`)
}
}
func TestDeleteRegistryValuePropagatesFailure(t *testing.T) {
originalDelete := registryDeleteValue
failure := errors.New("delete failed")
+9 -16
View File
@@ -4,17 +4,17 @@ set -eu
ROOT="$(CDPATH= cd -- "$(dirname -- "$0")/../.." && pwd)"
BASE_REF=""
OVERALL_PROFILE="coverage.txt"
ADDITIONAL_PROFILE="${COVERAGE_ADDITIONAL_PROFILE:-}"
ADDITIONAL_DIFF_PROFILE="${COVERAGE_ADDITIONAL_DIFF_PROFILE:-${COVERAGE_ADDITIONAL_PROFILE:-}}"
BASELINE_PROFILE="coverage-base.txt"
DIFF_PROFILE="coverage-policy.txt"
TARGET="${COVERAGE_TARGET:-80}"
OVERALL_TOLERANCE="${COVERAGE_OVERALL_TOLERANCE:-0.1}"
TARGET="${COVERAGE_TARGET:-100}"
OVERALL_TOLERANCE="${COVERAGE_OVERALL_TOLERANCE:-0}"
ENFORCE_OVERALL="${COVERAGE_ENFORCE_OVERALL:-false}"
CHANGED_ONLY="false"
SCOPE_BUILDABLE="false"
usage() {
printf '%s\n' "usage: $0 --base-ref <ref> [--changed-only] [--scope-buildable] [--overall-profile <file>] [--additional-profile <file>] [--baseline-profile <file>] [--diff-profile <file>]" >&2
printf '%s\n' "usage: $0 --base-ref <ref> [--changed-only] [--scope-buildable] [--overall-profile <file>] [--additional-diff-profile <file>] [--baseline-profile <file>] [--diff-profile <file>]" >&2
}
while [ "$#" -gt 0 ]; do
@@ -29,9 +29,9 @@ while [ "$#" -gt 0 ]; do
OVERALL_PROFILE="$2"
shift 2
;;
--additional-profile)
--additional-diff-profile|--additional-profile)
[ "$#" -ge 2 ] || { usage; exit 2; }
ADDITIONAL_PROFILE="$2"
ADDITIONAL_DIFF_PROFILE="$2"
shift 2
;;
--baseline-profile)
@@ -87,19 +87,12 @@ set -- "$CHECKER" \
if [ "$CHANGED_ONLY" = "true" ]; then
set -- "$@" --changed-only
else
baseline="$(go tool cover -func="$BASELINE_PROFILE" | awk '/^total:/ { gsub(/%/, "", $3); print $3 }')"
[ -n "$baseline" ] || {
printf 'error: cannot parse authoritative coverage from %s\n' "$BASELINE_PROFILE" >&2
exit 2
}
set -- "$@" \
--overall-profile "$OVERALL_PROFILE" \
--diff-profile "$OVERALL_PROFILE" \
--baseline-overall "$baseline"
if [ -n "$ADDITIONAL_PROFILE" ]; then
set -- "$@" \
--overall-profile "$ADDITIONAL_PROFILE" \
--diff-profile "$ADDITIONAL_PROFILE"
--baseline-profile "$BASELINE_PROFILE"
if [ -n "$ADDITIONAL_DIFF_PROFILE" ]; then
set -- "$@" --diff-profile "$ADDITIONAL_DIFF_PROFILE"
fi
fi
if [ "$SCOPE_BUILDABLE" = "true" ]; then
+51 -29
View File
@@ -13,7 +13,6 @@ import (
"flag"
"fmt"
"io"
"math"
"os"
"os/exec"
"path/filepath"
@@ -37,11 +36,13 @@ func (values *stringList) Set(value string) error {
}
type coverageBlock struct {
File string
StartLine int
EndLine int
Statements int
Count int
File string
StartLine int
StartColumn int
EndLine int
EndColumn int
Statements int
Count int
}
type lineRange struct {
@@ -78,6 +79,7 @@ func run(
buildableLoader func() (map[string]bool, error),
) int {
var overallPaths stringList
var baselinePaths stringList
var diffPaths stringList
var baseRef string
var modulePath string
@@ -90,12 +92,13 @@ func run(
flags := flag.NewFlagSet("coverage-gate", flag.ContinueOnError)
flags.SetOutput(stderr)
flags.Var(&overallPaths, "overall-profile", "coverage profile used for overall coverage (repeatable)")
flags.Var(&baselinePaths, "baseline-profile", "merge-base coverage profile evaluated with the same model as the candidate (repeatable)")
flags.Var(&diffPaths, "diff-profile", "coverage profile used for changed-code coverage (repeatable)")
flags.StringVar(&baseRef, "base-ref", "", "Git merge-base or previous main SHA")
flags.StringVar(&modulePath, "module", "", "Go module path used to normalize profile filenames")
flags.Float64Var(&baselineOverall, "baseline-overall", -1, "authoritative overall coverage percentage")
flags.Float64Var(&overallTolerance, "overall-tolerance", 0.1, "allowed overall coverage measurement variance in percentage points")
flags.Float64Var(&target, "target", 80, "required changed-code and eventual overall coverage percentage")
flags.Float64Var(&overallTolerance, "overall-tolerance", 0, "allowed overall coverage measurement variance in percentage points")
flags.Float64Var(&target, "target", 100, "required changed-code coverage percentage and optional overall floor")
flags.BoolVar(&enforceOverall, "enforce-overall-target", false, "require overall coverage to reach target")
flags.BoolVar(&changedOnly, "changed-only", false, "enforce changed-code coverage without an overall baseline")
flags.BoolVar(&scopeBuildable, "scope-buildable", false, "only evaluate changed files buildable on the current platform")
@@ -103,8 +106,13 @@ func run(
return 2
}
if len(diffPaths) == 0 || baseRef == "" || modulePath == "" || (!changedOnly && (len(overallPaths) == 0 || baselineOverall < 0)) {
fmt.Fprintln(stderr, "coverage-gate requires --diff-profile, --base-ref, and --module; overall mode also requires --overall-profile and --baseline-overall")
if len(diffPaths) == 0 || baseRef == "" || modulePath == "" ||
(!changedOnly && (len(overallPaths) == 0 || (len(baselinePaths) == 0 && baselineOverall < 0))) {
fmt.Fprintln(stderr, "coverage-gate requires --diff-profile, --base-ref, and --module; overall mode also requires --overall-profile and either --baseline-profile or --baseline-overall")
return 2
}
if len(baselinePaths) > 0 && baselineOverall >= 0 {
fmt.Fprintln(stderr, "coverage-gate accepts either --baseline-profile or --baseline-overall, not both")
return 2
}
var overall []coverageBlock
@@ -115,6 +123,14 @@ func run(
fmt.Fprintln(stderr, err)
return 2
}
if len(baselinePaths) > 0 {
baseline, baselineErr := readProfiles(baselinePaths, modulePath)
if baselineErr != nil {
fmt.Fprintln(stderr, baselineErr)
return 2
}
baselineOverall = coveragePercent(baseline)
}
}
diff, err := readProfiles(diffPaths, modulePath)
if err != nil {
@@ -150,12 +166,12 @@ func run(
if enforceOverall {
mode = "required"
}
fmt.Fprintf(stdout, "overall coverage: %.1f%% (merge-base %.1f%%; tolerance %.1fpp; target %.1f%%; %s)\n", result.Overall, baselineOverall, overallTolerance, target, mode)
fmt.Fprintf(stdout, "overall coverage: %.4f%% (merge-base %.4f%%; tolerance %.4fpp; target %.4f%%; %s)\n", result.Overall, baselineOverall, overallTolerance, target, mode)
}
if result.ChangedStatements == 0 {
fmt.Fprintf(stdout, "changed code coverage: n/a (no changed executable statements; target %.1f%%)\n", target)
fmt.Fprintf(stdout, "changed code coverage: n/a (no changed executable statements; target %.4f%%)\n", target)
} else {
fmt.Fprintf(stdout, "changed code coverage: %.1f%% (%d executable statements; target %.1f%%)\n", result.ChangedCoverage, result.ChangedStatements, target)
fmt.Fprintf(stdout, "changed code coverage: %.4f%% (%d executable statements; target %.4f%%)\n", result.ChangedCoverage, result.ChangedStatements, target)
}
if len(result.Failures) > 0 {
fmt.Fprintln(stderr, "coverage gate failed:")
@@ -171,13 +187,11 @@ func evaluate(input gateInput) gateResult {
result := gateResult{Failures: []string{}}
if !input.ChangedOnly {
result.Overall = coveragePercent(input.Overall)
baselineRounded := roundOne(input.BaselineOverall)
overallRounded := roundOne(result.Overall)
if overallRounded+input.OverallTolerance+1e-9 < baselineRounded {
result.Failures = append(result.Failures, fmt.Sprintf("overall coverage regressed from %.1f%% to %.1f%%", baselineRounded, overallRounded))
if result.Overall+input.OverallTolerance < input.BaselineOverall {
result.Failures = append(result.Failures, fmt.Sprintf("overall coverage regressed from %.4f%% to %.4f%%", input.BaselineOverall, result.Overall))
}
if input.EnforceOverall && overallRounded < input.Target {
result.Failures = append(result.Failures, fmt.Sprintf("overall coverage %.1f%% is below target %.1f%%", overallRounded, input.Target))
if input.EnforceOverall && result.Overall < input.Target {
result.Failures = append(result.Failures, fmt.Sprintf("overall coverage %.4f%% is below target %.4f%%", result.Overall, input.Target))
}
}
@@ -206,8 +220,8 @@ func evaluate(input gateInput) gateResult {
result.ChangedStatements = total
if total > 0 {
result.ChangedCoverage = float64(covered) * 100 / float64(total)
if result.ChangedCoverage+1e-9 < input.Target {
result.Failures = append(result.Failures, fmt.Sprintf("changed code coverage %.1f%% is below target %.1f%%", result.ChangedCoverage, input.Target))
if result.ChangedCoverage < input.Target {
result.Failures = append(result.Failures, fmt.Sprintf("changed code coverage %.4f%% is below target %.4f%%", result.ChangedCoverage, input.Target))
}
}
return result
@@ -253,11 +267,13 @@ func readProfiles(paths []string, modulePath string) ([]coverageBlock, error) {
values = append(values, value)
}
blocks = append(blocks, coverageBlock{
File: normalizeProfilePath(match[1], modulePath),
StartLine: values[0],
EndLine: values[2],
Statements: values[4],
Count: values[5],
File: normalizeProfilePath(match[1], modulePath),
StartLine: values[0],
StartColumn: values[1],
EndLine: values[2],
EndColumn: values[3],
Statements: values[4],
Count: values[5],
})
}
err = scanner.Err()
@@ -402,7 +418,15 @@ func coveragePercent(blocks []coverageBlock) float64 {
func mergeCoverageBlocks(blocks []coverageBlock) []coverageBlock {
merged := make(map[string]coverageBlock, len(blocks))
for _, block := range blocks {
key := fmt.Sprintf("%s:%d:%d:%d", block.File, block.StartLine, block.EndLine, block.Statements)
key := fmt.Sprintf(
"%s:%d:%d:%d:%d:%d",
block.File,
block.StartLine,
block.StartColumn,
block.EndLine,
block.EndColumn,
block.Statements,
)
current, ok := merged[key]
if !ok || block.Count > current.Count {
merged[key] = block
@@ -419,5 +443,3 @@ func mergeCoverageBlocks(blocks []coverageBlock) []coverageBlock {
}
return result
}
func roundOne(value float64) float64 { return math.Round(value*10) / 10 }
+198 -4
View File
@@ -39,7 +39,7 @@ func TestRun(t *testing.T) {
if code := run(args, &stdout, &stderr, loader, nil); code != 0 {
t.Fatalf("run code=%d stderr=%s", code, stderr.String())
}
if !strings.Contains(stdout.String(), "overall coverage: 100.0%") || !strings.Contains(stdout.String(), "changed code coverage: 100.0%") {
if !strings.Contains(stdout.String(), "overall coverage: 100.0000%") || !strings.Contains(stdout.String(), "changed code coverage: 100.0000%") {
t.Fatalf("unexpected output %q", stdout.String())
}
@@ -54,6 +54,115 @@ func TestRun(t *testing.T) {
}
}
func TestRunDefaultsToFullChangedCodeCoverage(t *testing.T) {
profile := filepath.Join(t.TempDir(), "coverage.out")
partialBody := "mode: atomic\n" +
"example.com/project/internal/a.go:10.1,12.2 9 1\n" +
"example.com/project/internal/a.go:13.1,13.2 1 0\n"
if err := os.WriteFile(profile, []byte(partialBody), 0o600); err != nil {
t.Fatal(err)
}
args := []string{
"--changed-only",
"--diff-profile", profile,
"--base-ref", "base",
"--module", "example.com/project",
}
loader := func(string) (map[string][]lineRange, error) {
return map[string][]lineRange{"internal/a.go": {{Start: 10, End: 13}}}, nil
}
var stdout, stderr bytes.Buffer
if code := run(args, &stdout, &stderr, loader, nil); code != 1 {
t.Fatalf("run code=%d, want 1; stdout=%s stderr=%s", code, stdout.String(), stderr.String())
}
if !strings.Contains(stderr.String(), "changed code coverage 90.0000% is below target 100.0000%") {
t.Fatalf("default target was not enforced: %q", stderr.String())
}
fullBody := "mode: atomic\n" +
"example.com/project/internal/a.go:10.1,12.2 9 1\n" +
"example.com/project/internal/a.go:13.1,13.2 1 1\n"
if err := os.WriteFile(profile, []byte(fullBody), 0o600); err != nil {
t.Fatal(err)
}
stdout.Reset()
stderr.Reset()
if code := run(args, &stdout, &stderr, loader, nil); code != 0 {
t.Fatalf("100%% changed coverage code=%d; stdout=%s stderr=%s", code, stdout.String(), stderr.String())
}
}
func TestRunEvaluatesBaselineProfileWithCandidateCoverageModel(t *testing.T) {
profile := filepath.Join(t.TempDir(), "coverage.out")
body := "mode: atomic\n" +
"example.com/project/internal/a.go:10.1,12.2 5 0\n" +
"example.com/project/internal/a.go:10.1,12.2 5 1\n" +
"example.com/project/internal/a.go:20.1,22.2 5 0\n"
if err := os.WriteFile(profile, []byte(body), 0o600); err != nil {
t.Fatal(err)
}
args := []string{
"--overall-profile", profile,
"--baseline-profile", profile,
"--diff-profile", profile,
"--base-ref", "base",
"--module", "example.com/project",
}
loader := func(string) (map[string][]lineRange, error) {
return map[string][]lineRange{}, nil
}
var stdout, stderr bytes.Buffer
if code := run(args, &stdout, &stderr, loader, nil); code != 0 {
t.Fatalf("run code=%d stderr=%s", code, stderr.String())
}
if !strings.Contains(stdout.String(), "overall coverage: 50.0000% (merge-base 50.0000%") {
t.Fatalf("baseline and candidate did not share one coverage model: %q", stdout.String())
}
}
func TestRunFailsClosedWhenBaselineProfileIsMissing(t *testing.T) {
profile := filepath.Join(t.TempDir(), "coverage.out")
body := "mode: atomic\nexample.com/project/internal/a.go:10.1,12.2 5 1\n"
if err := os.WriteFile(profile, []byte(body), 0o600); err != nil {
t.Fatal(err)
}
args := []string{
"--overall-profile", profile,
"--baseline-profile", filepath.Join(t.TempDir(), "missing.out"),
"--diff-profile", profile,
"--base-ref", "base",
"--module", "example.com/project",
}
loader := func(string) (map[string][]lineRange, error) {
return map[string][]lineRange{}, nil
}
var stdout, stderr bytes.Buffer
if code := run(args, &stdout, &stderr, loader, nil); code != 2 {
t.Fatalf("missing baseline profile code=%d, want 2; stdout=%s stderr=%s", code, stdout.String(), stderr.String())
}
if !strings.Contains(stderr.String(), "open coverage profile") {
t.Fatalf("missing baseline profile did not fail closed: %q", stderr.String())
}
}
func TestRunRejectsConflictingBaselineSources(t *testing.T) {
args := []string{
"--overall-profile", "candidate.out",
"--baseline-profile", "baseline.out",
"--baseline-overall", "100",
"--diff-profile", "candidate.out",
"--base-ref", "base",
"--module", "example.com/project",
}
var stdout, stderr bytes.Buffer
if code := run(args, &stdout, &stderr, nil, nil); code != 2 {
t.Fatalf("conflicting baseline sources code=%d, want 2; stdout=%s stderr=%s", code, stdout.String(), stderr.String())
}
if !strings.Contains(stderr.String(), "either --baseline-profile or --baseline-overall, not both") {
t.Fatalf("conflicting baseline sources were not rejected clearly: %q", stderr.String())
}
}
func TestCrossPlatformCoverageRunUnionsRepeatedOverallProfiles(t *testing.T) {
dir := t.TempDir()
uncovered := filepath.Join(dir, "uncovered.out")
@@ -81,7 +190,7 @@ func TestCrossPlatformCoverageRunUnionsRepeatedOverallProfiles(t *testing.T) {
if code := run(args, &stdout, &stderr, loader, nil); code != 0 {
t.Fatalf("run code=%d stderr=%s", code, stderr.String())
}
if !strings.Contains(stdout.String(), "overall coverage: 100.0%") || !strings.Contains(stdout.String(), "changed code coverage: 100.0%") {
if !strings.Contains(stdout.String(), "overall coverage: 100.0000%") || !strings.Contains(stdout.String(), "changed code coverage: 100.0000%") {
t.Fatalf("repeated profiles were not unioned: %q", stdout.String())
}
}
@@ -113,7 +222,7 @@ func TestRunChangedOnlyWithBuildableScope(t *testing.T) {
if code := run(args, &stdout, &stderr, loader, buildableLoader); code != 0 {
t.Fatalf("run code=%d stderr=%s", code, stderr.String())
}
if strings.Contains(stdout.String(), "overall coverage") || !strings.Contains(stdout.String(), "changed code coverage: 100.0%") {
if strings.Contains(stdout.String(), "overall coverage") || !strings.Contains(stdout.String(), "changed code coverage: 100.0000%") {
t.Fatalf("unexpected changed-only output %q", stdout.String())
}
@@ -143,7 +252,11 @@ func TestReadProfiles(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if len(blocks) != 1 || blocks[0].File != "internal/a.go" || blocks[0].Statements != 3 {
if len(blocks) != 1 ||
blocks[0].File != "internal/a.go" ||
blocks[0].StartColumn != 1 ||
blocks[0].EndColumn != 2 ||
blocks[0].Statements != 3 {
t.Fatalf("blocks=%v", blocks)
}
if _, err := readProfiles([]string{filepath.Join(t.TempDir(), "missing")}, "example.com/project"); err == nil {
@@ -275,6 +388,41 @@ func TestCrossPlatformCoverageUnionsDuplicateCrossPackageCoverageBlocks(t *testi
}
}
func TestCoverageKeepsDistinctBlocksOnTheSameLine(t *testing.T) {
result := evaluate(gateInput{
Diff: []coverageBlock{
{
File: "internal/a.go",
StartLine: 10,
StartColumn: 1,
EndLine: 10,
EndColumn: 5,
Statements: 1,
Count: 1,
},
{
File: "internal/a.go",
StartLine: 10,
StartColumn: 6,
EndLine: 10,
EndColumn: 10,
Statements: 1,
Count: 0,
},
},
Changed: map[string][]lineRange{"internal/a.go": {{Start: 10, End: 10}}},
Target: 100,
ChangedOnly: true,
})
if result.ChangedCoverage != 50 || result.ChangedStatements != 2 {
t.Fatalf("same-line blocks were collapsed: %#v", result)
}
if len(result.Failures) != 1 ||
!strings.Contains(result.Failures[0], "changed code coverage 50.0000% is below target 100.0000%") {
t.Fatalf("same-line uncovered block did not fail the gate: %v", result.Failures)
}
}
func TestEvaluateAllowsMeasurementTolerance(t *testing.T) {
result := evaluate(gateInput{
Overall: []coverageBlock{{Statements: 403, Count: 1}, {Statements: 597, Count: 0}},
@@ -288,6 +436,52 @@ func TestEvaluateAllowsMeasurementTolerance(t *testing.T) {
}
}
func TestEvaluateDefaultsToZeroOverallTolerance(t *testing.T) {
result := evaluate(gateInput{
Overall: []coverageBlock{{Statements: 403, Count: 1}, {Statements: 597, Count: 0}},
Changed: map[string][]lineRange{},
BaselineOverall: 40.4,
OverallTolerance: 0,
Target: 100,
})
if len(result.Failures) != 1 ||
!strings.Contains(result.Failures[0], "overall coverage regressed from 40.4000% to 40.3000%") {
t.Fatalf("zero-tolerance regression failures = %v", result.Failures)
}
}
func TestEvaluateRejectsRegressionHiddenByOneDecimalRounding(t *testing.T) {
result := evaluate(gateInput{
Overall: []coverageBlock{
{Statements: 1999, Count: 1},
{Statements: 1, Count: 0},
},
Changed: map[string][]lineRange{},
BaselineOverall: 100,
Target: 100,
})
if len(result.Failures) != 1 ||
!strings.Contains(result.Failures[0], "overall coverage regressed from 100.0000% to 99.9500%") {
t.Fatalf("sub-tenth regression failures = %v", result.Failures)
}
}
func TestEvaluateRejectsChangedCoverageBelow100WithoutRounding(t *testing.T) {
result := evaluate(gateInput{
Diff: []coverageBlock{
{File: "internal/a.go", StartLine: 1, EndLine: 1999, Statements: 1999, Count: 1},
{File: "internal/a.go", StartLine: 2000, EndLine: 2000, Statements: 1, Count: 0},
},
Changed: map[string][]lineRange{"internal/a.go": {{Start: 1, End: 2000}}},
Target: 100,
ChangedOnly: true,
})
if len(result.Failures) != 1 ||
!strings.Contains(result.Failures[0], "changed code coverage 99.9500% is below target 100.0000%") {
t.Fatalf("sub-100 changed-code failures = %v", result.Failures)
}
}
func TestEvaluateFailsClosed(t *testing.T) {
result := evaluate(gateInput{
Overall: []coverageBlock{
+13
View File
@@ -47,6 +47,19 @@ if [ -n "$missing" ]; then
exit 0
fi
# ossutil v2 signs with V4 and requires an explicit region. Derive it from the
# endpoint host (oss-<region>[-internal].aliyuncs.com) unless OSS_REGION is set.
if [ -z "${OSS_REGION:-}" ]; then
OSS_REGION="$(printf '%s' "$OSS_ENDPOINT" \
| sed -n 's#^\(https\{0,1\}://\)\{0,1\}oss-\([a-z0-9-]*[a-z0-9]\)\.aliyuncs\.com.*#\2#p' \
| sed 's#-internal$##')"
fi
if [ -z "$OSS_REGION" ]; then
echo "❌ Could not derive OSS_REGION from OSS_ENDPOINT=${OSS_ENDPOINT}; set OSS_REGION explicitly." >&2
exit 1
fi
export OSS_REGION
# ── Resolve version ──────────────────────────────────────────────────────────
VERSION="${VERSION:-$(git describe --tags --always 2>/dev/null || echo dev)}"
CHANNEL="${DWS_RELEASE_CHANNEL:-$(release_channel_for_version "$VERSION")}"
+30 -1
View File
@@ -401,7 +401,7 @@ func TestChangelogPRFastPathWorkflowContract(t *testing.T) {
admission := readWorkflow(".github/workflows/ci.yml")
for _, want := range []string{
"name: Code Admission — PR 合入门禁",
"name: CI",
"files.length === 1",
"files[0].filename === 'CHANGELOG.md'",
"files[0].status === 'modified'",
@@ -417,6 +417,10 @@ func TestChangelogPRFastPathWorkflowContract(t *testing.T) {
"mode=--fast-path",
`"$mode" "$PR_BASE_SHA" HEAD`,
"needs.lint.outputs.platform_sensitive == 'true'",
`COVERAGE_TARGET: "100"`,
`COVERAGE_ENFORCE_OVERALL: "false"`,
`COVERAGE_OVERALL_TOLERANCE: "0"`,
`COVERAGE_ADDITIONAL_DIFF_PROFILE=coverage-shortcut.txt`,
} {
if !strings.Contains(admission, want) {
t.Errorf("Code Admission workflow missing contract %q", want)
@@ -451,6 +455,31 @@ func TestChangelogPRFastPathWorkflowContract(t *testing.T) {
t.Error("Code Admission must not suppress required contexts with paths-ignore")
}
notification := readWorkflow(".github/workflows/notify-wukong.yml")
if !strings.Contains(notification, "- CI") {
t.Error("Wukong notification must follow the renamed CI workflow")
}
if strings.Contains(notification, "Code Admission — PR 合入门禁") {
t.Error("Wukong notification still follows the retired workflow display name")
}
coverageGate := readWorkflow("scripts/policy/check-coverage-gate.sh")
if !strings.Contains(coverageGate, `TARGET="${COVERAGE_TARGET:-100}"`) {
t.Error("coverage gate must default to 100% changed-code coverage")
}
if !strings.Contains(coverageGate, `OVERALL_TOLERANCE="${COVERAGE_OVERALL_TOLERANCE:-0}"`) {
t.Error("coverage gate must reject any reported overall regression")
}
if !strings.Contains(coverageGate, `--baseline-profile "$BASELINE_PROFILE"`) {
t.Error("coverage gate must evaluate the merge-base profile with the candidate checker")
}
if strings.Contains(coverageGate, `--overall-profile "$ADDITIONAL_DIFF_PROFILE"`) {
t.Error("supporting changed-code coverage must not inflate candidate overall coverage")
}
if strings.Contains(coverageGate, `go tool cover -func="$BASELINE_PROFILE"`) {
t.Error("coverage baseline must not use a different coverage calculator")
}
aiBehavior := readWorkflow(".github/workflows/ai-behavior-check.yml")
for _, want := range []string{
"name: Code Admission — AI Behavior",
+3
View File
@@ -1311,6 +1311,7 @@ func TestReleaseMirrorUsesChannelSpecificPointer(t *testing.T) {
"OSS_ACCESS_KEY_ID=test-key",
"OSS_ACCESS_KEY_SECRET=test-secret",
"OSS_ENDPOINT=https://oss.example.com",
"OSS_REGION=cn-test",
"OSS_BUCKET=test-bucket",
"OSS_PREFIX=dws",
"OSSUTIL="+fakeOSSUtil,
@@ -1381,6 +1382,7 @@ func TestReleaseMirrorFailsClosedWhenPointerCannotBeRead(t *testing.T) {
"OSS_ACCESS_KEY_ID=test-key",
"OSS_ACCESS_KEY_SECRET=test-secret",
"OSS_ENDPOINT=https://oss.example.com",
"OSS_REGION=cn-test",
"OSS_BUCKET=test-bucket",
"OSS_PREFIX=dws",
"OSSUTIL="+fakeOSSUtil,
@@ -1417,6 +1419,7 @@ func TestReleaseMirrorRepairsHistoricalAssetsWithoutMovingNewerPointer(t *testin
"OSS_ACCESS_KEY_ID=test-key",
"OSS_ACCESS_KEY_SECRET=test-secret",
"OSS_ENDPOINT=https://oss.example.com",
"OSS_REGION=cn-test",
"OSS_BUCKET=test-bucket",
"OSS_PREFIX=dws",
"OSSUTIL="+fakeOSSUtil,