Compare commits

...
37 changed files with 312 additions and 2386 deletions
-1
View File
@@ -13,7 +13,6 @@ 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.
-223
View File
@@ -1,223 +0,0 @@
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)
}
}
+48 -244
View File
@@ -21,61 +21,19 @@ 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 {
discard := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, discard)
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, disc)
configureOAuthProviderCompatibility(provider, configDir)
return provider
}
@@ -86,218 +44,64 @@ var (
}
)
// 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
// 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
}
if strings.TrimSpace(configDir) == "" {
return AccessTokenSnapshot{}, fmt.Errorf("config directory is empty")
return "", fmt.Errorf("config directory is empty")
}
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
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
return tok, nil
}
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
return "", noCredentialsError()
}
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())
tok, err := resolveAccessTokenFromDir(ctx, configDir)
if err != nil {
return "", err
}
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
if tok != "" {
return tok, nil
}
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)
return "", noCredentialsError()
}
func noCredentialsError() error {
if edition.Get().IsEmbedded {
return fmt.Errorf("认证信息已失效,请重新认证: %w", authpkg.ErrTokenDataNotFound)
return fmt.Errorf("认证信息已失效,请重新认证")
}
return fmt.Errorf("no credentials found, run: dws auth login: %w", authpkg.ErrTokenDataNotFound)
return fmt.Errorf("no credentials found, run: dws auth login")
}
+8 -16
View File
@@ -34,10 +34,6 @@ 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
@@ -216,16 +212,14 @@ 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: authpkg.ErrTokenDataNotFound}
}
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{err: errors.New("missing")} }
newLegacyTokenManager = func(string) legacyTokenGetter { return fakeLegacyTokenGetter{token: " legacy "} }
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); err != nil || got != "legacy" {
t.Fatalf("legacy token = %q, %v", got, err)
}
authpkg.SetRuntimeProfile("corp:user")
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || err == nil || err.Error() != "missing" {
t.Fatalf("explicit profile fallback = token %q error %v, want profile error", got, err)
}
authpkg.SetRuntimeProfile("")
@@ -252,10 +246,10 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
}
func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T) {
oldLoad, oldFactory := loadRefreshTokenData, newRefreshProvider
oldMark, oldFactory := markAccessTokenStale, newRefreshProvider
oldStop := stopStdio
t.Cleanup(func() {
loadRefreshTokenData, newRefreshProvider = oldLoad, oldFactory
markAccessTokenStale, newRefreshProvider = oldMark, oldFactory
stopStdio = oldStop
stdioMu.Lock()
stdioClients = make(map[string]*transport.StdioClient)
@@ -263,13 +257,11 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
})
fail := errors.New("failure")
_ = oldFactory(t.TempDir())
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) { return nil, fail }
markAccessTokenStale = func(string) error { return fail }
if _, err := ForceRefreshAccessToken(context.Background(), "config"); !errors.Is(err, fail) {
t.Fatalf("load rejected token error = %v", err)
}
loadRefreshTokenData = func(string) (*authpkg.TokenData, error) {
return &authpkg.TokenData{AccessToken: "rejected"}, nil
t.Fatalf("mark stale error = %v", err)
}
markAccessTokenStale = func(string) error { return nil }
for _, tc := range []struct {
getter fakeAccessTokenGetter
want string
@@ -278,7 +270,7 @@ func TestCrossPlatformCoverageForceRefreshAndStdioFailureCoverage(t *testing.T)
{getter: fakeAccessTokenGetter{token: " "}, want: "empty"},
{getter: fakeAccessTokenGetter{token: " refreshed "}},
} {
newRefreshProvider = func(string) rejectedAccessTokenRefresher { return tc.getter }
newRefreshProvider = func(string) accessTokenGetter { 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,15 +15,6 @@ 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
@@ -32,35 +23,8 @@ 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 {
@@ -70,103 +34,3 @@ 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
}
@@ -1,354 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package 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)
}
}
+1 -3
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()), flags: &GlobalFlags{Token: "token"}}
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client())}
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,8 +1650,6 @@ 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",
+1 -2
View File
@@ -4,7 +4,6 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
@@ -165,7 +164,7 @@ func TestCrossPlatformCoverageRawAPIAndTokenCoverage(t *testing.T) {
}
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{} }
missing := t.TempDir()
if got, err := resolveAccessTokenFromDir(context.Background(), missing); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
if got, err := resolveAccessTokenFromDir(context.Background(), missing); err != nil || got != "" {
t.Fatalf("missing access token = %q, %v", got, err)
}
if _, err := ResolveAuxiliaryAccessToken(context.Background(), missing, ""); err == nil {
+11 -5
View File
@@ -413,7 +413,7 @@ func eventStreamBusID(streamOpts eventStreamTicketOptions) string {
return "portal-ticket-normal:" + sourceID
}
func newEventSource(_ context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
func newEventSource(ctx context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
if !streamOpts.enabled() {
return eventNewDingtalkSource(source.Config{
ClientID: clientID,
@@ -421,6 +421,14 @@ func newEventSource(_ context.Context, configDir, clientID, clientSecret string,
})
}
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() {
@@ -432,10 +440,8 @@ func newEventSource(_ context.Context, configDir, clientID, clientSecret string,
ClientID: portalClientID,
ClientSecret: portalClientSecret,
PortalTicket: &source.PortalTicketConfig{
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
AccessTokenProvider: func(ctx context.Context) (string, error) {
return eventResolveAccessToken(ctx, configDir, "")
},
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
AccessToken: token,
SourceID: eventStreamSourceID(streamOpts.SourceID),
Mode: streamOpts.Mode,
ClientID: portalClientID,
@@ -132,18 +132,14 @@ func TestCrossPlatformCoverageEventSourcesAndForegroundCoverage(t *testing.T) {
if _, err := newEventSource(context.Background(), "config", "client", "secret", eventStreamTicketOptions{}); err != nil {
t.Fatal(err)
}
stream := eventStreamTicketOptions{Mode: "custom"}
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 "", fail }
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err != nil {
t.Fatalf("stream source construction = %v", err)
stream := eventStreamTicketOptions{Mode: "custom"}
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); !errors.Is(err, fail) {
t.Fatalf("stream token error = %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 " ", 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 "token", nil }
for _, mode := range []string{"custom", "normal"} {
+5 -19
View File
@@ -259,7 +259,7 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
return personalConsumeRun(ctx, cfg)
}
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
client := personal.NewClient(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(newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
subs, err := personalListSubscriptions(personal.NewClient(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 := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
client := personal.NewClient(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,10 +723,7 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
if err != nil {
return personal.Identity{}, err
}
tokenData, err := personalLoadTokenData(configDir)
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
return personal.Identity{}, fmt.Errorf("load OAuth identity metadata: %w", err)
}
tokenData, _ := personalLoadTokenData(configDir)
var corpID, userID, clientID, refreshToken string
if tokenData != nil {
corpID = tokenData.CorpID
@@ -772,15 +769,6 @@ 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 == "" {
@@ -829,9 +817,7 @@ func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptio
}
_ = ctx
return source.NewPersonal(source.PersonalConfig{
AccessTokenProvider: func(ctx context.Context) (string, error) {
return personalResolveAuxiliaryAccessToken(ctx, opts.ConfigDir, "")
},
AccessToken: opts.Identity.AccessToken,
ClientID: clientID,
ClientSecret: clientSecret,
SourceID: opts.Identity.SourceID,
+16 -27
View File
@@ -27,13 +27,9 @@ type accessTokenGetter interface {
GetAccessToken(context.Context) (string, error)
}
type rejectedAccessTokenRefresher interface {
ForceRefreshRejectedToken(context.Context, string) (string, error)
}
var (
loadRefreshTokenData = authpkg.LoadTokenData
newRefreshProvider = func(configDir string) rejectedAccessTokenRefresher {
markAccessTokenStale = authpkg.MarkAccessTokenStale
newRefreshProvider = func(configDir string) accessTokenGetter {
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
provider := authpkg.NewOAuthProvider(configDir, disc)
configureOAuthProviderCompatibility(provider, configDir)
@@ -46,33 +42,26 @@ var (
// server-side rejection (HTTP 401 or business code such as
// TOKEN_VERIFIED_FAILED) on what locally appeared to be a still-valid token.
//
// 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.
// 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.
func ForceRefreshAccessToken(ctx context.Context, configDir string) (string, error) {
if strings.TrimSpace(configDir) == "" {
return "", fmt.Errorf("config directory is empty")
}
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")
if err := markAccessTokenStale(configDir); err != nil {
return "", fmt.Errorf("mark access token stale: %w", err)
}
provider := newRefreshProvider(configDir)
tok, err := provider.ForceRefreshRejectedToken(ctx, rejectedAccessToken)
tok, err := provider.GetAccessToken(ctx)
if err != nil {
return "", err
}
+20 -20
View File
@@ -56,7 +56,7 @@ var openBrowserFunc = tryOpenBrowser
var (
patAuthorizationTimeout = PatAuthRetryTimeout
patAuthorizationPollInterval = PatAuthPollInterval
patResolveAccessToken = ResolveAuxiliaryAccessToken
patLoadTokenData = authpkg.LoadTokenData
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, error) {
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
timeout := patAuthorizationTimeout
deadline := time.Now().Add(timeout)
pollTicker := time.NewTicker(patAuthorizationPollInterval)
@@ -290,26 +290,27 @@ func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Wr
select {
case <-ctx.Done():
fmt.Fprintf(output, "%s 操作已取消\n", tui.StateMark("error"))
return false, ctx.Err()
return false
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, nil
return false
case <-pollTicker.C:
pollCount++
elapsed := time.Since(start).Truncate(time.Second)
remaining := time.Until(deadline).Truncate(time.Second)
// 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)
// 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
}
}
// Show polling status
@@ -339,10 +340,7 @@ func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocati
PrintPatAuthError(output, scopeErr)
// Wait for user to complete authorization
authorized, waitErr := patWaitForAuthorization(ctx, configDir, output)
if waitErr != nil {
return executor.Result{}, waitErr
}
authorized := patWaitForAuthorization(ctx, configDir, output)
if !authorized {
return executor.Result{}, apperrors.NewAuth(
"等待用户授权超时",
@@ -796,6 +794,12 @@ 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 {
@@ -824,10 +828,6 @@ 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,41 +52,39 @@ func TestCrossPlatformCoveragePATRetryRemainingPureAndWaitCoverage(t *testing.T)
oldTimeout := patAuthorizationTimeout
oldInterval := patAuthorizationPollInterval
oldResolve := patResolveAccessToken
oldLoad := patLoadTokenData
t.Cleanup(func() {
patAuthorizationTimeout = oldTimeout
patAuthorizationPollInterval = oldInterval
patResolveAccessToken = oldResolve
patLoadTokenData = oldLoad
})
patAuthorizationTimeout = 50 * time.Millisecond
patAuthorizationPollInterval = time.Millisecond
patResolveAccessToken = func(context.Context, string, string) (string, error) {
return "token", nil
patLoadTokenData = func(string) (*authpkg.TokenData, error) {
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
}
out.Reset()
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || !ok {
if !WaitForPatAuthorization(context.Background(), "", &out) {
t.Fatal("valid token did not authorize")
}
ctx, cancel := context.WithCancel(context.Background())
cancel()
out.Reset()
if ok, err := WaitForPatAuthorization(ctx, "", &out); ok || !errors.Is(err, context.Canceled) {
t.Fatalf("cancelled authorization = %v, %v", ok, err)
if WaitForPatAuthorization(ctx, "", &out) {
t.Fatal("cancelled authorization succeeded")
}
patAuthorizationTimeout = time.Millisecond
patAuthorizationPollInterval = time.Hour
out.Reset()
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok {
t.Fatalf("timed out authorization = %v, %v", ok, err)
if WaitForPatAuthorization(context.Background(), "", &out) {
t.Fatal("timed out authorization succeeded")
}
patAuthorizationTimeout = 5 * time.Millisecond
patAuthorizationPollInterval = time.Millisecond
patResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
out.Reset()
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())
if WaitForPatAuthorization(context.Background(), "", &out) || !strings.Contains(out.String(), "等待授权中") {
t.Fatalf("invalid-token polling output = %q", out.String())
}
}
@@ -111,12 +109,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, error) { return false, nil }
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return false }
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, error) { return true, nil }
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return true }
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)
}
@@ -217,15 +215,15 @@ func patRaw(flowID, clientID, secret string) string {
func TestCrossPlatformCoveragePATRetryRemainingPollAndBrowserCoverage(t *testing.T) {
oldDo := patPollHTTPDo
oldRequest := patPollNewRequest
oldResolve := patResolveAccessToken
oldLoad := patLoadTokenData
oldBrowser := patBrowserOpenCommand
t.Cleanup(func() {
patPollHTTPDo = oldDo
patPollNewRequest = oldRequest
patResolveAccessToken = oldResolve
patLoadTokenData = oldLoad
patBrowserOpenCommand = oldBrowser
})
patResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return &authpkg.TokenData{AccessToken: "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 {
+1 -5
View File
@@ -333,11 +333,7 @@ func (r *recoveryRuntime) CallToolDirect(ctx context.Context, serverID, toolName
if err != nil {
return nil, err
}
authToken, err := resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags))
if err != nil {
return nil, tokenResolutionError(err)
}
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
tc := r.transport.WithAuth(resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags)), resolveIdentityHeaders())
result, err := tc.CallTool(ctx, endpoint, toolName, args)
if err != nil {
return nil, err
+56 -66
View File
@@ -235,9 +235,7 @@ 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 func() {
_, _ = runnerGetCachedRuntimeToken(ctx)
}()
go runnerGetCachedRuntimeToken(ctx)
}
if shouldUseDirectRuntime(invocation) {
@@ -536,12 +534,8 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
authToken := ""
if hasPluginAuth {
authToken = pluginAuth.Token
} 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)
}
} else {
authToken = r.resolveAuthToken(ctx)
}
var timeoutSec int
@@ -623,12 +617,6 @@ 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
}
@@ -637,15 +625,9 @@ 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 isRefreshableTransportAuthError(err) {
if isAuthError(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
}
@@ -670,12 +652,6 @@ 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
}
}
@@ -696,12 +672,6 @@ 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
}
@@ -826,49 +796,67 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
}, nil
}
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) (string, error) {
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
explicitToken := ""
if r != nil && r.globalFlags != nil {
explicitToken = r.globalFlags.Token
}
return resolveRuntimeAuthToken(ctx, explicitToken)
}
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) (string, error) {
snapshot, err := runtimeTokenManager.Get(ctx, defaultConfigDir(), explicitToken)
if err != nil {
return "", err
if token := strings.TrimSpace(explicitToken); token != "" {
return token
}
return snapshot.AccessToken, nil
if tp := edition.Get().TokenProvider; tp != nil {
token, _ := tp(ctx, func() (string, error) {
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
})
return token
}
return getCachedRuntimeToken(ctx)
}
// 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) {
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
if token := strings.TrimSpace(explicitToken); token != "" {
return token
}
// Use cached token to avoid repeated Keychain access (~70ms per call)
return getCachedRuntimeToken(ctx)
}
// 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()
loadStart := time.Now()
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
return resolveRuntimeAuthToken(ctx, "")
}
func tokenResolutionError(err error) error {
if err == nil {
return nil
configDir := defaultConfigDir()
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
slog.Error(tokenErr.Error())
return ""
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return err
if token == "" {
return ""
}
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)
cachedRuntimeTokenMu.Lock()
cachedRuntimeTokens[cacheKey] = token
cachedRuntimeTokenMu.Unlock()
return token
}
// generateExecutionID returns a random 16-char hex string used to correlate
@@ -883,7 +871,9 @@ func generateExecutionID() string {
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
// This should be called after login/logout operations.
func ResetRuntimeTokenCache() {
runtimeTokenManager.Invalidate()
cachedRuntimeTokenMu.Lock()
defer cachedRuntimeTokenMu.Unlock()
cachedRuntimeTokens = map[string]string{}
}
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, error) {
runnerGetCachedRuntimeToken = func(context.Context) string {
prefetched <- struct{}{}
return "", nil
return ""
}
r := &runtimeRunner{
loader: cli.CatalogLoaderFrom(cli.Catalog{}, wantErr),
@@ -197,7 +197,7 @@ func TestCrossPlatformCoverageRunnerRemainingExecutionCoverage(t *testing.T) {
return nil
}
authErr := apperrors.NewAuth("expired", apperrors.WithReason("http_401"))
authErr := apperrors.NewAuth("expired")
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, err := r.resolveAuthToken(context.Background()); err != nil || got != "explicit" {
t.Fatalf("explicit auth token = %q, %v", got, err)
if got := r.resolveAuthToken(context.Background()); got != "explicit" {
t.Fatalf("explicit auth token = %q", got)
}
edition.Override(&edition.Hooks{TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
_, _ = fallback()
return "provided", nil
}})
r.globalFlags.Token = ""
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "provided" {
t.Fatalf("provided auth token = %q, %v", got, err)
if got := r.resolveAuthToken(context.Background()); got != "provided" {
t.Fatalf("provided auth token = %q", got)
}
if got, err := resolveRuntimeAuthToken(context.Background(), " runtime "); err != nil || got != "runtime" {
t.Fatalf("runtime explicit token = %q, %v", got, err)
if got := resolveRuntimeAuthToken(context.Background(), " runtime "); got != "runtime" {
t.Fatalf("runtime explicit token = %q", got)
}
t.Setenv(envDWSChannel, "channel")
+26 -30
View File
@@ -17,7 +17,6 @@ import (
"archive/zip"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"mime"
@@ -37,25 +36,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
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() }
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() }
)
func init() {
@@ -297,7 +296,7 @@ func newSkillAddHintCommand() *cobra.Command {
func runSkillGet(cmd *cobra.Command, args []string) error {
skillID, _ := cmd.Flags().GetString("skill-id")
accessToken, err := skillLoadAccessToken(cmd.Context())
accessToken, err := skillLoadAccessToken()
if err != nil {
return err
}
@@ -320,7 +319,7 @@ func runSkillFind(cmd *cobra.Command, args []string) error {
if source == "" {
source, _ = cmd.Flags().GetString("scopes")
}
accessToken, err := skillLoadAccessToken(cmd.Context())
accessToken, err := skillLoadAccessToken()
if err != nil {
return err
}
@@ -389,7 +388,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(cmd.Context())
accessToken, err := skillLoadAccessToken()
if err != nil {
return err
}
@@ -442,16 +441,13 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
return nil
}
func loadSkillAccessToken(ctx context.Context) (string, error) {
func loadSkillAccessToken() (string, error) {
configDir := defaultConfigDir()
token, err := skillResolveAccessToken(ctx, configDir, "")
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
tokenData, err := skillLoadTokenData(configDir)
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
return "", skillAuthError()
}
if err != nil {
return "", fmt.Errorf("resolve skill access token: %w", err)
}
return token, nil
return tokenData.AccessToken, nil
}
func skillAuthError() error {
@@ -57,11 +57,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
})
fail := errors.New("failure")
cmd := skillCoverageCommand()
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
skillLoadAccessToken = func() (string, error) { return "", fail }
if err := runSkillGet(cmd, nil); !errors.Is(err, fail) {
t.Fatalf("skill get auth error = %v", err)
}
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
skillLoadAccessToken = func() (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(context.Context) (string, error) { return "", fail }
skillLoadAccessToken = func() (string, error) { return "", fail }
if err := runSkillFind(cmd, nil); !errors.Is(err, fail) {
t.Fatalf("skill find auth error = %v", err)
}
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
skillLoadAccessToken = func() (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(context.Context) (string, error) { return "", fail }
skillLoadAccessToken = func() (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(context.Context) (string, error) { return "token", nil }
skillLoadAccessToken = func() (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,35 +152,25 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
func TestCrossPlatformCoverageSkillCommandLowLevelRemainingCoverage(t *testing.T) {
oldHTTP := skillHTTPDo
oldNewRequest, oldResolveToken := skillNewRequest, skillResolveAccessToken
oldNewRequest, oldLoadToken := skillNewRequest, skillLoadTokenData
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, skillResolveAccessToken = oldNewRequest, oldResolveToken
skillNewRequest, skillLoadTokenData = oldNewRequest, oldLoadToken
skillUserHomeDir = oldHome
skillMkdirTemp, skillCreate, skillCreateTemp = oldMkdirTemp, oldCreate, oldCreateTemp
skillRemoveAll, skillRemove, skillMkdirAll = oldRemoveAll, oldRemove, oldMkdir
skillOpenFile, skillCopy, skillOpenZipFile = oldOpen, oldCopy, oldZipOpen
})
fail := errors.New("failure")
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
if _, err := loadSkillAccessToken(context.Background()); err == nil {
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
if _, err := loadSkillAccessToken(); err == nil {
t.Fatal("invalid skill access token succeeded")
}
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
skillLoadTokenData = oldLoadToken
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")
+13 -9
View File
@@ -18,6 +18,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"os"
@@ -396,11 +397,9 @@ func TestSkillInstallRequiresAuth(t *testing.T) {
configDir := filepath.Join(tempDir, "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
t.Cleanup(CloseFileLogger)
originalResolveToken := skillResolveAccessToken
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "", authpkg.ErrTokenDataNotFound
}
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
originalLoadToken := skillLoadTokenData
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, errors.New("missing") }
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
// Ensure the config directory exists but has no token
if err := os.MkdirAll(configDir, 0755); err != nil {
@@ -680,11 +679,16 @@ func TestSkillSearchUsesSourceQueryAndKeepsScopesCompat(t *testing.T) {
configDir := filepath.Join(t.TempDir(), "config")
t.Setenv("DWS_CONFIG_DIR", configDir)
t.Cleanup(CloseFileLogger)
originalResolveToken := skillResolveAccessToken
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
return "test-token", nil
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
}
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
var gotSources []string
var gotScopes []string
-10
View File
@@ -52,10 +52,6 @@ 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())
@@ -83,12 +79,6 @@ 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 {
@@ -1,70 +0,0 @@
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)
}
}
+1 -233
View File
@@ -13,50 +13,7 @@
package auth
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,
}
import "time"
// MarkAccessTokenStale loads the persisted TokenData, sets ExpiresAt to a past
// instant (preserving access_token and refresh_token), and writes it back. The
@@ -82,192 +39,3 @@ 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
}
+2 -4
View File
@@ -47,10 +47,8 @@ func (m *Manager) GetToken() (string, string, error) {
}
return token, "file", nil
}
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)
return "", "", fmt.Errorf("%s", i18n.T("未找到认证信息,请运行 dws auth login"))
}
func (m *Manager) GetMCPURL() (string, error) {
+10 -30
View File
@@ -115,12 +115,6 @@ 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() {
@@ -642,50 +636,36 @@ continueLogin:
return tokenData, nil
}
// 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) {
// 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 := oauthLoadToken(p.configDir)
if err != nil {
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)
return "", errors.New(i18n.T("未登录,请运行 dws auth login"))
}
// Fast path: access_token still valid — no lock needed.
if data.IsAccessTokenValid() {
return data, nil
return data.AccessToken, nil
}
// Slow path: token expired — try locked refresh.
if data.IsRefreshTokenValid() {
refreshed, rErr := p.lockedRefresh(ctx)
if rErr == nil {
return refreshed, nil
return refreshed.AccessToken, nil
}
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
if p.logger != nil {
p.logger.Warn(i18n.T("refresh_token 刷新失败"), "error", rErr)
}
return nil, fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
return "", fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
} else {
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
}
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
return "", errors.New(i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"))
}
// lockedRefresh attempts to refresh the token while holding dual-layer locks.
@@ -717,7 +697,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 := loadOAuthTokenUnderHeldLock(p.configDir, RuntimeProfile())
data, err := oauthLoadTokenLocked(p.configDir, RuntimeProfile())
if err != nil {
return nil, err
}
@@ -1,510 +0,0 @@
// Copyright 2026 Alibaba Group
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package auth
import (
"context"
"encoding/json"
"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,7 +14,6 @@
package auth
import (
"fmt"
"os"
"testing"
@@ -35,17 +34,7 @@ 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)
}
+11 -37
View File
@@ -127,10 +127,6 @@ 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
@@ -151,7 +147,6 @@ 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 {
@@ -164,27 +159,6 @@ 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 {
@@ -210,10 +184,13 @@ func DeleteTokenMarker(configDir string) error {
return nil
}
// 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.
// SaveTokenData persists TokenData. When an edition hook (SaveToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to the default keychain-based storage.
func SaveTokenData(configDir string, data *TokenData) error {
if h := edition.Get(); h.SaveToken != nil {
return saveTokenViaHook(h, configDir, data)
}
return withProfilesLock(configDir, func() error {
return saveTokenDataLocked(configDir, data)
})
@@ -487,8 +464,9 @@ func tokenLoadProfileIdentity(profile Profile) (*TokenData, error) {
return orgData, nil
}
// DeleteTokenData removes token data. Edition hooks and the default keychain
// path are both serialized with refresh through the auth dual lock.
// DeleteTokenData removes token data. When an edition hook (DeleteToken) is
// registered, it delegates entirely to the hook; otherwise it falls back
// to keychain + legacy cleanup.
func DeleteTokenData(configDir string) error {
return DeleteTokenDataForProfile(configDir, RuntimeProfile())
}
@@ -500,9 +478,7 @@ func DeleteTokenDataForProfile(configDir, profile string) error {
if strings.TrimSpace(profile) != "" {
return fmt.Errorf("profile selection is not supported by the current auth backend")
}
return withProfilesLock(configDir, func() error {
return h.DeleteToken(configDir)
})
return h.DeleteToken(configDir)
}
return withProfilesLock(configDir, func() error {
return deleteTokenDataForProfileLocked(configDir, profile)
@@ -937,9 +913,7 @@ 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 withProfilesLock(configDir, func() error {
return h.DeleteToken(configDir)
})
return h.DeleteToken(configDir)
}
return withProfilesLock(configDir, func() error {
var firstErr error
+9 -28
View File
@@ -63,10 +63,9 @@ func (i Identity) Key() string {
}
type Client struct {
BaseURL string
HTTPClient *http.Client
Identity Identity
AccessTokenProvider func(context.Context) (string, error)
BaseURL string
HTTPClient *http.Client
Identity Identity
}
type CreateSubscriptionRequest struct {
@@ -338,9 +337,8 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
if c == nil {
return errors.New("personal event: nil client")
}
accessToken, err := c.resolveAccessToken(ctx)
if err != nil {
return err
if c.Identity.AccessToken == "" {
return errors.New("personal event: access token is required")
}
u := strings.TrimRight(c.BaseURL, "/") + path
if len(q) > 0 {
@@ -360,7 +358,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, accessToken)
c.decorate(req)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
@@ -426,26 +424,9 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
return json.Unmarshal(data, out)
}
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)
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)
req.Header.Set("X-DWS-Client-Id", c.Identity.ClientID)
req.Header.Set("X-DWS-Source-Id", c.Identity.SourceID)
if c.Identity.CorpID != "" {
@@ -1,56 +0,0 @@
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)
}
}
+15 -39
View File
@@ -41,22 +41,19 @@ const (
)
type PersonalConfig struct {
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
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
}
type AccessTokenProvider func(context.Context) (string, error)
type PersonalSource struct {
cfg PersonalConfig
machine *Machine
@@ -76,8 +73,8 @@ type ticketResponse struct {
}
func NewPersonal(cfg PersonalConfig) (*PersonalSource, error) {
if cfg.AccessTokenProvider == nil && strings.TrimSpace(cfg.AccessToken) == "" {
return nil, errors.New("personal source: AccessToken or AccessTokenProvider is required")
if strings.TrimSpace(cfg.AccessToken) == "" {
return nil, errors.New("personal source: AccessToken is required")
}
if strings.TrimSpace(cfg.ClientID) == "" {
return nil, errors.New("personal source: ClientID is required")
@@ -195,10 +192,6 @@ 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,
@@ -214,8 +207,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", accessToken)
req.Header.Set("Authorization", "Bearer "+accessToken)
req.Header.Set("x-user-access-token", s.cfg.AccessToken)
req.Header.Set("Authorization", "Bearer "+s.cfg.AccessToken)
req.Header.Set("X-DWS-Client-Id", s.cfg.ClientID)
req.Header.Set("X-DWS-Source-Id", s.cfg.SourceID)
@@ -245,23 +238,6 @@ 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 {
+11 -16
View File
@@ -39,15 +39,14 @@ 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
AccessTokenProvider AccessTokenProvider
SourceID string
Mode string
ClientID string
ClientSecret string
UserAgent string
HTTPClient *http.Client
TicketURL string
AccessToken string
SourceID string
Mode string
ClientID string
ClientSecret string
UserAgent string
HTTPClient *http.Client
}
var portalWriteMessage = func(conn *websocket.Conn, messageType int, data []byte) error {
@@ -61,8 +60,8 @@ func (c *PortalTicketConfig) Valid() error {
if strings.TrimSpace(c.TicketURL) == "" {
return errors.New("source: portal ticket URL 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.AccessToken) == "" {
return errors.New("source: portal access token is required")
}
if strings.TrimSpace(c.SourceID) == "" {
return errors.New("source: portal sourceId is required")
@@ -162,10 +161,6 @@ 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}
@@ -189,7 +184,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", accessToken)
req.Header.Set("x-user-access-token", cfg.AccessToken)
resp, err := httpClient.Do(req)
if err != nil {
@@ -1,66 +0,0 @@
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,11 +36,6 @@ 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,7 +14,6 @@
package keychain
import (
"fmt"
"os"
"path/filepath"
"testing"
@@ -26,15 +25,8 @@ 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)
}
+4 -19
View File
@@ -16,7 +16,6 @@
package keychain
import (
"crypto/sha256"
"encoding/base64"
"errors"
"fmt"
@@ -48,8 +47,9 @@ 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, the storage root is taken from
// that env var; otherwise it defaults to %LocalAppData%\<service>.
// 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>.
func StorageDir(service string) string {
if override := os.Getenv(StorageDirEnv); override != "" {
return filepath.Join(override, service)
@@ -66,22 +66,7 @@ func StorageDir(service string) string {
}
func registryPathForService(service string) string {
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])
return regRootPath + `\` + safeRegistryComponent(service)
}
var safeRegRe = regexp.MustCompile(`[^a-zA-Z0-9._-]`)
@@ -17,37 +17,12 @@ 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")