diff --git a/go.mod b/go.mod index d041cdd..b0c6c34 100644 --- a/go.mod +++ b/go.mod @@ -8,7 +8,6 @@ require ( github.com/charmbracelet/bubbletea v1.3.10 github.com/charmbracelet/lipgloss v1.1.0 github.com/fatih/color v1.19.0 - github.com/go-jose/go-jose/v4 v4.1.4 github.com/itchyny/gojq v0.12.19 github.com/joho/godotenv v1.5.1 github.com/mattn/go-runewidth v0.0.27 diff --git a/go.sum b/go.sum index 313f2ec..e5e9540 100644 --- a/go.sum +++ b/go.sum @@ -49,8 +49,6 @@ github.com/fsnotify/fsnotify v1.4.9 h1:hsms1Qyu0jgnwNXIxa+/V/PDsU6CfLf6CNO8H7IWo github.com/fsnotify/fsnotify v1.4.9/go.mod h1:znqG4EE+3YCdAaPaxE2ZRY/06pZUdp0tY4IgpuI1SZQ= github.com/getkin/kin-openapi v0.133.0 h1:pJdmNohVIJ97r4AUFtEXRXwESr8b0bD721u/Tz6k8PQ= github.com/getkin/kin-openapi v0.133.0/go.mod h1:boAciF6cXk5FhPqe/NQeBTeenbjqU4LhWBf09ILVvWE= -github.com/go-jose/go-jose/v4 v4.1.4 h1:moDMcTHmvE6Groj34emNPLs/qtYXRVcd6S7NHbHz3kA= -github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/go-openapi/jsonpointer v0.21.0 h1:YgdVicSA9vH5RiHs9TZW5oyafXZFc6+2Vc1rr/O9oNQ= github.com/go-openapi/jsonpointer v0.21.0/go.mod h1:IUyH9l/+uyhIYQ/PXVA41Rexl+kOkAPDdXEYns6fzUY= github.com/go-openapi/swag v0.23.0 h1:vsEVJDUo2hPJ2tu0/Xc+4noaxyEffXNIs3cOULZ+GrE= diff --git a/internal/auth/auth.go b/internal/auth/auth.go index 102619b..f4a027a 100644 --- a/internal/auth/auth.go +++ b/internal/auth/auth.go @@ -102,14 +102,17 @@ func oauthScopesFromMetadata(meta map[string]any) ([]string, error) { return nil, fmt.Errorf("server OAuth metadata contains an invalid supported scope") } scope = strings.TrimSpace(scope) + if scope == "openid" || scope == "profile" || scope == "email" { + continue + } if _, duplicate := seen[scope]; duplicate { continue } seen[scope] = struct{}{} scopes = append(scopes, scope) } - if _, ok := seen["openid"]; !ok { - return nil, fmt.Errorf("server OAuth metadata does not support the required openid scope") + if len(scopes) == 0 { + return nil, fmt.Errorf("server OAuth metadata does not advertise usable API scopes") } return scopes, nil } @@ -391,18 +394,8 @@ func Login(server string) (*config.Credential, error) { return nil, fmt.Errorf("token exchange failed: %w", err) } - vt := newVerifiedToken(tok) - if err := requireIDTokenForOpenID(effectiveTokenScope(vt, scope), vt.IDToken); err != nil { - return nil, err - } - issuer := stringFromMap(meta, "issuer") - if issuer == "" { - issuer = server - } - if err := vt.ValidateIDToken(issuer, clientID); err != nil { - return nil, err - } - return verifiedTokenToCredential(clientID, resource, vt, "", scope, time.Now()) + token := newOAuthToken(tok) + return oauthTokenToCredential(clientID, resource, token, "", scope, time.Now()) } // RefreshToken attempts to refresh the access token. @@ -427,13 +420,6 @@ func RefreshToken(server string, cred *config.Credential) (*config.Credential, e return nil, fmt.Errorf("refresh failed: %w", err) } - vt := newVerifiedToken(tok) - issuer := stringFromMap(meta, "issuer") - if issuer == "" { - issuer = server - } - if err := vt.ValidateIDToken(issuer, cred.ClientID); err != nil { - return nil, err - } - return verifiedTokenToCredential(cred.ClientID, resource, vt, cred.RefreshToken, cred.Scope, time.Now()) + token := newOAuthToken(tok) + return oauthTokenToCredential(cred.ClientID, resource, token, cred.RefreshToken, cred.Scope, time.Now()) } diff --git a/internal/auth/auth_test.go b/internal/auth/auth_test.go index 05b6035..24880fb 100644 --- a/internal/auth/auth_test.go +++ b/internal/auth/auth_test.go @@ -1,8 +1,6 @@ package auth import ( - "crypto/rand" - "crypto/rsa" "encoding/json" "net" "net/http" @@ -13,8 +11,7 @@ import ( "time" "github.com/Life-USTC/CLI/internal/config" - "github.com/go-jose/go-jose/v4" - "github.com/go-jose/go-jose/v4/jwt" + "golang.org/x/oauth2" ) func TestRegisterPublicClientUsesNativeApplicationType(t *testing.T) { @@ -34,7 +31,7 @@ func TestRegisterPublicClientUsesNativeApplicationType(t *testing.T) { _, err := registerPublicClient( server.URL, - []string{"openid", "profile", "workspace.todo:read"}, + []string{"offline_access", "workspace.todo:read"}, []string{"http://127.0.0.1:46289/callback"}, []string{"authorization_code", "refresh_token"}, []string{"code"}, @@ -46,7 +43,7 @@ func TestRegisterPublicClientUsesNativeApplicationType(t *testing.T) { if body["application_type"] != "native" { t.Fatalf("application_type = %#v, want native", body["application_type"]) } - if body["scope"] != "openid profile workspace.todo:read" { + if body["scope"] != "offline_access workspace.todo:read" { t.Fatalf("scope = %#v", body["scope"]) } redirectURIs, ok := body["redirect_uris"].([]any) @@ -71,7 +68,7 @@ func TestRegisterPublicClientOmitsUnusedDeviceRedirectMetadata(t *testing.T) { _, err := registerPublicClient( server.URL, - []string{"openid", "profile", "workspace.todo:read"}, + []string{"offline_access", "workspace.todo:read"}, nil, []string{"urn:ietf:params:oauth:grant-type:device_code", "refresh_token"}, nil, @@ -83,6 +80,9 @@ func TestRegisterPublicClientOmitsUnusedDeviceRedirectMetadata(t *testing.T) { if body["application_type"] != "native" { t.Fatalf("application_type = %#v, want native", body["application_type"]) } + if body["scope"] != "offline_access workspace.todo:read" { + t.Fatalf("scope = %#v", body["scope"]) + } if _, ok := body["redirect_uris"]; ok { t.Fatalf("redirect_uris should be omitted, got %#v", body["redirect_uris"]) } @@ -96,6 +96,8 @@ func TestOAuthScopesFromMetadata(t *testing.T) { "scopes_supported": []any{ "openid", "profile", + "email", + "offline_access", "workspace.todo:read", "workspace.todo:write", "workspace.todo:read", @@ -105,8 +107,7 @@ func TestOAuthScopesFromMetadata(t *testing.T) { t.Fatal(err) } want := []string{ - "openid", - "profile", + "offline_access", "workspace.todo:read", "workspace.todo:write", } @@ -128,9 +129,9 @@ func TestOAuthScopesFromMetadataRejectsMissingOrInvalidScopes(t *testing.T) { {name: "missing", meta: map[string]any{}}, {name: "wrong type", meta: map[string]any{"scopes_supported": "openid profile"}}, {name: "empty", meta: map[string]any{"scopes_supported": []any{}}}, - {name: "invalid item", meta: map[string]any{"scopes_supported": []any{"openid", 42}}}, - {name: "blank item", meta: map[string]any{"scopes_supported": []any{"openid", " "}}}, - {name: "missing openid", meta: map[string]any{"scopes_supported": []any{"profile"}}}, + {name: "invalid item", meta: map[string]any{"scopes_supported": []any{"offline_access", 42}}}, + {name: "blank item", meta: map[string]any{"scopes_supported": []any{"offline_access", " "}}}, + {name: "identity only", meta: map[string]any{"scopes_supported": []any{"openid", "profile", "email"}}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -141,6 +142,95 @@ func TestOAuthScopesFromMetadataRejectsMissingOrInvalidScopes(t *testing.T) { } } +func TestLoginDeviceCodeAcceptsOAuthTokenWithoutIDToken(t *testing.T) { + t.Setenv("PATH", t.TempDir()) + registrationScopes := make(chan string, 1) + deviceScopes := make(chan string, 1) + tokenScopes := make(chan string, 1) + + var server *httptest.Server + server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch r.URL.Path { + case "/.well-known/oauth-authorization-server/api/auth": + _ = json.NewEncoder(w).Encode(map[string]any{ + "issuer": server.URL + "/api/auth", + "registration_endpoint": server.URL + "/api/auth/oauth2/register", + "device_authorization_endpoint": server.URL + "/api/auth/oauth2/device-authorization", + "token_endpoint": server.URL + "/api/auth/oauth2/token", + "scopes_supported": []string{ + "openid", + "profile", + "email", + "offline_access", + "workspace.todo:read", + }, + }) + case "/api/auth/oauth2/register": + var body map[string]any + if err := json.NewDecoder(r.Body).Decode(&body); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + registrationScope, _ := body["scope"].(string) + registrationScopes <- registrationScope + w.WriteHeader(http.StatusCreated) + _, _ = w.Write([]byte(`{"client_id":"device-client"}`)) + case "/api/auth/oauth2/device-authorization": + if err := r.ParseForm(); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + deviceScopes <- r.Form.Get("scope") + _ = json.NewEncoder(w).Encode(map[string]any{ + "device_code": "device-code", + "user_code": "TEST-CODE", + "verification_uri": server.URL + "/oauth/device", + "verification_uri_complete": server.URL + "/oauth/device?code=TEST-CODE", + "expires_in": 60, + "interval": 1, + }) + case "/api/auth/oauth2/token": + if err := r.ParseForm(); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + tokenScopes <- r.Form.Get("scope") + _ = json.NewEncoder(w).Encode(map[string]any{ + "access_token": "device-access", + "refresh_token": "device-refresh", + "token_type": "Bearer", + "expires_in": 3600, + "scope": "offline_access workspace.todo:read", + }) + default: + http.NotFound(w, r) + } + })) + t.Cleanup(server.Close) + + cred, err := LoginDeviceCode(server.URL) + if err != nil { + t.Fatal(err) + } + if cred.AccessToken != "device-access" || cred.RefreshToken != "device-refresh" { + t.Fatalf("credential = %#v", cred) + } + const wantScope = "offline_access workspace.todo:read" + if got := <-registrationScopes; got != wantScope { + t.Fatalf("registration scope = %q, want %q", got, wantScope) + } + if got := <-deviceScopes; got != wantScope { + t.Fatalf("device authorization scope = %q, want %q", got, wantScope) + } + if got := <-tokenScopes; got != wantScope { + t.Fatalf("token scope = %q, want %q", got, wantScope) + } + if cred.Scope != wantScope { + t.Fatalf("credential scope = %q, want %q", cred.Scope, wantScope) + } +} + func TestOAuthCallbackHandlerDeliversOnlyFirstRequest(t *testing.T) { results := make(chan callbackResult, 1) handler := oauthCallbackHandler(results) @@ -211,12 +301,12 @@ func TestCallbackRedirectURIMatchesLoopbackListener(t *testing.T) { } } -func TestVerifiedTokenToCredentialUsesFallbacks(t *testing.T) { - vt := &VerifiedToken{ +func TestOAuthTokenToCredentialUsesFallbacks(t *testing.T) { + token := &oauthToken{ AccessToken: "access", ExpiresIn: 120, } - cred, err := verifiedTokenToCredential("client", "https://example.test", vt, "refresh", "openid", time.Now()) + cred, err := oauthTokenToCredential("client", "https://example.test", token, "refresh", "openid", time.Now()) if err != nil { t.Fatal(err) } @@ -231,39 +321,24 @@ func TestVerifiedTokenToCredentialUsesFallbacks(t *testing.T) { } } -func TestVerifiedTokenToCredentialRequiresAccessToken(t *testing.T) { - vt := &VerifiedToken{} - if _, err := verifiedTokenToCredential("client", "resource", vt, "", "", time.Now()); err == nil { +func TestOAuthTokenToCredentialRequiresAccessToken(t *testing.T) { + token := &oauthToken{} + if _, err := oauthTokenToCredential("client", "resource", token, "", "", time.Now()); err == nil { t.Fatal("expected missing access token error") } } -func TestVerifiedTokenToCredentialNilGuard(t *testing.T) { - if _, err := verifiedTokenToCredential("client", "resource", nil, "", "", time.Now()); err == nil { +func TestOAuthTokenToCredentialNilGuard(t *testing.T) { + if _, err := oauthTokenToCredential("client", "resource", nil, "", "", time.Now()); err == nil { t.Fatal("expected error for nil token") } } -func TestRequireIDTokenForOpenID(t *testing.T) { - if err := requireIDTokenForOpenID("openid profile", ""); err == nil { - t.Fatal("expected error when openid scope requested without id_token") - } - if err := requireIDTokenForOpenID("profile email", ""); err != nil { - t.Fatalf("unexpected error when openid not requested: %v", err) - } - if err := requireIDTokenForOpenID("openid", "idtoken"); err != nil { - t.Fatalf("unexpected error when id_token present: %v", err) - } -} - func TestEffectiveTokenScopePrefersGrantedScope(t *testing.T) { - vt := &VerifiedToken{Scope: "profile workspace.todo:read"} - if got := effectiveTokenScope(vt, "openid profile workspace.todo:read"); got != "profile workspace.todo:read" { + token := &oauthToken{Scope: "profile workspace.todo:read"} + if got := effectiveTokenScope(token, "openid profile workspace.todo:read"); got != "profile workspace.todo:read" { t.Fatalf("effective scope = %q", got) } - if err := requireIDTokenForOpenID(effectiveTokenScope(vt, "openid profile"), ""); err != nil { - t.Fatalf("reduced grant without openid should not require an ID token: %v", err) - } } func TestRefreshTokenDoesNotRequireNewIDToken(t *testing.T) { @@ -299,7 +374,7 @@ func TestRefreshTokenDoesNotRequireNewIDToken(t *testing.T) { cred, err := RefreshToken(server.URL, &config.Credential{ ClientID: "client-1", RefreshToken: "refresh-1", - Scope: "openid profile workspace.todo:read", + Scope: "offline_access workspace.todo:read", }) if err != nil { t.Fatal(err) @@ -307,46 +382,35 @@ func TestRefreshTokenDoesNotRequireNewIDToken(t *testing.T) { if cred.AccessToken != "next-access" || cred.RefreshToken != "refresh-1" { t.Fatalf("credential = %#v", cred) } - if cred.Scope != "openid profile workspace.todo:read" { + if cred.Scope != "offline_access workspace.todo:read" { t.Fatalf("scope = %q", cred.Scope) } } -func TestValidateIDTokenAudienceIsClientID(t *testing.T) { - key, err := rsa.GenerateKey(rand.Reader, 2048) - if err != nil { - t.Fatal(err) - } - signer, err := jose.NewSigner(jose.SigningKey{Algorithm: jose.RS256, Key: key}, nil) +func TestOAuthAccessTokenDoesNotRequireIDToken(t *testing.T) { + token := (&oauth2.Token{ + AccessToken: "device-access", + RefreshToken: "device-refresh", + TokenType: "Bearer", + }).WithExtra(map[string]any{ + "expires_in": 3600, + "scope": "openid profile workspace.todo:read", + }) + cred, err := oauthTokenToCredential( + "device-client", + "https://life.example/api/auth", + newOAuthToken(token), + "", + "", + time.Now(), + ) if err != nil { t.Fatal(err) } - build := func(aud any) string { - t.Helper() - claims := map[string]any{ - "iss": "https://issuer.test", - "aud": aud, - "exp": time.Now().Add(time.Hour).Unix(), - } - s, err := jwt.Signed(signer).Claims(claims).Serialize() - if err != nil { - t.Fatal(err) - } - return s - } - - vt := &VerifiedToken{IDToken: build("client-id-123")} - if err := vt.ValidateIDToken("https://issuer.test", "client-id-123"); err != nil { - t.Fatalf("expected client_id audience to validate: %v", err) - } - - vt.IDToken = build("https://server.test") - if err := vt.ValidateIDToken("https://issuer.test", "client-id-123"); err == nil { - t.Fatal("expected server URL audience to fail against client_id expectation") + if cred.AccessToken != "device-access" || cred.RefreshToken != "device-refresh" { + t.Fatalf("credential = %#v", cred) } - - vt.IDToken = build([]string{"client-id-123", "other"}) - if err := vt.ValidateIDToken("https://issuer.test", "client-id-123"); err != nil { - t.Fatalf("expected audience list containing client_id to validate: %v", err) + if cred.Scope != "openid profile workspace.todo:read" { + t.Fatalf("scope = %q", cred.Scope) } } diff --git a/internal/auth/device.go b/internal/auth/device.go index 8385f78..38b4887 100644 --- a/internal/auth/device.go +++ b/internal/auth/device.go @@ -93,16 +93,6 @@ func LoginDeviceCode(server string) (*config.Credential, error) { return nil, fmt.Errorf("device authorization failed: %w", err) } - vt := newVerifiedToken(tok) - if err := requireIDTokenForOpenID(effectiveTokenScope(vt, scope), vt.IDToken); err != nil { - return nil, err - } - issuer := stringFromMap(meta, "issuer") - if issuer == "" { - issuer = server - } - if err := vt.ValidateIDToken(issuer, clientID); err != nil { - return nil, err - } - return verifiedTokenToCredential(clientID, resource, vt, "", scope, time.Now()) + token := newOAuthToken(tok) + return oauthTokenToCredential(clientID, resource, token, "", scope, time.Now()) } diff --git a/internal/auth/oauth.go b/internal/auth/oauth.go index c5894b3..086edaa 100644 --- a/internal/auth/oauth.go +++ b/internal/auth/oauth.go @@ -3,41 +3,35 @@ package auth import ( "context" "errors" - "fmt" "net/http" "strconv" "strings" "time" "github.com/Life-USTC/CLI/internal/config" - "github.com/go-jose/go-jose/v4" - "github.com/go-jose/go-jose/v4/jwt" "golang.org/x/oauth2" ) -// VerifiedToken wraps an oauth2.Token and preserves the optional ID token. -type VerifiedToken struct { +type oauthToken struct { AccessToken string RefreshToken string TokenType string Expiry time.Time ExpiresIn int Scope string - IDToken string } -func newVerifiedToken(tok *oauth2.Token) *VerifiedToken { +func newOAuthToken(tok *oauth2.Token) *oauthToken { if tok == nil { return nil } - return &VerifiedToken{ + return &oauthToken{ AccessToken: tok.AccessToken, RefreshToken: tok.RefreshToken, TokenType: tok.TokenType, Expiry: tok.Expiry, ExpiresIn: tokenExpiresIn(tok, 0), Scope: tokenExtraString(tok, "scope"), - IDToken: tokenExtraString(tok, "id_token"), } } @@ -51,71 +45,6 @@ func tokenExtraString(tok *oauth2.Token, key string) string { return "" } -// ValidateIDToken checks the ID token's issuer and audience claims. -// It does not verify the JWT signature; callers should fetch the issuer's -// JWKS and verify the signature when required. -// Empty issuer or audience is treated as an error so checks are never -// silently skipped. -func (t *VerifiedToken) ValidateIDToken(issuer, audience string) error { - if t == nil || t.IDToken == "" { - return nil - } - if issuer == "" { - return errors.New("id_token issuer required") - } - if audience == "" { - return errors.New("id_token audience required") - } - parsed, err := jwt.ParseSigned(t.IDToken, []jose.SignatureAlgorithm{jose.RS256, jose.ES256, jose.EdDSA}) - if err != nil { - return fmt.Errorf("invalid id_token: %w", err) - } - claims := map[string]any{} - if err := parsed.UnsafeClaimsWithoutVerification(&claims); err != nil { - return fmt.Errorf("invalid id_token claims: %w", err) - } - if iss, _ := claims["iss"].(string); strings.TrimSpace(iss) != issuer { - return fmt.Errorf("invalid issuer %q, expected %q", iss, issuer) - } - if !audienceMatches(claims["aud"], audience) { - return fmt.Errorf("invalid audience, expected %q", audience) - } - if exp, ok := expiresAtFromClaim(claims["exp"]); ok && !exp.After(time.Now()) { - return errors.New("id_token expired") - } - return nil -} - -func audienceMatches(audClaim any, expected string) bool { - if s, ok := audClaim.(string); ok { - return strings.TrimSpace(s) == expected - } - if list, ok := audClaim.([]any); ok { - for _, item := range list { - if s, ok := item.(string); ok && strings.TrimSpace(s) == expected { - return true - } - } - } - return false -} - -func expiresAtFromClaim(expClaim any) (time.Time, bool) { - switch v := expClaim.(type) { - case float64: - return time.Unix(int64(v), 0), true - case int: - return time.Unix(int64(v), 0), true - case int64: - return time.Unix(v, 0), true - case string: - if n, err := parseIntString(v); err == nil { - return time.Unix(n, 0), true - } - } - return time.Time{}, false -} - func parseIntString(s string) (int64, error) { s = strings.TrimSpace(s) if s == "" { @@ -124,54 +53,35 @@ func parseIntString(s string) (int64, error) { return strconv.ParseInt(s, 10, 64) } -// scopeIncludes reports whether a space-delimited OAuth scope contains value. -func scopeIncludes(scope, value string) bool { - for _, s := range strings.Fields(scope) { - if s == value { - return true - } - } - return false -} - -// requireIDTokenForOpenID returns an error when the openid scope was requested -// but the token response does not contain an ID token. -func requireIDTokenForOpenID(scope, idToken string) error { - if scopeIncludes(scope, "openid") && strings.TrimSpace(idToken) == "" { - return errors.New("openid scope requested but token response missing id_token") - } - return nil -} - -func effectiveTokenScope(token *VerifiedToken, fallback string) string { +func effectiveTokenScope(token *oauthToken, fallback string) string { if token != nil && strings.TrimSpace(token.Scope) != "" { return strings.TrimSpace(token.Scope) } return strings.TrimSpace(fallback) } -func verifiedTokenToCredential(clientID, resource string, vt *VerifiedToken, fallbackRefresh, fallbackScope string, now time.Time) (*config.Credential, error) { - if vt == nil { +func oauthTokenToCredential(clientID, resource string, token *oauthToken, fallbackRefresh, fallbackScope string, now time.Time) (*config.Credential, error) { + if token == nil { return nil, errors.New("token response is nil") } - accessToken := strings.TrimSpace(vt.AccessToken) + accessToken := strings.TrimSpace(token.AccessToken) if accessToken == "" { return nil, errors.New("token response missing access_token") } - expiresIn := vt.ExpiresIn + expiresIn := token.ExpiresIn if expiresIn <= 0 { expiresIn = 3600 } - refreshToken := strings.TrimSpace(vt.RefreshToken) + refreshToken := strings.TrimSpace(token.RefreshToken) if refreshToken == "" { refreshToken = fallbackRefresh } - scope := effectiveTokenScope(vt, fallbackScope) + scope := effectiveTokenScope(token, fallbackScope) return &config.Credential{ ClientID: clientID, AccessToken: accessToken, RefreshToken: refreshToken, - TokenType: strings.TrimSpace(vt.TokenType), + TokenType: strings.TrimSpace(token.TokenType), ExpiresAt: float64(now.Add(time.Duration(expiresIn) * time.Second).Unix()), Scope: scope, Resource: resource,