fix(auth): unify access token resolution
This commit is contained in:
@@ -13,6 +13,7 @@ The format is inspired by [Keep a Changelog](https://keepachangelog.com/) and th
|
||||
|
||||
### Fixed
|
||||
|
||||
- **Consistent access-token caching and errors** — runtime, recovery, Skill, PAT polling, and personal/portal event clients now resolve user access tokens through one expiry- and publication-aware manager, so long-running processes reload rotated credentials while keychain, refresh, parse, permission, and cancellation failures remain observable instead of being collapsed into “not authenticated.”
|
||||
- **Tag-push GitHub Release publication** — Draft publication now locks one GitHub Release database ID, verifies its exact tag, channel, notes, recovery marker, asset set, and uploaded bytes, then publishes and rechecks that same ID as immutable. Recovery runs use the trusted default-branch release helpers instead of the sealed tag's historical scripts, fixing the Draft-only `GET /releases/tags/{tag}` 404 without allowing the release identity to drift during recovery.
|
||||
- **Release preflight reliability** — source-mode installer tests now use isolated temporary checkouts and HOME directories instead of overwriting and deleting the real repository `dws` binary, release preflight explicitly rebuilds before policy checks, and the full-suite runner gives the growing script package a non-flaky five-minute per-suite budget.
|
||||
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
type tokenManagerSnapshotProvider struct {
|
||||
load func() (*authpkg.TokenData, error)
|
||||
}
|
||||
|
||||
func (p tokenManagerSnapshotProvider) GetAccessToken(context.Context) (string, error) {
|
||||
data, err := p.load()
|
||||
if err != nil || data == nil {
|
||||
return "", err
|
||||
}
|
||||
return data.AccessToken, nil
|
||||
}
|
||||
|
||||
func (p tokenManagerSnapshotProvider) GetTokenSnapshot(context.Context) (*authpkg.TokenData, error) {
|
||||
return p.load()
|
||||
}
|
||||
|
||||
type tokenManagerLegacyGetter struct {
|
||||
token string
|
||||
err error
|
||||
}
|
||||
|
||||
func (g tokenManagerLegacyGetter) GetToken() (string, string, error) {
|
||||
return g.token, "file", g.err
|
||||
}
|
||||
|
||||
func installTokenManagerFakes(t *testing.T, load func() (*authpkg.TokenData, error)) {
|
||||
t.Helper()
|
||||
oldProvider, oldLegacy := newAccessTokenProvider, newLegacyTokenManager
|
||||
oldEdition := edition.Get()
|
||||
edition.Override(&edition.Hooks{})
|
||||
newAccessTokenProvider = func(string) accessTokenGetter {
|
||||
return tokenManagerSnapshotProvider{load: load}
|
||||
}
|
||||
newLegacyTokenManager = func(string) legacyTokenGetter {
|
||||
return tokenManagerLegacyGetter{err: authpkg.ErrTokenDataNotFound}
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
newAccessTokenProvider, newLegacyTokenManager = oldProvider, oldLegacy
|
||||
edition.Override(oldEdition)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerCachesUntilMarkerRevisionChanges(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
token := "token-a"
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: token, ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
|
||||
manager := NewTokenManager()
|
||||
first, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || first.AccessToken != "token-a" {
|
||||
t.Fatalf("first token = %#v, %v", first, err)
|
||||
}
|
||||
second, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || second.AccessToken != "token-a" || calls.Load() != 1 {
|
||||
t.Fatalf("cached token = %#v, %v, calls=%d", second, err, calls.Load())
|
||||
}
|
||||
|
||||
token = "token-b"
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
rotated, err := manager.Get(context.Background(), configDir, "")
|
||||
if err != nil || rotated.AccessToken != "token-b" || calls.Load() != 2 {
|
||||
t.Fatalf("rotated token = %#v, %v, calls=%d", rotated, err, calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerDoesNotCacheWithoutExpiryOrRevision(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token"}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerTreatsMalformedMarkerAsUncacheable(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := os.WriteFile(filepath.Join(configDir, "token.json"), []byte("{"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if snapshot, err := manager.Get(context.Background(), configDir, ""); err != nil || snapshot.AccessToken != "token" {
|
||||
t.Fatalf("snapshot = %#v, error = %v", snapshot, err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerDoesNotCacheOpaqueEditionStorageWithProviderFallback(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
edition.Override(&edition.Hooks{
|
||||
LoadToken: func(string) ([]byte, error) { return nil, nil },
|
||||
TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
|
||||
return fallback()
|
||||
},
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
for range 2 {
|
||||
if _, err := manager.Get(context.Background(), configDir, ""); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerCoalescesConcurrentLoads(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := authpkg.WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var calls atomic.Int32
|
||||
release := make(chan struct{})
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) {
|
||||
calls.Add(1)
|
||||
<-release
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
})
|
||||
manager := NewTokenManager()
|
||||
const workers = 8
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(workers)
|
||||
errs := make(chan error, workers)
|
||||
for range workers {
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_, err := manager.Get(context.Background(), configDir, "")
|
||||
errs <- err
|
||||
}()
|
||||
}
|
||||
for calls.Load() == 0 {
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
close(release)
|
||||
wg.Wait()
|
||||
close(errs)
|
||||
for err := range errs {
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls.Load() != 1 {
|
||||
t.Fatalf("provider calls = %d, want 1", calls.Load())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenManagerPreservesProviderFailure(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
want := errors.New("keychain permission denied")
|
||||
installTokenManagerFakes(t, func() (*authpkg.TokenData, error) { return nil, want })
|
||||
_, err := NewTokenManager().Get(context.Background(), configDir, "")
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("provider failure was misclassified as missing credentials: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenResolutionErrorOnlyClassifiesTrueMissingCredential(t *testing.T) {
|
||||
missing := tokenResolutionError(authpkg.ErrTokenDataNotFound)
|
||||
var typed interface{ Unwrap() error }
|
||||
if !errors.As(missing, &typed) || !errors.Is(missing, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("missing error = %v", missing)
|
||||
}
|
||||
want := errors.New("decrypt failed")
|
||||
if got := tokenResolutionError(want); !errors.Is(got, want) || errors.Is(got, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("storage error = %v", got)
|
||||
}
|
||||
if got := tokenResolutionError(context.Canceled); !errors.Is(got, context.Canceled) {
|
||||
t.Fatalf("cancellation = %v", got)
|
||||
}
|
||||
}
|
||||
@@ -21,19 +21,61 @@ import (
|
||||
"log/slog"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
authpkg "github.com/DingTalk-Real-AI/dingtalk-workspace-cli/internal/auth"
|
||||
"github.com/DingTalk-Real-AI/dingtalk-workspace-cli/pkg/edition"
|
||||
)
|
||||
|
||||
const accessTokenRefreshWindow = 5 * time.Minute
|
||||
|
||||
type legacyTokenGetter interface {
|
||||
GetToken() (string, string, error)
|
||||
}
|
||||
|
||||
type accessTokenSnapshotGetter interface {
|
||||
GetTokenSnapshot(context.Context) (*authpkg.TokenData, error)
|
||||
}
|
||||
|
||||
// AccessTokenSnapshot is the minimal bearer view needed by the process cache.
|
||||
// Refresh-token material never leaves the auth package.
|
||||
type AccessTokenSnapshot struct {
|
||||
AccessToken string
|
||||
ExpiresAt time.Time
|
||||
Source string
|
||||
}
|
||||
|
||||
type tokenManagerKey struct {
|
||||
configDir string
|
||||
profile string
|
||||
}
|
||||
|
||||
type tokenManagerEntry struct {
|
||||
mu sync.Mutex
|
||||
snapshot AccessTokenSnapshot
|
||||
revision string
|
||||
}
|
||||
|
||||
// TokenManager is the only process cache for user access tokens. Cache entries
|
||||
// are isolated by config directory and profile, expiry-aware, and invalidated
|
||||
// by the credential publication marker written by auth storage.
|
||||
type TokenManager struct {
|
||||
mu sync.Mutex
|
||||
entries map[tokenManagerKey]*tokenManagerEntry
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
func NewTokenManager() *TokenManager {
|
||||
return &TokenManager{entries: make(map[tokenManagerKey]*tokenManagerEntry), now: time.Now}
|
||||
}
|
||||
|
||||
var runtimeTokenManager = NewTokenManager()
|
||||
|
||||
var (
|
||||
newAccessTokenProvider = func(configDir string) accessTokenGetter {
|
||||
disc := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, disc)
|
||||
discard := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
provider := authpkg.NewOAuthProvider(configDir, discard)
|
||||
configureOAuthProviderCompatibility(provider, configDir)
|
||||
return provider
|
||||
}
|
||||
@@ -44,64 +86,218 @@ var (
|
||||
}
|
||||
)
|
||||
|
||||
// resolveAccessTokenFromDir loads OAuth then legacy token from configDir, applying
|
||||
// the same host compatibility hooks as MCP. It mirrors the former body of
|
||||
// getCachedRuntimeToken (excluding process-level cache and timing).
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
provider := newAccessTokenProvider(configDir)
|
||||
token, tokenErr := provider.GetAccessToken(ctx)
|
||||
if tokenErr == nil && strings.TrimSpace(token) != "" {
|
||||
return strings.TrimSpace(token), nil
|
||||
}
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
return "", tokenErr
|
||||
}
|
||||
if strings.TrimSpace(authpkg.RuntimeProfile()) != "" {
|
||||
if tokenErr != nil {
|
||||
return "", tokenErr
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
manager := newLegacyTokenManager(configDir)
|
||||
if leg, _, err := manager.GetToken(); err == nil && strings.TrimSpace(leg) != "" {
|
||||
return strings.TrimSpace(leg), nil
|
||||
}
|
||||
if tokenErr != nil {
|
||||
return "", tokenErr
|
||||
}
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves a bearer token for HTTP clients that should
|
||||
// align with MCP tool calls. Non-empty explicitToken wins. When configDir matches
|
||||
// the active edition config directory, the same process-cached path as MCP is used.
|
||||
// Otherwise tokens are loaded from configDir with host compatibility hooks applied.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
if t := strings.TrimSpace(explicitToken); t != "" {
|
||||
return t, nil
|
||||
// Get resolves an access token for the active runtime profile.
|
||||
func (m *TokenManager) Get(ctx context.Context, configDir, explicitToken string) (AccessTokenSnapshot, error) {
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return AccessTokenSnapshot{AccessToken: token, Source: "explicit"}, nil
|
||||
}
|
||||
if strings.TrimSpace(configDir) == "" {
|
||||
return "", fmt.Errorf("config directory is empty")
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("config directory is empty")
|
||||
}
|
||||
if filepath.Clean(configDir) == filepath.Clean(defaultConfigDir()) {
|
||||
if tok := resolveRuntimeAuthToken(ctx, ""); tok != "" {
|
||||
return tok, nil
|
||||
key := tokenManagerKey{
|
||||
configDir: canonicalTokenConfigDir(configDir),
|
||||
profile: strings.TrimSpace(authpkg.RuntimeProfile()),
|
||||
}
|
||||
entry := m.entry(key)
|
||||
entry.mu.Lock()
|
||||
defer entry.mu.Unlock()
|
||||
|
||||
now := time.Now()
|
||||
if m != nil && m.now != nil {
|
||||
now = m.now()
|
||||
}
|
||||
revision, present, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if tokenSnapshotUsable(entry.snapshot, now) && present && revision != "" && revision == entry.revision {
|
||||
return entry.snapshot, nil
|
||||
}
|
||||
|
||||
// Treat the marker and credential as one optimistic snapshot. A concurrent
|
||||
// login/refresh between the reads causes a retry instead of caching stale A
|
||||
// under the publication marker for B.
|
||||
for attempt := 0; attempt < 4; attempt++ {
|
||||
beforeRevision, beforePresent, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
snapshot, err := resolveTokenSnapshotWithEdition(ctx, configDir, key.profile)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
afterRevision, afterPresent, err := authpkg.ReadTokenMarkerRevision(configDir)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if beforePresent != afterPresent || beforeRevision != afterRevision {
|
||||
continue
|
||||
}
|
||||
if strings.TrimSpace(snapshot.AccessToken) == "" {
|
||||
return AccessTokenSnapshot{}, noCredentialsError()
|
||||
}
|
||||
if tokenSnapshotUsable(snapshot, now) && afterPresent && afterRevision != "" {
|
||||
entry.snapshot = snapshot
|
||||
entry.revision = afterRevision
|
||||
} else {
|
||||
entry.snapshot = AccessTokenSnapshot{}
|
||||
entry.revision = ""
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
tok, err := resolveAccessTokenFromDir(ctx, configDir)
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("token publication changed repeatedly while resolving credentials")
|
||||
}
|
||||
|
||||
func (m *TokenManager) entry(key tokenManagerKey) *tokenManagerEntry {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if m.entries == nil {
|
||||
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
|
||||
}
|
||||
entry := m.entries[key]
|
||||
if entry == nil {
|
||||
entry = &tokenManagerEntry{}
|
||||
m.entries[key] = entry
|
||||
}
|
||||
return entry
|
||||
}
|
||||
|
||||
func (m *TokenManager) Invalidate() {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
m.mu.Lock()
|
||||
m.entries = make(map[tokenManagerKey]*tokenManagerEntry)
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
func resolveTokenSnapshotWithEdition(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
|
||||
hooks := edition.Get()
|
||||
opaqueStorage := hooks.LoadToken != nil || hooks.SaveToken != nil || hooks.DeleteToken != nil
|
||||
provider := hooks.TokenProvider
|
||||
if provider == nil {
|
||||
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
// Opaque edition storage hooks have no publication-revision contract.
|
||||
// Resolve them on every logical request instead of caching a token that
|
||||
// may be replaced outside the default auth store.
|
||||
if opaqueStorage {
|
||||
snapshot.ExpiresAt = time.Time{}
|
||||
}
|
||||
return snapshot, nil
|
||||
}
|
||||
var fallbackSnapshot AccessTokenSnapshot
|
||||
var fallbackCalled bool
|
||||
token, err := provider(ctx, func() (string, error) {
|
||||
fallbackCalled = true
|
||||
var fallbackErr error
|
||||
fallbackSnapshot, fallbackErr = resolveAccessTokenSnapshotFromDir(ctx, configDir, profile)
|
||||
if fallbackErr != nil {
|
||||
return "", fallbackErr
|
||||
}
|
||||
return fallbackSnapshot.AccessToken, nil
|
||||
})
|
||||
if err != nil {
|
||||
return AccessTokenSnapshot{}, fmt.Errorf("edition token provider: %w", err)
|
||||
}
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
return AccessTokenSnapshot{}, noCredentialsError()
|
||||
}
|
||||
if fallbackCalled && token == fallbackSnapshot.AccessToken {
|
||||
if opaqueStorage {
|
||||
fallbackSnapshot.ExpiresAt = time.Time{}
|
||||
}
|
||||
return fallbackSnapshot, nil
|
||||
}
|
||||
// Edition providers expose no lifetime metadata, so resolve them on every
|
||||
// logical request instead of recreating a process-lifetime string cache.
|
||||
return AccessTokenSnapshot{AccessToken: token, Source: "edition"}, nil
|
||||
}
|
||||
|
||||
func resolveAccessTokenSnapshotFromDir(ctx context.Context, configDir, profile string) (AccessTokenSnapshot, error) {
|
||||
provider := newAccessTokenProvider(configDir)
|
||||
if snapshotProvider, ok := provider.(accessTokenSnapshotGetter); ok {
|
||||
data, err := snapshotProvider.GetTokenSnapshot(ctx)
|
||||
if err == nil && data != nil && strings.TrimSpace(data.AccessToken) != "" {
|
||||
return AccessTokenSnapshot{
|
||||
AccessToken: strings.TrimSpace(data.AccessToken),
|
||||
ExpiresAt: data.ExpiresAt,
|
||||
Source: "oauth",
|
||||
}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
return resolveLegacyToken(configDir, err)
|
||||
}
|
||||
|
||||
token, err := provider.GetAccessToken(ctx)
|
||||
if err == nil && strings.TrimSpace(token) != "" {
|
||||
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: "oauth_compat"}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if strings.TrimSpace(profile) != "" {
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
return resolveLegacyToken(configDir, err)
|
||||
}
|
||||
|
||||
func resolveLegacyToken(configDir string, oauthErr error) (AccessTokenSnapshot, error) {
|
||||
token, source, err := newLegacyTokenManager(configDir).GetToken()
|
||||
if err == nil && strings.TrimSpace(token) != "" {
|
||||
return AccessTokenSnapshot{AccessToken: strings.TrimSpace(token), Source: source}, nil
|
||||
}
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return AccessTokenSnapshot{}, err
|
||||
}
|
||||
if oauthErr != nil {
|
||||
return AccessTokenSnapshot{}, oauthErr
|
||||
}
|
||||
return AccessTokenSnapshot{}, authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
|
||||
func resolveAccessTokenFromDir(ctx context.Context, configDir string) (string, error) {
|
||||
snapshot, err := resolveAccessTokenSnapshotFromDir(ctx, configDir, authpkg.RuntimeProfile())
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if tok != "" {
|
||||
return tok, nil
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
// ResolveAuxiliaryAccessToken resolves every non-runner bearer token through
|
||||
// the same TokenManager used by MCP tool calls.
|
||||
func ResolveAuxiliaryAccessToken(ctx context.Context, configDir, explicitToken string) (string, error) {
|
||||
snapshot, err := runtimeTokenManager.Get(ctx, configDir, explicitToken)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return "", noCredentialsError()
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
func tokenSnapshotUsable(snapshot AccessTokenSnapshot, now time.Time) bool {
|
||||
return strings.TrimSpace(snapshot.AccessToken) != "" &&
|
||||
!snapshot.ExpiresAt.IsZero() &&
|
||||
now.Before(snapshot.ExpiresAt.Add(-accessTokenRefreshWindow))
|
||||
}
|
||||
|
||||
func canonicalTokenConfigDir(configDir string) string {
|
||||
if absolute, err := filepath.Abs(configDir); err == nil {
|
||||
return filepath.Clean(absolute)
|
||||
}
|
||||
return filepath.Clean(configDir)
|
||||
}
|
||||
|
||||
func noCredentialsError() error {
|
||||
if edition.Get().IsEmbedded {
|
||||
return fmt.Errorf("认证信息已失效,请重新认证")
|
||||
return fmt.Errorf("认证信息已失效,请重新认证: %w", authpkg.ErrTokenDataNotFound)
|
||||
}
|
||||
return fmt.Errorf("no credentials found, run: dws auth login")
|
||||
return fmt.Errorf("no credentials found, run: dws auth login: %w", authpkg.ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
@@ -212,14 +212,16 @@ func TestCrossPlatformCoverageConfigAndTokenSeamsCoverage(t *testing.T) {
|
||||
if _, err := resolveAccessTokenFromDir(context.Background(), "unused"); !errors.Is(err, authpkg.ErrTokenDecryption) {
|
||||
t.Fatalf("decryption error = %v", err)
|
||||
}
|
||||
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{err: errors.New("missing")} }
|
||||
newAccessTokenProvider = func(string) accessTokenGetter {
|
||||
return fakeAccessTokenGetter{err: authpkg.ErrTokenDataNotFound}
|
||||
}
|
||||
newLegacyTokenManager = func(string) legacyTokenGetter { return fakeLegacyTokenGetter{token: " legacy "} }
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); err != nil || got != "legacy" {
|
||||
t.Fatalf("legacy token = %q, %v", got, err)
|
||||
}
|
||||
authpkg.SetRuntimeProfile("corp:user")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || err == nil || err.Error() != "missing" {
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), "unused"); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("explicit profile fallback = token %q error %v, want profile error", got, err)
|
||||
}
|
||||
authpkg.SetRuntimeProfile("")
|
||||
|
||||
@@ -554,7 +554,7 @@ func TestCrossPlatformCoverageRecoveryRuntimeHTTP(t *testing.T) {
|
||||
defer server.Close()
|
||||
SetDynamicServers([]mcptypes.ServerDescriptor{{Endpoint: server.URL, CLI: mcptypes.CLIOverlay{ID: "devdoc", Tools: []mcptypes.CLITool{{Name: "search_open_platform_docs_rag"}}}}})
|
||||
t.Cleanup(func() { SetDynamicServers(nil) })
|
||||
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client())}
|
||||
runtime := &recoveryRuntime{transport: transport.NewClient(server.Client()), flags: &GlobalFlags{Token: "token"}}
|
||||
got, err := runtime.Search(context.Background(), "query", recovery.RecoveryContext{ToolName: "search"})
|
||||
if err != nil || got.DocSearch.Status != "success" || len(got.KBHits) == 0 {
|
||||
t.Fatalf("recovery search = %#v %v", got, err)
|
||||
@@ -1650,6 +1650,8 @@ func TestCrossPlatformCoveragePersonalSubscriptionAndSourceCoverage(t *testing.T
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalEventCommandRuntimeCoverage(t *testing.T) {
|
||||
authpkg.SetRuntimeProfile("")
|
||||
t.Cleanup(func() { authpkg.SetRuntimeProfile("") })
|
||||
configDir := setupPersonalIdentityToken(t, &authpkg.TokenData{
|
||||
AccessToken: "access", RefreshToken: "refresh", ExpiresAt: time.Now().Add(time.Hour),
|
||||
CorpID: "corp", UserID: "user", ClientID: "client",
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
@@ -164,7 +165,7 @@ func TestCrossPlatformCoverageRawAPIAndTokenCoverage(t *testing.T) {
|
||||
}
|
||||
newAccessTokenProvider = func(string) accessTokenGetter { return fakeAccessTokenGetter{} }
|
||||
missing := t.TempDir()
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), missing); err != nil || got != "" {
|
||||
if got, err := resolveAccessTokenFromDir(context.Background(), missing); got != "" || !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
t.Fatalf("missing access token = %q, %v", got, err)
|
||||
}
|
||||
if _, err := ResolveAuxiliaryAccessToken(context.Background(), missing, ""); err == nil {
|
||||
|
||||
@@ -413,7 +413,7 @@ func eventStreamBusID(streamOpts eventStreamTicketOptions) string {
|
||||
return "portal-ticket-normal:" + sourceID
|
||||
}
|
||||
|
||||
func newEventSource(ctx context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
|
||||
func newEventSource(_ context.Context, configDir, clientID, clientSecret string, streamOpts eventStreamTicketOptions) (*source.DingtalkSource, error) {
|
||||
if !streamOpts.enabled() {
|
||||
return eventNewDingtalkSource(source.Config{
|
||||
ClientID: clientID,
|
||||
@@ -421,14 +421,6 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
|
||||
})
|
||||
}
|
||||
|
||||
token, err := eventResolveAccessToken(ctx, configDir, "")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("event stream ticket: resolve user token: %w", err)
|
||||
}
|
||||
if strings.TrimSpace(token) == "" {
|
||||
return nil, errors.New("event stream ticket: empty user token")
|
||||
}
|
||||
|
||||
portalClientID := clientID
|
||||
portalClientSecret := clientSecret
|
||||
if streamOpts.usesPortalNormalMode() {
|
||||
@@ -440,8 +432,10 @@ func newEventSource(ctx context.Context, configDir, clientID, clientSecret strin
|
||||
ClientID: portalClientID,
|
||||
ClientSecret: portalClientSecret,
|
||||
PortalTicket: &source.PortalTicketConfig{
|
||||
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
|
||||
AccessToken: token,
|
||||
TicketURL: eventStreamTicketURL(streamOpts.TicketURL),
|
||||
AccessTokenProvider: func(ctx context.Context) (string, error) {
|
||||
return eventResolveAccessToken(ctx, configDir, "")
|
||||
},
|
||||
SourceID: eventStreamSourceID(streamOpts.SourceID),
|
||||
Mode: streamOpts.Mode,
|
||||
ClientID: portalClientID,
|
||||
|
||||
@@ -132,14 +132,18 @@ func TestCrossPlatformCoverageEventSourcesAndForegroundCoverage(t *testing.T) {
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", eventStreamTicketOptions{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
|
||||
stream := eventStreamTicketOptions{Mode: "custom"}
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); !errors.Is(err, fail) {
|
||||
t.Fatalf("stream token error = %v", err)
|
||||
var captured source.Config
|
||||
eventNewDingtalkSource = func(cfg source.Config, _ ...source.SourceOption) (*source.DingtalkSource, error) {
|
||||
captured = cfg
|
||||
return &source.DingtalkSource{}, nil
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return " ", nil }
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err == nil {
|
||||
t.Fatal("empty stream token succeeded")
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "", fail }
|
||||
if _, err := newEventSource(context.Background(), "config", "client", "secret", stream); err != nil {
|
||||
t.Fatalf("stream source construction = %v", err)
|
||||
}
|
||||
if _, err := captured.PortalTicket.AccessTokenProvider(context.Background()); !errors.Is(err, fail) {
|
||||
t.Fatalf("stream token provider error = %v", err)
|
||||
}
|
||||
eventResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
|
||||
for _, mode := range []string{"custom", "normal"} {
|
||||
|
||||
@@ -259,7 +259,7 @@ func runPersonalEventConsume(c *cobra.Command, opts personalConsumeOptions) erro
|
||||
return personalConsumeRun(ctx, cfg)
|
||||
}
|
||||
|
||||
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
sub, eventKey, ruleType, err := personalEnsureSubscription(ctx, client, identity, opts)
|
||||
if err != nil {
|
||||
return fmt.Errorf("event consume --as user: %w", err)
|
||||
@@ -498,7 +498,7 @@ func runPersonalEventStatus(c *cobra.Command, opts personalStatusOptions) error
|
||||
if status == "" || status == "all" {
|
||||
status = ""
|
||||
}
|
||||
subs, err := personalListSubscriptions(personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
|
||||
subs, err := personalListSubscriptions(newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity), ctx, personal.ListOptions{
|
||||
Status: status,
|
||||
EventKey: opts.EventKey,
|
||||
SubscribeID: opts.SubscribeID,
|
||||
@@ -613,7 +613,7 @@ func runPersonalEventStop(c *cobra.Command, opts personalStopOptions) error {
|
||||
if err != nil {
|
||||
return fmt.Errorf("event stop --as user: %w", err)
|
||||
}
|
||||
client := personal.NewClient(personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
client := newPersonalEventControlClient(configDir, personalEventControlBaseURL(opts.ControlBaseURL, configDir), identity)
|
||||
for _, id := range subscribeIDs {
|
||||
if err := personalDeleteSubscription(client, ctx, id); err != nil {
|
||||
return fmt.Errorf("event stop --as user: cancel subscription %s: %w", id, err)
|
||||
@@ -723,7 +723,10 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
|
||||
if err != nil {
|
||||
return personal.Identity{}, err
|
||||
}
|
||||
tokenData, _ := personalLoadTokenData(configDir)
|
||||
tokenData, err := personalLoadTokenData(configDir)
|
||||
if err != nil && !errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return personal.Identity{}, fmt.Errorf("load OAuth identity metadata: %w", err)
|
||||
}
|
||||
var corpID, userID, clientID, refreshToken string
|
||||
if tokenData != nil {
|
||||
corpID = tokenData.CorpID
|
||||
@@ -769,6 +772,15 @@ func resolvePersonalEventIdentity(ctx context.Context, configDir string, sourceI
|
||||
}, nil
|
||||
}
|
||||
|
||||
func newPersonalEventControlClient(configDir, baseURL string, identity personal.Identity) *personal.Client {
|
||||
identity.AccessToken = ""
|
||||
client := personal.NewClient(baseURL, identity)
|
||||
client.AccessTokenProvider = func(ctx context.Context) (string, error) {
|
||||
return personalResolveAuxiliaryAccessToken(ctx, configDir, "")
|
||||
}
|
||||
return client
|
||||
}
|
||||
|
||||
func personalTokenSubject(kind, token string) string {
|
||||
token = strings.TrimSpace(token)
|
||||
if token == "" {
|
||||
@@ -817,7 +829,9 @@ func newPersonalStreamSource(ctx context.Context, opts personalStreamSourceOptio
|
||||
}
|
||||
_ = ctx
|
||||
return source.NewPersonal(source.PersonalConfig{
|
||||
AccessToken: opts.Identity.AccessToken,
|
||||
AccessTokenProvider: func(ctx context.Context) (string, error) {
|
||||
return personalResolveAuxiliaryAccessToken(ctx, opts.ConfigDir, "")
|
||||
},
|
||||
ClientID: clientID,
|
||||
ClientSecret: clientSecret,
|
||||
SourceID: opts.Identity.SourceID,
|
||||
|
||||
@@ -56,7 +56,7 @@ var openBrowserFunc = tryOpenBrowser
|
||||
var (
|
||||
patAuthorizationTimeout = PatAuthRetryTimeout
|
||||
patAuthorizationPollInterval = PatAuthPollInterval
|
||||
patLoadTokenData = authpkg.LoadTokenData
|
||||
patResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
patWaitForAuthorization = WaitForPatAuthorization
|
||||
patPollDeviceFlowWithInterval = pollPatDeviceFlowWithInterval
|
||||
patSaveAppConfig = authpkg.SaveAppConfig
|
||||
@@ -272,7 +272,7 @@ func patAuthorizationURIFromData(data map[string]any) string {
|
||||
|
||||
// WaitForPatAuthorization polls until the user completes authorization or timeout.
|
||||
// It returns true if authorization was completed, false if timed out or cancelled.
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) bool {
|
||||
func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Writer) (bool, error) {
|
||||
timeout := patAuthorizationTimeout
|
||||
deadline := time.Now().Add(timeout)
|
||||
pollTicker := time.NewTicker(patAuthorizationPollInterval)
|
||||
@@ -290,27 +290,26 @@ func WaitForPatAuthorization(ctx context.Context, configDir string, output io.Wr
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fmt.Fprintf(output, "%s 操作已取消\n", tui.StateMark("error"))
|
||||
return false
|
||||
return false, ctx.Err()
|
||||
|
||||
case <-time.After(time.Until(deadline)):
|
||||
fmt.Fprintf(output, "%s 等待授权超时 (%s)\n", tui.StateMark("error"), timeout)
|
||||
fmt.Fprintf(output, " %s 请重新执行命令\n", tui.Dim("ℹ"))
|
||||
return false
|
||||
return false, nil
|
||||
|
||||
case <-pollTicker.C:
|
||||
pollCount++
|
||||
elapsed := time.Since(start).Truncate(time.Second)
|
||||
remaining := time.Until(deadline).Truncate(time.Second)
|
||||
|
||||
// Check if token is now valid
|
||||
tokenData, err := patLoadTokenData(configDir)
|
||||
if err == nil && tokenData != nil {
|
||||
if tokenData.IsAccessTokenValid() || tokenData.IsRefreshTokenValid() {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true
|
||||
}
|
||||
// Check the same resolver used by every outbound bearer request.
|
||||
if _, err := patResolveAccessToken(ctx, configDir, ""); err == nil {
|
||||
fmt.Fprintf(output, "\r%s %s (%s 已用, %s 剩余) \n",
|
||||
tui.StateMark("ok"), tui.Bold("授权成功!"), elapsed, remaining)
|
||||
fmt.Fprintln(output)
|
||||
return true, nil
|
||||
} else if !stderrors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return false, fmt.Errorf("check authorization token: %w", err)
|
||||
}
|
||||
|
||||
// Show polling status
|
||||
@@ -340,7 +339,10 @@ func retryWithPatAuthRetry(ctx context.Context, runner executor.Runner, invocati
|
||||
PrintPatAuthError(output, scopeErr)
|
||||
|
||||
// Wait for user to complete authorization
|
||||
authorized := patWaitForAuthorization(ctx, configDir, output)
|
||||
authorized, waitErr := patWaitForAuthorization(ctx, configDir, output)
|
||||
if waitErr != nil {
|
||||
return executor.Result{}, waitErr
|
||||
}
|
||||
if !authorized {
|
||||
return executor.Result{}, apperrors.NewAuth(
|
||||
"等待用户授权超时",
|
||||
@@ -794,12 +796,6 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
|
||||
pollURL := fmt.Sprintf("%s%s?flowId=%s",
|
||||
authpkg.GetMCPBaseURL(), authpkg.DevicePollPath, url.QueryEscape(flowID))
|
||||
|
||||
// Load user access token for the poll request header.
|
||||
var accessToken string
|
||||
if tokenData, err := authpkg.LoadTokenData(configDir); err == nil && tokenData != nil {
|
||||
accessToken = tokenData.AccessToken
|
||||
}
|
||||
|
||||
// Use a client that does NOT follow redirects, so we can detect SSO 302.
|
||||
noRedirectClient := &http.Client{
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
@@ -828,6 +824,10 @@ func pollPatDeviceFlowWithInterval(ctx context.Context, flowID string, configDir
|
||||
slog.Debug("PAT poll: failed to create request", "error", err)
|
||||
continue
|
||||
}
|
||||
accessToken, tokenErr := patResolveAccessToken(ctx, configDir, "")
|
||||
if tokenErr != nil && !stderrors.Is(tokenErr, authpkg.ErrTokenDataNotFound) {
|
||||
return "", "", fmt.Errorf("resolve PAT poll access token: %w", tokenErr)
|
||||
}
|
||||
if accessToken != "" {
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
}
|
||||
|
||||
@@ -52,39 +52,41 @@ func TestCrossPlatformCoveragePATRetryRemainingPureAndWaitCoverage(t *testing.T)
|
||||
|
||||
oldTimeout := patAuthorizationTimeout
|
||||
oldInterval := patAuthorizationPollInterval
|
||||
oldLoad := patLoadTokenData
|
||||
oldResolve := patResolveAccessToken
|
||||
t.Cleanup(func() {
|
||||
patAuthorizationTimeout = oldTimeout
|
||||
patAuthorizationPollInterval = oldInterval
|
||||
patLoadTokenData = oldLoad
|
||||
patResolveAccessToken = oldResolve
|
||||
})
|
||||
patAuthorizationTimeout = 50 * time.Millisecond
|
||||
patAuthorizationPollInterval = time.Millisecond
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{AccessToken: "token", ExpiresAt: time.Now().Add(time.Hour)}, nil
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "token", nil
|
||||
}
|
||||
out.Reset()
|
||||
if !WaitForPatAuthorization(context.Background(), "", &out) {
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || !ok {
|
||||
t.Fatal("valid token did not authorize")
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(ctx, "", &out) {
|
||||
t.Fatal("cancelled authorization succeeded")
|
||||
if ok, err := WaitForPatAuthorization(ctx, "", &out); ok || !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("cancelled authorization = %v, %v", ok, err)
|
||||
}
|
||||
patAuthorizationTimeout = time.Millisecond
|
||||
patAuthorizationPollInterval = time.Hour
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(context.Background(), "", &out) {
|
||||
t.Fatal("timed out authorization succeeded")
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok {
|
||||
t.Fatalf("timed out authorization = %v, %v", ok, err)
|
||||
}
|
||||
patAuthorizationTimeout = 5 * time.Millisecond
|
||||
patAuthorizationPollInterval = time.Millisecond
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
out.Reset()
|
||||
if WaitForPatAuthorization(context.Background(), "", &out) || !strings.Contains(out.String(), "等待授权中") {
|
||||
t.Fatalf("invalid-token polling output = %q", out.String())
|
||||
if ok, err := WaitForPatAuthorization(context.Background(), "", &out); err != nil || ok || !strings.Contains(out.String(), "等待授权中") {
|
||||
t.Fatalf("invalid-token polling = %v, %v, output %q", ok, err, out.String())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -109,12 +111,12 @@ func TestCrossPlatformCoveragePATRetryRemainingOrchestrationCoverage(t *testing.
|
||||
t.Setenv("DWS_CONFIG_DIR", t.TempDir())
|
||||
|
||||
scope := &PatScopeError{OriginalError: "missing", Identity: "user", ErrorType: "missing_scope", Message: "missing", Hint: "login", MissingScope: "calendar:read"}
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return false }
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return false, nil }
|
||||
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{}, executor.Invocation{}, scope, t.TempDir(), io.Discard); err == nil {
|
||||
t.Fatal("PAT retry timeout succeeded")
|
||||
}
|
||||
wantErr := errors.New("runner failed")
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) bool { return true }
|
||||
patWaitForAuthorization = func(context.Context, string, io.Writer) (bool, error) { return true, nil }
|
||||
if _, err := retryWithPatAuthRetry(context.Background(), runnerCoverageFallback{err: wantErr}, executor.Invocation{}, scope, t.TempDir(), io.Discard); !errors.Is(err, wantErr) {
|
||||
t.Fatalf("authorized retry = %v", err)
|
||||
}
|
||||
@@ -215,15 +217,15 @@ func patRaw(flowID, clientID, secret string) string {
|
||||
func TestCrossPlatformCoveragePATRetryRemainingPollAndBrowserCoverage(t *testing.T) {
|
||||
oldDo := patPollHTTPDo
|
||||
oldRequest := patPollNewRequest
|
||||
oldLoad := patLoadTokenData
|
||||
oldResolve := patResolveAccessToken
|
||||
oldBrowser := patBrowserOpenCommand
|
||||
t.Cleanup(func() {
|
||||
patPollHTTPDo = oldDo
|
||||
patPollNewRequest = oldRequest
|
||||
patLoadTokenData = oldLoad
|
||||
patResolveAccessToken = oldResolve
|
||||
patBrowserOpenCommand = oldBrowser
|
||||
})
|
||||
patLoadTokenData = func(string) (*authpkg.TokenData, error) { return &authpkg.TokenData{AccessToken: "token"}, nil }
|
||||
patResolveAccessToken = func(context.Context, string, string) (string, error) { return "token", nil }
|
||||
cancelled, cancelNow := context.WithCancel(context.Background())
|
||||
cancelNow()
|
||||
if status, _, err := pollPatDeviceFlowWithInterval(cancelled, "flow", t.TempDir(), io.Discard, 0); err != nil || status != authpkg.StatusCancelled {
|
||||
|
||||
@@ -333,7 +333,11 @@ func (r *recoveryRuntime) CallToolDirect(ctx context.Context, serverID, toolName
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tc := r.transport.WithAuth(resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags)), resolveIdentityHeaders())
|
||||
authToken, err := resolveRuntimeAuthToken(ctx, recoveryRuntimeToken(r.flags))
|
||||
if err != nil {
|
||||
return nil, tokenResolutionError(err)
|
||||
}
|
||||
tc := r.transport.WithAuth(authToken, resolveIdentityHeaders())
|
||||
result, err := tc.CallTool(ctx, endpoint, toolName, args)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
+39
-53
@@ -235,7 +235,9 @@ func (r *runtimeRunner) runSingle(ctx context.Context, invocation executor.Invoc
|
||||
// ~70ms on macOS; starting it here lets the load overlap with endpoint
|
||||
// resolution and catalog loading below.
|
||||
if prefetchToken {
|
||||
go runnerGetCachedRuntimeToken(ctx)
|
||||
go func() {
|
||||
_, _ = runnerGetCachedRuntimeToken(ctx)
|
||||
}()
|
||||
}
|
||||
|
||||
if shouldUseDirectRuntime(invocation) {
|
||||
@@ -534,8 +536,12 @@ func (r *runtimeRunner) executeInvocation(ctx context.Context, endpoint string,
|
||||
authToken := ""
|
||||
if hasPluginAuth {
|
||||
authToken = pluginAuth.Token
|
||||
} else {
|
||||
authToken = r.resolveAuthToken(ctx)
|
||||
} else if !invocation.DryRun && (r.globalFlags == nil || !r.globalFlags.Mock) {
|
||||
var tokenErr error
|
||||
authToken, tokenErr = r.resolveAuthToken(ctx)
|
||||
if tokenErr != nil {
|
||||
return executor.Result{}, tokenResolutionError(tokenErr)
|
||||
}
|
||||
}
|
||||
|
||||
var timeoutSec int
|
||||
@@ -796,67 +802,49 @@ func (r *runtimeRunner) executeStdioInvocation(ctx context.Context, invocation e
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) string {
|
||||
func (r *runtimeRunner) resolveAuthToken(ctx context.Context) (string, error) {
|
||||
explicitToken := ""
|
||||
if r != nil && r.globalFlags != nil {
|
||||
explicitToken = r.globalFlags.Token
|
||||
}
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
}
|
||||
if tp := edition.Get().TokenProvider; tp != nil {
|
||||
token, _ := tp(ctx, func() (string, error) {
|
||||
return resolveAccessTokenFromDir(ctx, defaultConfigDir())
|
||||
})
|
||||
return token
|
||||
}
|
||||
return getCachedRuntimeToken(ctx)
|
||||
return resolveRuntimeAuthToken(ctx, explicitToken)
|
||||
}
|
||||
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) string {
|
||||
if token := strings.TrimSpace(explicitToken); token != "" {
|
||||
return token
|
||||
func resolveRuntimeAuthToken(ctx context.Context, explicitToken string) (string, error) {
|
||||
snapshot, err := runtimeTokenManager.Get(ctx, defaultConfigDir(), explicitToken)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// Use cached token to avoid repeated Keychain access (~70ms per call)
|
||||
return getCachedRuntimeToken(ctx)
|
||||
return snapshot.AccessToken, nil
|
||||
}
|
||||
|
||||
// Cached token state for process lifetime
|
||||
var (
|
||||
cachedRuntimeTokenMu sync.Mutex
|
||||
cachedRuntimeTokens = map[string]string{}
|
||||
)
|
||||
|
||||
// getCachedRuntimeToken returns a cached access token, loading it only once per process.
|
||||
// This avoids repeated Keychain access which takes ~70ms each time.
|
||||
func getCachedRuntimeToken(ctx context.Context) string {
|
||||
cacheKey := strings.TrimSpace(authpkg.RuntimeProfile())
|
||||
if cacheKey == "" {
|
||||
cacheKey = "__default__"
|
||||
}
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
if token := cachedRuntimeTokens[cacheKey]; token != "" {
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
return token
|
||||
}
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
|
||||
// getCachedRuntimeToken is kept as the prefetch seam used by runner tests. The
|
||||
// cache itself lives exclusively in TokenManager.
|
||||
func getCachedRuntimeToken(ctx context.Context) (string, error) {
|
||||
loadStart := time.Now()
|
||||
defer func() { RecordTiming(ctx, "auth_keychain", time.Since(loadStart)) }()
|
||||
return resolveRuntimeAuthToken(ctx, "")
|
||||
}
|
||||
|
||||
configDir := defaultConfigDir()
|
||||
token, tokenErr := resolveAccessTokenFromDir(ctx, configDir)
|
||||
if tokenErr != nil && errors.Is(tokenErr, authpkg.ErrTokenDecryption) {
|
||||
slog.Error(tokenErr.Error())
|
||||
return ""
|
||||
func tokenResolutionError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if token == "" {
|
||||
return ""
|
||||
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
||||
return err
|
||||
}
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
cachedRuntimeTokens[cacheKey] = token
|
||||
cachedRuntimeTokenMu.Unlock()
|
||||
return token
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return apperrors.NewAuth(
|
||||
"未登录,请先执行 dws auth login",
|
||||
apperrors.WithReason("not_authenticated"),
|
||||
apperrors.WithHint("运行 'dws auth login' 完成登录后重试"),
|
||||
apperrors.WithActions("dws auth login"),
|
||||
apperrors.WithCause(err),
|
||||
)
|
||||
}
|
||||
// Keychain, parse, permission, lock, and refresh failures are real local or
|
||||
// network errors. Preserve their cause instead of disguising them as logout.
|
||||
return fmt.Errorf("resolve access token: %w", err)
|
||||
}
|
||||
|
||||
// generateExecutionID returns a random 16-char hex string used to correlate
|
||||
@@ -871,9 +859,7 @@ func generateExecutionID() string {
|
||||
// ResetRuntimeTokenCache clears the cached token, forcing a reload on next access.
|
||||
// This should be called after login/logout operations.
|
||||
func ResetRuntimeTokenCache() {
|
||||
cachedRuntimeTokenMu.Lock()
|
||||
defer cachedRuntimeTokenMu.Unlock()
|
||||
cachedRuntimeTokens = map[string]string{}
|
||||
runtimeTokenManager.Invalidate()
|
||||
}
|
||||
|
||||
func newRuntimeContentScanner() safety.Scanner {
|
||||
|
||||
@@ -50,9 +50,9 @@ func TestCrossPlatformCoverageRunnerRemainingRoutingCoverage(t *testing.T) {
|
||||
|
||||
inv := executor.Invocation{CanonicalProduct: "product", Tool: "tool"}
|
||||
prefetched := make(chan struct{}, 1)
|
||||
runnerGetCachedRuntimeToken = func(context.Context) string {
|
||||
runnerGetCachedRuntimeToken = func(context.Context) (string, error) {
|
||||
prefetched <- struct{}{}
|
||||
return ""
|
||||
return "", nil
|
||||
}
|
||||
r := &runtimeRunner{
|
||||
loader: cli.CatalogLoaderFrom(cli.Catalog{}, wantErr),
|
||||
@@ -329,19 +329,19 @@ func TestCrossPlatformCoverageRunnerRemainingStdioAuthAndHeadersCoverage(t *test
|
||||
}
|
||||
|
||||
r.globalFlags.Token = " explicit "
|
||||
if got := r.resolveAuthToken(context.Background()); got != "explicit" {
|
||||
t.Fatalf("explicit auth token = %q", got)
|
||||
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "explicit" {
|
||||
t.Fatalf("explicit auth token = %q, %v", got, err)
|
||||
}
|
||||
edition.Override(&edition.Hooks{TokenProvider: func(_ context.Context, fallback func() (string, error)) (string, error) {
|
||||
_, _ = fallback()
|
||||
return "provided", nil
|
||||
}})
|
||||
r.globalFlags.Token = ""
|
||||
if got := r.resolveAuthToken(context.Background()); got != "provided" {
|
||||
t.Fatalf("provided auth token = %q", got)
|
||||
if got, err := r.resolveAuthToken(context.Background()); err != nil || got != "provided" {
|
||||
t.Fatalf("provided auth token = %q, %v", got, err)
|
||||
}
|
||||
if got := resolveRuntimeAuthToken(context.Background(), " runtime "); got != "runtime" {
|
||||
t.Fatalf("runtime explicit token = %q", got)
|
||||
if got, err := resolveRuntimeAuthToken(context.Background(), " runtime "); err != nil || got != "runtime" {
|
||||
t.Fatalf("runtime explicit token = %q, %v", got, err)
|
||||
}
|
||||
|
||||
t.Setenv(envDWSChannel, "channel")
|
||||
|
||||
@@ -17,6 +17,7 @@ import (
|
||||
"archive/zip"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime"
|
||||
@@ -36,25 +37,25 @@ import (
|
||||
)
|
||||
|
||||
var (
|
||||
skillLoadAccessToken = loadSkillAccessToken
|
||||
skillDownloadToTmp = downloadSkillToTmpDir
|
||||
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
|
||||
skillNewRequest = http.NewRequestWithContext
|
||||
skillLoadTokenData = authpkg.LoadTokenData
|
||||
skillResolveTargetPath = resolveSkillTargetPath
|
||||
skillFetchDownloadInfo = fetchSkillDownloadInfo
|
||||
skillDownloadFile = downloadSkillFile
|
||||
skillExtractZip = extractSkillZip
|
||||
skillUserHomeDir = os.UserHomeDir
|
||||
skillMkdirTemp = os.MkdirTemp
|
||||
skillCreate = os.Create
|
||||
skillCreateTemp = os.CreateTemp
|
||||
skillRemoveAll = os.RemoveAll
|
||||
skillRemove = os.Remove
|
||||
skillMkdirAll = os.MkdirAll
|
||||
skillOpenFile = os.OpenFile
|
||||
skillCopy = io.Copy
|
||||
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
|
||||
skillLoadAccessToken = loadSkillAccessToken
|
||||
skillDownloadToTmp = downloadSkillToTmpDir
|
||||
skillHTTPDo = func(client *http.Client, req *http.Request) (*http.Response, error) { return client.Do(req) }
|
||||
skillNewRequest = http.NewRequestWithContext
|
||||
skillResolveAccessToken = ResolveAuxiliaryAccessToken
|
||||
skillResolveTargetPath = resolveSkillTargetPath
|
||||
skillFetchDownloadInfo = fetchSkillDownloadInfo
|
||||
skillDownloadFile = downloadSkillFile
|
||||
skillExtractZip = extractSkillZip
|
||||
skillUserHomeDir = os.UserHomeDir
|
||||
skillMkdirTemp = os.MkdirTemp
|
||||
skillCreate = os.Create
|
||||
skillCreateTemp = os.CreateTemp
|
||||
skillRemoveAll = os.RemoveAll
|
||||
skillRemove = os.Remove
|
||||
skillMkdirAll = os.MkdirAll
|
||||
skillOpenFile = os.OpenFile
|
||||
skillCopy = io.Copy
|
||||
skillOpenZipFile = func(file *zip.File) (io.ReadCloser, error) { return file.Open() }
|
||||
)
|
||||
|
||||
func init() {
|
||||
@@ -296,7 +297,7 @@ func newSkillAddHintCommand() *cobra.Command {
|
||||
|
||||
func runSkillGet(cmd *cobra.Command, args []string) error {
|
||||
skillID, _ := cmd.Flags().GetString("skill-id")
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -319,7 +320,7 @@ func runSkillFind(cmd *cobra.Command, args []string) error {
|
||||
if source == "" {
|
||||
source, _ = cmd.Flags().GetString("scopes")
|
||||
}
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -388,7 +389,7 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return apperrors.NewValidation(fmt.Sprintf("invalid target '%s': %v. Supported targets: %s", target, err, supportedTargets()))
|
||||
}
|
||||
|
||||
accessToken, err := skillLoadAccessToken()
|
||||
accessToken, err := skillLoadAccessToken(cmd.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -441,13 +442,16 @@ func runSkillAdd(cmd *cobra.Command, args []string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func loadSkillAccessToken() (string, error) {
|
||||
func loadSkillAccessToken(ctx context.Context) (string, error) {
|
||||
configDir := defaultConfigDir()
|
||||
tokenData, err := skillLoadTokenData(configDir)
|
||||
if err != nil || tokenData == nil || !tokenData.IsAccessTokenValid() {
|
||||
token, err := skillResolveAccessToken(ctx, configDir, "")
|
||||
if errors.Is(err, authpkg.ErrTokenDataNotFound) {
|
||||
return "", skillAuthError()
|
||||
}
|
||||
return tokenData.AccessToken, nil
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("resolve skill access token: %w", err)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func skillAuthError() error {
|
||||
|
||||
@@ -57,11 +57,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
})
|
||||
fail := errors.New("failure")
|
||||
cmd := skillCoverageCommand()
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillGet(cmd, nil); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill get auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
|
||||
if err := runSkillFind(cmd, nil); err == nil {
|
||||
t.Fatal("skill find request failure should propagate")
|
||||
@@ -72,11 +72,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
t.Fatalf("skill get download error = %v", err)
|
||||
}
|
||||
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillFind(cmd, nil); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill find auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillHTTPDo = func(*http.Client, *http.Request) (*http.Response, error) { return nil, fail }
|
||||
if err := runSkillFind(cmd, nil); err == nil {
|
||||
t.Fatal("skill find network failure should propagate")
|
||||
@@ -111,11 +111,11 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
t.Fatal("invalid skill target should fail")
|
||||
}
|
||||
skillResolveTargetPath = func(string) (string, error) { return "dest", nil }
|
||||
skillLoadAccessToken = func() (string, error) { return "", fail }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "", fail }
|
||||
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill add auth error = %v", err)
|
||||
}
|
||||
skillLoadAccessToken = func() (string, error) { return "token", nil }
|
||||
skillLoadAccessToken = func(context.Context) (string, error) { return "token", nil }
|
||||
skillFetchDownloadInfo = func(context.Context, string, string) (*downloadSkillResponse, error) { return nil, fail }
|
||||
if err := runSkillAdd(cmd, []string{"id", "target"}); !errors.Is(err, fail) {
|
||||
t.Fatalf("skill info error = %v", err)
|
||||
@@ -152,25 +152,35 @@ func TestCrossPlatformCoverageSkillCommandHighLevelRemainingCoverage(t *testing.
|
||||
|
||||
func TestCrossPlatformCoverageSkillCommandLowLevelRemainingCoverage(t *testing.T) {
|
||||
oldHTTP := skillHTTPDo
|
||||
oldNewRequest, oldLoadToken := skillNewRequest, skillLoadTokenData
|
||||
oldNewRequest, oldResolveToken := skillNewRequest, skillResolveAccessToken
|
||||
oldHome := skillUserHomeDir
|
||||
oldMkdirTemp, oldCreate, oldCreateTemp := skillMkdirTemp, skillCreate, skillCreateTemp
|
||||
oldRemoveAll, oldRemove, oldMkdir := skillRemoveAll, skillRemove, skillMkdirAll
|
||||
oldOpen, oldCopy, oldZipOpen := skillOpenFile, skillCopy, skillOpenZipFile
|
||||
t.Cleanup(func() {
|
||||
skillHTTPDo = oldHTTP
|
||||
skillNewRequest, skillLoadTokenData = oldNewRequest, oldLoadToken
|
||||
skillNewRequest, skillResolveAccessToken = oldNewRequest, oldResolveToken
|
||||
skillUserHomeDir = oldHome
|
||||
skillMkdirTemp, skillCreate, skillCreateTemp = oldMkdirTemp, oldCreate, oldCreateTemp
|
||||
skillRemoveAll, skillRemove, skillMkdirAll = oldRemoveAll, oldRemove, oldMkdir
|
||||
skillOpenFile, skillCopy, skillOpenZipFile = oldOpen, oldCopy, oldZipOpen
|
||||
})
|
||||
fail := errors.New("failure")
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, nil }
|
||||
if _, err := loadSkillAccessToken(); err == nil {
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
if _, err := loadSkillAccessToken(context.Background()); err == nil {
|
||||
t.Fatal("invalid skill access token succeeded")
|
||||
}
|
||||
skillLoadTokenData = oldLoadToken
|
||||
canceled, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
skillResolveAccessToken = func(ctx context.Context, _, _ string) (string, error) {
|
||||
return "", ctx.Err()
|
||||
}
|
||||
if _, err := loadSkillAccessToken(canceled); !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("skill token cancellation = %v", err)
|
||||
}
|
||||
skillResolveAccessToken = oldResolveToken
|
||||
skillNewRequest = func(context.Context, string, string, io.Reader) (*http.Request, error) { return nil, fail }
|
||||
if _, err := fetchSkillDownloadInfo(context.Background(), "token", "id"); err == nil {
|
||||
t.Fatal("download-info request failure should propagate")
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
@@ -397,9 +396,11 @@ func TestSkillInstallRequiresAuth(t *testing.T) {
|
||||
configDir := filepath.Join(tempDir, "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
t.Cleanup(CloseFileLogger)
|
||||
originalLoadToken := skillLoadTokenData
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) { return nil, errors.New("missing") }
|
||||
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
|
||||
originalResolveToken := skillResolveAccessToken
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "", authpkg.ErrTokenDataNotFound
|
||||
}
|
||||
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
|
||||
|
||||
// Ensure the config directory exists but has no token
|
||||
if err := os.MkdirAll(configDir, 0755); err != nil {
|
||||
@@ -679,16 +680,11 @@ func TestSkillSearchUsesSourceQueryAndKeepsScopesCompat(t *testing.T) {
|
||||
configDir := filepath.Join(t.TempDir(), "config")
|
||||
t.Setenv("DWS_CONFIG_DIR", configDir)
|
||||
t.Cleanup(CloseFileLogger)
|
||||
originalLoadToken := skillLoadTokenData
|
||||
skillLoadTokenData = func(string) (*authpkg.TokenData, error) {
|
||||
return &authpkg.TokenData{
|
||||
AccessToken: "test-token",
|
||||
RefreshToken: "refresh-token",
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
RefreshExpAt: time.Now().Add(24 * time.Hour),
|
||||
}, nil
|
||||
originalResolveToken := skillResolveAccessToken
|
||||
skillResolveAccessToken = func(context.Context, string, string) (string, error) {
|
||||
return "test-token", nil
|
||||
}
|
||||
t.Cleanup(func() { skillLoadTokenData = originalLoadToken })
|
||||
t.Cleanup(func() { skillResolveAccessToken = originalResolveToken })
|
||||
|
||||
var gotSources []string
|
||||
var gotScopes []string
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotPreservesLoadFailure(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
want := errors.New("keychain permission denied")
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
_, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
if errors.Is(err, ErrTokenDataNotFound) {
|
||||
t.Fatalf("load failure was misclassified as missing credentials: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderLoginPreservesLoadFailure(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
want := errors.New("keychain permission denied")
|
||||
oauthLoadToken = func(string) (*TokenData, error) { return nil, want }
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
_, err := NewOAuthProvider(t.TempDir(), nil).Login(context.Background(), false)
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want cause %v", err, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageOAuthProviderTokenSnapshotReturnsExpiryMetadata(t *testing.T) {
|
||||
oldLoad := oauthLoadToken
|
||||
expiresAt := time.Now().Add(time.Hour)
|
||||
oauthLoadToken = func(string) (*TokenData, error) {
|
||||
return &TokenData{AccessToken: "token", ExpiresAt: expiresAt}, nil
|
||||
}
|
||||
t.Cleanup(func() { oauthLoadToken = oldLoad })
|
||||
|
||||
snapshot, err := NewOAuthProvider(t.TempDir(), nil).GetTokenSnapshot(context.Background())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if snapshot.AccessToken != "token" || !snapshot.ExpiresAt.Equal(expiresAt) {
|
||||
t.Fatalf("snapshot = %#v", snapshot)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageTokenMarkerRevisionChangesOnEveryPublication(t *testing.T) {
|
||||
configDir := t.TempDir()
|
||||
if err := WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
first, present, err := ReadTokenMarkerRevision(configDir)
|
||||
if err != nil || !present || first == "" {
|
||||
t.Fatalf("first marker = %q, %v, %v", first, present, err)
|
||||
}
|
||||
if err := WriteTokenMarker(configDir); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
second, present, err := ReadTokenMarkerRevision(configDir)
|
||||
if err != nil || !present || second == "" || second == first {
|
||||
t.Fatalf("second marker = %q, %v, %v; first=%q", second, present, err, first)
|
||||
}
|
||||
}
|
||||
@@ -47,8 +47,10 @@ func (m *Manager) GetToken() (string, string, error) {
|
||||
}
|
||||
return token, "file", nil
|
||||
}
|
||||
|
||||
return "", "", fmt.Errorf("%s", i18n.T("未找到认证信息,请运行 dws auth login"))
|
||||
if err != nil && !os.IsNotExist(err) {
|
||||
return "", "", fmt.Errorf("load legacy token: %w", err)
|
||||
}
|
||||
return "", "", fmt.Errorf("%s: %w", i18n.T("未找到认证信息,请运行 dws auth login"), ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
func (m *Manager) GetMCPURL() (string, error) {
|
||||
|
||||
@@ -115,6 +115,12 @@ func (p *OAuthProvider) Login(ctx context.Context, force bool) (*TokenData, erro
|
||||
// Smart degradation: try silent refresh before opening browser.
|
||||
if !force {
|
||||
data, err := oauthLoadToken(p.configDir)
|
||||
if err != nil && !errors.Is(err, ErrTokenDataNotFound) && !os.IsNotExist(err) {
|
||||
if preflightErr := preflightTokenPersistence(p.configDir); preflightErr != nil {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("本地登录态无法安全更新"), preflightErr)
|
||||
}
|
||||
return nil, fmt.Errorf("load existing access token: %w", err)
|
||||
}
|
||||
if err == nil {
|
||||
// Case 1: access_token still valid — no action needed.
|
||||
if data.IsAccessTokenValid() {
|
||||
@@ -636,36 +642,50 @@ continueLogin:
|
||||
return tokenData, nil
|
||||
}
|
||||
|
||||
// GetAccessToken returns a valid access token, auto-refreshing if needed.
|
||||
// Uses a file lock with double-check pattern to prevent concurrent refresh
|
||||
// from multiple CLI processes.
|
||||
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
|
||||
// GetTokenSnapshot returns a valid token together with its expiry metadata.
|
||||
// Storage and refresh failures retain their original cause; only a confirmed
|
||||
// missing credential is reported as ErrTokenDataNotFound.
|
||||
func (p *OAuthProvider) GetTokenSnapshot(ctx context.Context) (*TokenData, error) {
|
||||
data, err := oauthLoadToken(p.configDir)
|
||||
if err != nil {
|
||||
return "", errors.New(i18n.T("未登录,请运行 dws auth login"))
|
||||
if errors.Is(err, ErrTokenDataNotFound) || os.IsNotExist(err) {
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("未登录,请运行 dws auth login"), ErrTokenDataNotFound)
|
||||
}
|
||||
return nil, fmt.Errorf("load access token: %w", err)
|
||||
}
|
||||
|
||||
// Fast path: access_token still valid — no lock needed.
|
||||
if data.IsAccessTokenValid() {
|
||||
return data.AccessToken, nil
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// Slow path: token expired — try locked refresh.
|
||||
if data.IsRefreshTokenValid() {
|
||||
refreshed, rErr := p.lockedRefresh(ctx)
|
||||
if rErr == nil {
|
||||
return refreshed.AccessToken, nil
|
||||
return refreshed, nil
|
||||
}
|
||||
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
|
||||
if p.logger != nil {
|
||||
p.logger.Warn(i18n.T("refresh_token 刷新失败"), "error", rErr)
|
||||
}
|
||||
return "", fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("refresh_token 刷新失败"), rErr)
|
||||
} else {
|
||||
_ = oauthMarkProfile(p.configDir, TokenProfileSelector(data), ProfileStatusExpired)
|
||||
}
|
||||
|
||||
return "", errors.New(i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"))
|
||||
return nil, fmt.Errorf("%s: %w", i18n.T("所有凭证已失效,请运行 dws auth login 重新登录"), ErrTokenDataNotFound)
|
||||
}
|
||||
|
||||
// GetAccessToken returns a valid access token, auto-refreshing if needed.
|
||||
// Uses a file lock with double-check pattern to prevent concurrent refresh
|
||||
// from multiple CLI processes.
|
||||
func (p *OAuthProvider) GetAccessToken(ctx context.Context) (string, error) {
|
||||
data, err := p.GetTokenSnapshot(ctx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return strings.TrimSpace(data.AccessToken), nil
|
||||
}
|
||||
|
||||
// lockedRefresh attempts to refresh the token while holding dual-layer locks.
|
||||
|
||||
@@ -127,6 +127,10 @@ const tokenJSONFile = "token.json"
|
||||
type TokenMarker struct {
|
||||
UpdatedAt string `json:"updated_at"`
|
||||
ManualToken bool `json:"manual_token,omitempty"`
|
||||
// Revision changes on every credential publication. Runtime token caches
|
||||
// use it as a cheap cross-process invalidation signal without reading the
|
||||
// platform keychain on every request.
|
||||
Revision string `json:"revision,omitempty"`
|
||||
}
|
||||
|
||||
// WriteTokenMarker writes a token.json marker containing only an updated_at
|
||||
@@ -147,6 +151,7 @@ func writeTokenMarker(configDir string, manual bool) error {
|
||||
marker := TokenMarker{
|
||||
UpdatedAt: time.Now().Format(time.RFC3339),
|
||||
ManualToken: manual,
|
||||
Revision: uuid.NewString(),
|
||||
}
|
||||
data, _ := tokenJSONMarshalIndent(marker, "", " ")
|
||||
if err := tokenMkdirAll(configDir, 0o700); err != nil {
|
||||
@@ -159,6 +164,27 @@ func writeTokenMarker(configDir string, manual bool) error {
|
||||
return tokenRename(tmp, filepath.Join(configDir, tokenJSONFile))
|
||||
}
|
||||
|
||||
// ReadTokenMarkerRevision returns the current credential publication revision.
|
||||
// Existing markers without a revision remain readable, but callers must avoid
|
||||
// caching them because they cannot prove that the credential is unchanged.
|
||||
func ReadTokenMarkerRevision(configDir string) (revision string, present bool, err error) {
|
||||
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return "", false, nil
|
||||
}
|
||||
return "", false, fmt.Errorf("read token marker: %w", err)
|
||||
}
|
||||
var marker TokenMarker
|
||||
if err := json.Unmarshal(data, &marker); err != nil {
|
||||
// The marker is only a cache-coherency hint. A malformed historical or
|
||||
// externally modified marker must disable caching, not make an otherwise
|
||||
// valid credential unusable.
|
||||
return "", true, nil
|
||||
}
|
||||
return strings.TrimSpace(marker.Revision), true, nil
|
||||
}
|
||||
|
||||
func manualTokenMarkerActive(configDir string) (bool, error) {
|
||||
data, err := tokenReadFile(filepath.Join(configDir, tokenJSONFile))
|
||||
if err != nil {
|
||||
|
||||
@@ -63,9 +63,10 @@ func (i Identity) Key() string {
|
||||
}
|
||||
|
||||
type Client struct {
|
||||
BaseURL string
|
||||
HTTPClient *http.Client
|
||||
Identity Identity
|
||||
BaseURL string
|
||||
HTTPClient *http.Client
|
||||
Identity Identity
|
||||
AccessTokenProvider func(context.Context) (string, error)
|
||||
}
|
||||
|
||||
type CreateSubscriptionRequest struct {
|
||||
@@ -337,8 +338,9 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
|
||||
if c == nil {
|
||||
return errors.New("personal event: nil client")
|
||||
}
|
||||
if c.Identity.AccessToken == "" {
|
||||
return errors.New("personal event: access token is required")
|
||||
accessToken, err := c.resolveAccessToken(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
u := strings.TrimRight(c.BaseURL, "/") + path
|
||||
if len(q) > 0 {
|
||||
@@ -358,7 +360,7 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
|
||||
if err != nil {
|
||||
return fmt.Errorf("personal event: create request: %w", err)
|
||||
}
|
||||
c.decorate(req)
|
||||
c.decorate(req, accessToken)
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
@@ -424,9 +426,26 @@ func (c *Client) do(ctx context.Context, method, path string, q url.Values, body
|
||||
return json.Unmarshal(data, out)
|
||||
}
|
||||
|
||||
func (c *Client) decorate(req *http.Request) {
|
||||
req.Header.Set("Authorization", "Bearer "+c.Identity.AccessToken)
|
||||
req.Header.Set("x-user-access-token", c.Identity.AccessToken)
|
||||
func (c *Client) resolveAccessToken(ctx context.Context) (string, error) {
|
||||
if c.AccessTokenProvider != nil {
|
||||
token, err := c.AccessTokenProvider(ctx)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("personal event: resolve access token: %w", err)
|
||||
}
|
||||
if token = strings.TrimSpace(token); token != "" {
|
||||
return token, nil
|
||||
}
|
||||
return "", errors.New("personal event: access token provider returned empty token")
|
||||
}
|
||||
if token := strings.TrimSpace(c.Identity.AccessToken); token != "" {
|
||||
return token, nil
|
||||
}
|
||||
return "", errors.New("personal event: access token is required")
|
||||
}
|
||||
|
||||
func (c *Client) decorate(req *http.Request, accessToken string) {
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
req.Header.Set("X-DWS-Client-Id", c.Identity.ClientID)
|
||||
req.Header.Set("X-DWS-Source-Id", c.Identity.SourceID)
|
||||
if c.Identity.CorpID != "" {
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package personal
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type accessTokenRoundTripper func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f accessTokenRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageClientResolvesAccessTokenPerRequest(t *testing.T) {
|
||||
tokens := []string{"token-a", "token-b"}
|
||||
calls := 0
|
||||
client := NewClient("https://control.test", Identity{AccessToken: "stale", ClientID: "client", SourceID: "source"})
|
||||
client.AccessTokenProvider = func(context.Context) (string, error) {
|
||||
token := tokens[calls]
|
||||
calls++
|
||||
return token, nil
|
||||
}
|
||||
client.HTTPClient = &http.Client{Transport: accessTokenRoundTripper(func(req *http.Request) (*http.Response, error) {
|
||||
want := tokens[calls-1]
|
||||
if got := req.Header.Get("Authorization"); got != "Bearer "+want {
|
||||
t.Fatalf("Authorization = %q, want Bearer %s", got, want)
|
||||
}
|
||||
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"success":true,"result":{"items":[]}}`)), Header: make(http.Header)}, nil
|
||||
})}
|
||||
for range 2 {
|
||||
if _, err := client.ListSubscriptions(context.Background(), ListOptions{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoverageClientDoesNotFallBackAfterProviderFailure(t *testing.T) {
|
||||
want := errors.New("keychain failed")
|
||||
client := NewClient("https://control.test", Identity{AccessToken: "stale", ClientID: "client", SourceID: "source"})
|
||||
client.AccessTokenProvider = func(context.Context) (string, error) { return "", want }
|
||||
client.HTTPClient = &http.Client{Transport: accessTokenRoundTripper(func(*http.Request) (*http.Response, error) {
|
||||
t.Fatal("HTTP must not run after token provider failure")
|
||||
return nil, nil
|
||||
})}
|
||||
_, err := client.ListSubscriptions(context.Background(), ListOptions{})
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("error = %v, want %v", err, want)
|
||||
}
|
||||
}
|
||||
@@ -41,19 +41,22 @@ const (
|
||||
)
|
||||
|
||||
type PersonalConfig struct {
|
||||
AccessToken string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
SourceID string
|
||||
TicketURL string
|
||||
TicketMode string
|
||||
HTTPClient *http.Client
|
||||
WebSocketDialer *websocket.Dialer
|
||||
Now func() time.Time
|
||||
ReconnectMin time.Duration
|
||||
ReconnectMax time.Duration
|
||||
AccessToken string
|
||||
AccessTokenProvider AccessTokenProvider
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
SourceID string
|
||||
TicketURL string
|
||||
TicketMode string
|
||||
HTTPClient *http.Client
|
||||
WebSocketDialer *websocket.Dialer
|
||||
Now func() time.Time
|
||||
ReconnectMin time.Duration
|
||||
ReconnectMax time.Duration
|
||||
}
|
||||
|
||||
type AccessTokenProvider func(context.Context) (string, error)
|
||||
|
||||
type PersonalSource struct {
|
||||
cfg PersonalConfig
|
||||
machine *Machine
|
||||
@@ -73,8 +76,8 @@ type ticketResponse struct {
|
||||
}
|
||||
|
||||
func NewPersonal(cfg PersonalConfig) (*PersonalSource, error) {
|
||||
if strings.TrimSpace(cfg.AccessToken) == "" {
|
||||
return nil, errors.New("personal source: AccessToken is required")
|
||||
if cfg.AccessTokenProvider == nil && strings.TrimSpace(cfg.AccessToken) == "" {
|
||||
return nil, errors.New("personal source: AccessToken or AccessTokenProvider is required")
|
||||
}
|
||||
if strings.TrimSpace(cfg.ClientID) == "" {
|
||||
return nil, errors.New("personal source: ClientID is required")
|
||||
@@ -192,6 +195,10 @@ func (s *PersonalSource) runAttempt(ctx context.Context, emit dwsevent.EmitFn) (
|
||||
}
|
||||
|
||||
func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, error) {
|
||||
accessToken, err := resolveSourceAccessToken(ctx, s.cfg.AccessTokenProvider, s.cfg.AccessToken, "personal source")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
body := map[string]any{
|
||||
"sourceId": s.cfg.SourceID,
|
||||
"mode": s.cfg.TicketMode,
|
||||
@@ -207,8 +214,8 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
req.Header.Set("x-user-access-token", s.cfg.AccessToken)
|
||||
req.Header.Set("Authorization", "Bearer "+s.cfg.AccessToken)
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
req.Header.Set("Authorization", "Bearer "+accessToken)
|
||||
req.Header.Set("X-DWS-Client-Id", s.cfg.ClientID)
|
||||
req.Header.Set("X-DWS-Source-Id", s.cfg.SourceID)
|
||||
|
||||
@@ -238,6 +245,23 @@ func (s *PersonalSource) fetchTicket(ctx context.Context) (*ticketResponse, erro
|
||||
return ticket, nil
|
||||
}
|
||||
|
||||
func resolveSourceAccessToken(ctx context.Context, provider AccessTokenProvider, fallback, component string) (string, error) {
|
||||
if provider != nil {
|
||||
token, err := provider(ctx)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%s: resolve access token: %w", component, err)
|
||||
}
|
||||
if token = strings.TrimSpace(token); token != "" {
|
||||
return token, nil
|
||||
}
|
||||
return "", fmt.Errorf("%s: access token provider returned empty token", component)
|
||||
}
|
||||
if token := strings.TrimSpace(fallback); token != "" {
|
||||
return token, nil
|
||||
}
|
||||
return "", fmt.Errorf("%s: access token is required", component)
|
||||
}
|
||||
|
||||
func (s *PersonalSource) handleFrame(conn *websocket.Conn, data []byte, emit dwsevent.EmitFn) error {
|
||||
df, err := payload.DecodeDataFrame(data)
|
||||
if err != nil {
|
||||
|
||||
@@ -39,14 +39,15 @@ const (
|
||||
// normal mode uses portal-side managed credentials; custom mode asks portal to
|
||||
// open the user connection with the caller-provided clientId/clientSecret.
|
||||
type PortalTicketConfig struct {
|
||||
TicketURL string
|
||||
AccessToken string
|
||||
SourceID string
|
||||
Mode string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
UserAgent string
|
||||
HTTPClient *http.Client
|
||||
TicketURL string
|
||||
AccessToken string
|
||||
AccessTokenProvider AccessTokenProvider
|
||||
SourceID string
|
||||
Mode string
|
||||
ClientID string
|
||||
ClientSecret string
|
||||
UserAgent string
|
||||
HTTPClient *http.Client
|
||||
}
|
||||
|
||||
var portalWriteMessage = func(conn *websocket.Conn, messageType int, data []byte) error {
|
||||
@@ -60,8 +61,8 @@ func (c *PortalTicketConfig) Valid() error {
|
||||
if strings.TrimSpace(c.TicketURL) == "" {
|
||||
return errors.New("source: portal ticket URL is required")
|
||||
}
|
||||
if strings.TrimSpace(c.AccessToken) == "" {
|
||||
return errors.New("source: portal access token is required")
|
||||
if c.AccessTokenProvider == nil && strings.TrimSpace(c.AccessToken) == "" {
|
||||
return errors.New("source: portal access token or provider is required")
|
||||
}
|
||||
if strings.TrimSpace(c.SourceID) == "" {
|
||||
return errors.New("source: portal sourceId is required")
|
||||
@@ -161,6 +162,10 @@ type portalStreamTicket struct {
|
||||
}
|
||||
|
||||
func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalStreamTicket, error) {
|
||||
accessToken, err := resolveSourceAccessToken(ctx, cfg.AccessTokenProvider, cfg.AccessToken, "source: portal ticket")
|
||||
if err != nil {
|
||||
return portalStreamTicket{}, err
|
||||
}
|
||||
httpClient := cfg.HTTPClient
|
||||
if httpClient == nil {
|
||||
httpClient = &http.Client{Timeout: 20 * time.Second}
|
||||
@@ -184,7 +189,7 @@ func requestPortalTicket(ctx context.Context, cfg *PortalTicketConfig) (portalSt
|
||||
if ua := strings.TrimSpace(cfg.UserAgent); ua != "" {
|
||||
req.Header.Set("User-Agent", ua)
|
||||
}
|
||||
req.Header.Set("x-user-access-token", cfg.AccessToken)
|
||||
req.Header.Set("x-user-access-token", accessToken)
|
||||
|
||||
resp, err := httpClient.Do(req)
|
||||
if err != nil {
|
||||
|
||||
@@ -0,0 +1,66 @@
|
||||
package source
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type tokenProviderRoundTripper func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f tokenProviderRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
return f(req)
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePersonalSourceResolvesTokenForEveryTicketRequest(t *testing.T) {
|
||||
tokens := []string{"token-a", "token-b"}
|
||||
calls := 0
|
||||
source, err := NewPersonal(PersonalConfig{
|
||||
AccessTokenProvider: func(context.Context) (string, error) {
|
||||
token := tokens[calls]
|
||||
calls++
|
||||
return token, nil
|
||||
},
|
||||
ClientID: "client",
|
||||
SourceID: "source",
|
||||
TicketURL: "https://ticket.test",
|
||||
HTTPClient: &http.Client{Transport: tokenProviderRoundTripper(func(req *http.Request) (*http.Response, error) {
|
||||
want := tokens[calls-1]
|
||||
if got := req.Header.Get("x-user-access-token"); got != want {
|
||||
t.Fatalf("token header = %q, want %q", got, want)
|
||||
}
|
||||
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(`{"endpoint":"wss://stream.test","ticket":"ticket"}`)), Header: make(http.Header)}, nil
|
||||
})},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for range 2 {
|
||||
if _, err := source.fetchTicket(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if calls != 2 {
|
||||
t.Fatalf("provider calls = %d, want 2", calls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCrossPlatformCoveragePortalTicketProviderFailureStopsBeforeHTTP(t *testing.T) {
|
||||
want := errors.New("token store failed")
|
||||
httpCalled := false
|
||||
_, err := requestPortalTicket(context.Background(), &PortalTicketConfig{
|
||||
TicketURL: "https://ticket.test",
|
||||
AccessTokenProvider: func(context.Context) (string, error) { return "", want },
|
||||
SourceID: "source",
|
||||
HTTPClient: &http.Client{Transport: tokenProviderRoundTripper(func(*http.Request) (*http.Response, error) {
|
||||
httpCalled = true
|
||||
return nil, errors.New("unexpected HTTP")
|
||||
})},
|
||||
})
|
||||
if !errors.Is(err, want) || httpCalled {
|
||||
t.Fatalf("request error = %v, httpCalled=%v", err, httpCalled)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user