diff --git a/README.md b/README.md index 3018a17..8b5dd5e 100644 --- a/README.md +++ b/README.md @@ -52,6 +52,15 @@ http.ListenAndServe(":8080", handler) - **Fast token caching** - 5-min cache, <5ms validation - **Production ready** - Security hardened, battle-tested - **Multiple providers** - HMAC, Okta, Google, Azure AD +- **Headless service tokens** - Optional asymmetric JWTs for non-interactive agents without changing your OAuth provider + +--- + +## Service Tokens for Agents + +Remote MCP servers can opt in to service-token auth for headless agents while keeping the existing human OAuth flow. Service tokens are asymmetric JWTs verified with a public key; the MCP server never needs the private key. + +See [Service Tokens](docs/SERVICE-TOKENS.md) for the token contract, config fields, and rotation notes. --- diff --git a/config.go b/config.go index 89f3745..cb205ec 100644 --- a/config.go +++ b/config.go @@ -1,6 +1,7 @@ package oauth import ( + "encoding/base64" "fmt" "github.com/Vungle/oauth-mcp-proxy/provider" @@ -28,6 +29,13 @@ type Config struct { // Security JWTSecret []byte // For HMAC provider and state signing + // Service token settings for non-interactive clients + ServiceTokenEnabled bool + ServiceTokenIssuer string + ServiceTokenAudience string + ServiceTokenPublicKeyPEM string + ServiceTokenSubjectPrefix string + // Optional - Logging // Logger allows custom logging implementation. If nil, uses default logger // that outputs to log.Printf with level prefixes ([INFO], [ERROR], etc.). @@ -76,6 +84,24 @@ func (c *Config) Validate() error { return fmt.Errorf("audience is required") } + if c.ServiceTokenEnabled { + if c.ServiceTokenSubjectPrefix == "" { + c.ServiceTokenSubjectPrefix = "svc-" + } + if c.ServiceTokenIssuer == "" { + return fmt.Errorf("service token issuer is required when service token auth is enabled") + } + if c.ServiceTokenAudience == "" { + return fmt.Errorf("service token audience is required when service token auth is enabled") + } + if c.ServiceTokenPublicKeyPEM == "" { + return fmt.Errorf("service token public key PEM is required when service token auth is enabled") + } + if c.Issuer != "" && c.ServiceTokenIssuer == c.Issuer { + return fmt.Errorf("service token issuer must be distinct from OAuth issuer") + } + } + // Validate proxy mode requirements if c.Mode == "proxy" { if c.ClientID == "" { @@ -122,11 +148,15 @@ func SetupOAuth(cfg *Config) (provider.TokenValidator, error) { func createValidator(cfg *Config, logger Logger) (provider.TokenValidator, error) { // Convert root Config to provider.Config providerCfg := &provider.Config{ - Provider: cfg.Provider, - Issuer: cfg.Issuer, - Audience: cfg.Audience, - JWTSecret: cfg.JWTSecret, - Logger: logger, + Provider: cfg.Provider, + Issuer: cfg.Issuer, + Audience: cfg.Audience, + JWTSecret: cfg.JWTSecret, + ServiceTokenIssuer: cfg.ServiceTokenIssuer, + ServiceTokenAudience: cfg.ServiceTokenAudience, + ServiceTokenPublicKeyPEM: cfg.ServiceTokenPublicKeyPEM, + ServiceTokenSubjectPrefix: cfg.ServiceTokenSubjectPrefix, + Logger: logger, } var validator provider.TokenValidator @@ -143,6 +173,14 @@ func createValidator(cfg *Config, logger Logger) (provider.TokenValidator, error return nil, err } + if cfg.ServiceTokenEnabled { + serviceTokenValidator := &provider.ServiceTokenValidator{} + if err := serviceTokenValidator.Initialize(providerCfg); err != nil { + return nil, err + } + validator = provider.NewRoutingValidator(validator, serviceTokenValidator, cfg.ServiceTokenIssuer) + } + return validator, nil } @@ -227,6 +265,16 @@ func (b *ConfigBuilder) WithJWTSecret(secret []byte) *ConfigBuilder { return b } +// WithServiceToken enables asymmetric service-token validation. +func (b *ConfigBuilder) WithServiceToken(issuer, audience, publicKeyPEM, subjectPrefix string) *ConfigBuilder { + b.config.ServiceTokenEnabled = true + b.config.ServiceTokenIssuer = issuer + b.config.ServiceTokenAudience = audience + b.config.ServiceTokenPublicKeyPEM = publicKeyPEM + b.config.ServiceTokenSubjectPrefix = subjectPrefix + return b +} + // WithLogger sets the logger func (b *ConfigBuilder) WithLogger(logger Logger) *ConfigBuilder { b.config.Logger = logger @@ -290,8 +338,11 @@ func FromEnv() (*Config, error) { } jwtSecret := getEnv("JWT_SECRET", "") + serviceTokenEnabled := getEnv("SERVICE_TOKEN_ENABLED", "") == "true" || + getEnv("SERVICE_TOKEN_ENABLED", "") == "1" || + getEnv("SERVICE_TOKEN_ENABLED", "") == "yes" - return NewConfigBuilder(). + builder := NewConfigBuilder(). WithMode(getEnv("OAUTH_MODE", "")). WithProvider(getEnv("OAUTH_PROVIDER", "")). WithRedirectURIs(getEnv("OAUTH_REDIRECT_URIS", "")). @@ -301,6 +352,33 @@ func FromEnv() (*Config, error) { WithClientID(getEnv("OIDC_CLIENT_ID", "")). WithClientSecret(getEnv("OIDC_CLIENT_SECRET", "")). WithServerURL(serverURL). - WithJWTSecret([]byte(jwtSecret)). - Build() + WithJWTSecret([]byte(jwtSecret)) + + if serviceTokenEnabled { + publicKeyPEM := getEnv("SERVICE_TOKEN_PUBLIC_KEY_PEM", "") + if publicKeyPEM == "" { + publicKeyPEM = decodeBase64Env("SERVICE_TOKEN_PUBLIC_KEY_PEM_B64") + } + + builder.WithServiceToken( + getEnv("SERVICE_TOKEN_ISSUER", ""), + getEnv("SERVICE_TOKEN_AUDIENCE", ""), + publicKeyPEM, + getEnv("SERVICE_TOKEN_SUBJECT_PREFIX", ""), + ) + } + + return builder.Build() +} + +func decodeBase64Env(key string) string { + value := getEnv(key, "") + if value == "" { + return "" + } + decoded, err := base64.StdEncoding.DecodeString(value) + if err != nil { + return "" + } + return string(decoded) } diff --git a/docs/CONFIGURATION.md b/docs/CONFIGURATION.md index d67d5c6..ff59b50 100644 --- a/docs/CONFIGURATION.md +++ b/docs/CONFIGURATION.md @@ -25,6 +25,13 @@ type Config struct { ServerURL string // Your server's public URL RedirectURIs string // Allowed redirect URIs + // Optional - Service tokens for non-interactive clients + ServiceTokenEnabled bool + ServiceTokenIssuer string + ServiceTokenAudience string + ServiceTokenPublicKeyPEM string + ServiceTokenSubjectPrefix string // default: "svc-" + // Optional - Logging Logger Logger // Custom logger implementation } @@ -93,6 +100,12 @@ _, oauthOption, _ := oauth.WithOAuth(mux, cfg) - `MCP_PORT` - Server port (default: 8080) - `HTTPS_CERT_FILE` - TLS cert file (enables HTTPS) - `HTTPS_KEY_FILE` - TLS key file (enables HTTPS) +- `SERVICE_TOKEN_ENABLED` - Enables asymmetric service-token validation +- `SERVICE_TOKEN_ISSUER` - Expected service-token issuer +- `SERVICE_TOKEN_AUDIENCE` - Expected service-token audience +- `SERVICE_TOKEN_PUBLIC_KEY_PEM` - Public key PEM for service-token verification +- `SERVICE_TOKEN_PUBLIC_KEY_PEM_B64` - Base64-encoded public key PEM +- `SERVICE_TOKEN_SUBJECT_PREFIX` - Required service-token subject prefix (default: `svc-`) **Benefits:** @@ -118,6 +131,28 @@ Provider: "okta" // Use Okta OIDC validation **See:** [Provider Guides](providers/) for setup instructions +## Optional Service Tokens + +Service-token auth runs alongside the configured OAuth provider. It lets headless agents authenticate with an asymmetric JWT instead of an interactive browser OAuth flow. + +```go +cfg := &oauth.Config{ + Provider: "okta", + Issuer: "https://company.okta.com", + Audience: "https://company.okta.com", + + ServiceTokenEnabled: true, + ServiceTokenIssuer: "agent-auth-service", + ServiceTokenAudience: "api://example-mcp-server", + ServiceTokenPublicKeyPEM: publicKeyPEM, + ServiceTokenSubjectPrefix: "svc-", +} +``` + +Use `SERVICE_TOKEN_PUBLIC_KEY_PEM_B64` for Kubernetes/Terraform values to avoid multiline PEM escaping issues. + +**See:** [Service Tokens](SERVICE-TOKENS.md) for token claims, supported algorithms, and key handling. + ### Audience **Type:** `string` diff --git a/docs/SERVICE-TOKENS.md b/docs/SERVICE-TOKENS.md new file mode 100644 index 0000000..45fd802 --- /dev/null +++ b/docs/SERVICE-TOKENS.md @@ -0,0 +1,74 @@ +# Service Tokens + +Service tokens let non-interactive agents call an MCP server without a browser OAuth flow. They are optional and run alongside the existing OAuth/OIDC provider. + +## Token Contract + +Service tokens are asymmetric JWTs. The MCP server verifies them with a public key; it never receives the private key. + +Required claims: + +- `iss`: must match `ServiceTokenIssuer` +- `aud`: must match `ServiceTokenAudience` +- `sub`: must start with `ServiceTokenSubjectPrefix` (default: `svc-`) +- `exp`: required; expired tokens are rejected + +Supported signing algorithms: + +- `EdDSA` with an Ed25519 public key +- `RS256` with an RSA public key + +Do not use HMAC/HS256 for service tokens. HMAC requires the MCP server to hold a shared secret that can also mint tokens. + +## Configuration + +```go +oauthServer, oauthOption, err := mark3labs.WithOAuth(mux, &oauth.Config{ + Provider: "okta", + Issuer: "https://your-company.okta.com", + Audience: "https://your-company.okta.com", + + ServiceTokenEnabled: true, + ServiceTokenIssuer: "agent-auth-service", + ServiceTokenAudience: "api://example-mcp-server", + ServiceTokenPublicKeyPEM: publicKeyPEM, + ServiceTokenSubjectPrefix: "svc-", +}) +``` + +`ServiceTokenPublicKeyPEM` should contain only a public key. Keep the matching private key in a separate, restricted secret store used only by the token minting workflow. + +## Environment Variables + +`FromEnv()` also supports: + +```text +SERVICE_TOKEN_ENABLED=true +SERVICE_TOKEN_ISSUER=agent-auth-service +SERVICE_TOKEN_AUDIENCE=api://example-mcp-server +SERVICE_TOKEN_PUBLIC_KEY_PEM= +SERVICE_TOKEN_PUBLIC_KEY_PEM_B64= +SERVICE_TOKEN_SUBJECT_PREFIX=svc- +``` + +Use `SERVICE_TOKEN_PUBLIC_KEY_PEM_B64` for Kubernetes and Terraform values to avoid multiline PEM escaping issues. + +## Validation Flow + +The validator first peeks at the unsigned `iss` claim only to choose the validator: + +- `iss == ServiceTokenIssuer`: validate with the service-token public key and required claims. +- Any other issuer: validate with the configured OAuth/OIDC provider. + +The peeked claim is not trusted for authorization. If a token claims the service issuer but fails service-token validation, it is rejected and does not fall back to the OAuth provider. + +## Key Rotation + +First version supports one active public key. Rotation is: + +1. Mint new tokens with a new private key. +2. Deploy the matching new public key to MCP servers. +3. Replace agent tokens. +4. Stop accepting the old public key. + +Keep service-token TTLs short enough that this rotation window is acceptable. diff --git a/middleware.go b/middleware.go index b223f3a..8488eb6 100644 --- a/middleware.go +++ b/middleware.go @@ -60,13 +60,14 @@ func (s *Server) Middleware() func(server.ToolHandlerFunc) server.ToolHandlerFun return nil, fmt.Errorf("authentication failed: %w", err) } - // Cache the validation result (expire in 5 minutes) - expiresAt := time.Now().Add(5 * time.Minute) - s.cache.setCachedToken(tokenHash, user, expiresAt) + // Cache the validation result, but never beyond the token's own expiry. + if expiresAt, ok := cacheExpiresAtForToken(tokenString, time.Now()); ok { + s.cache.setCachedToken(tokenHash, user, expiresAt) + } // Add user to context for downstream handlers ctx = context.WithValue(ctx, userContextKey, user) - s.logger.Info("Authenticated user %s for tool: %s (cached for 5 minutes)", user.Username, req.Params.Name) + s.logger.Info("Authenticated user %s for tool: %s", user.Username, req.Params.Name) return next(ctx, req) } diff --git a/oauth.go b/oauth.go index f8673c2..9dbc4c3 100644 --- a/oauth.go +++ b/oauth.go @@ -10,9 +10,12 @@ import ( "time" "github.com/Vungle/oauth-mcp-proxy/provider" + "github.com/golang-jwt/jwt/v5" mcpserver "github.com/mark3labs/mcp-go/server" ) +const tokenCacheTTL = 5 * time.Minute + // Server represents an OAuth authentication server instance. // Each Server maintains its own token cache and validator, allowing // multiple independent OAuth configurations in the same application. @@ -126,13 +129,38 @@ func (s *Server) ValidateTokenCached(ctx context.Context, token string) (*User, return nil, fmt.Errorf("authentication failed: %w", err) } - expiresAt := time.Now().Add(5 * time.Minute) - s.cache.setCachedToken(tokenHash, user, expiresAt) + now := time.Now() + if expiresAt, ok := cacheExpiresAtForToken(token, now); ok { + s.cache.setCachedToken(tokenHash, user, expiresAt) + s.logger.Info("Authenticated user %s (cached until %s)", user.Username, expiresAt.Format(time.RFC3339)) + } else { + s.logger.Info("Authenticated user %s (not cached because token is expired)", user.Username) + } - s.logger.Info("Authenticated user %s (cached for 5 minutes)", user.Username) return user, nil } +func cacheExpiresAtForToken(tokenString string, now time.Time) (time.Time, bool) { + expiresAt := now.Add(tokenCacheTTL) + + tokenString = strings.TrimPrefix(tokenString, "Bearer ") + claims := jwt.RegisteredClaims{} + if _, _, err := jwt.NewParser().ParseUnverified(tokenString, &claims); err != nil { + return expiresAt, true + } + if claims.ExpiresAt == nil { + return expiresAt, true + } + if !claims.ExpiresAt.After(now) { + return time.Time{}, false + } + if claims.ExpiresAt.Before(expiresAt) { + return claims.ExpiresAt.Time, true + } + + return expiresAt, true +} + // GetAuthorizationServerMetadataURL returns the OAuth 2.0 authorization server metadata URL func (s *Server) GetAuthorizationServerMetadataURL() string { return fmt.Sprintf("%s/.well-known/oauth-authorization-server", s.config.ServerURL) diff --git a/provider/provider.go b/provider/provider.go index 39c5847..40b9504 100644 --- a/provider/provider.go +++ b/provider/provider.go @@ -2,7 +2,11 @@ package provider import ( "context" + "crypto/ed25519" + "crypto/rsa" "crypto/tls" + "crypto/x509" + "encoding/pem" "fmt" "net/http" "strings" @@ -30,11 +34,15 @@ type Logger interface { // Config holds OAuth configuration (subset needed by provider) type Config struct { - Provider string - Issuer string - Audience string - JWTSecret []byte - Logger Logger + Provider string + Issuer string + Audience string + JWTSecret []byte + ServiceTokenIssuer string + ServiceTokenAudience string + ServiceTokenPublicKeyPEM string + ServiceTokenSubjectPrefix string + Logger Logger } // TokenValidator interface for OAuth token validation @@ -58,6 +66,170 @@ type OIDCValidator struct { logger Logger } +// ServiceTokenValidator validates asymmetric JWTs for non-interactive service access. +type ServiceTokenValidator struct { + issuer string + audience string + subjectPrefix string + publicKey any +} + +// RoutingValidator routes service-token JWTs to ServiceTokenValidator and all +// other tokens to the primary OAuth provider validator. +type RoutingValidator struct { + primary TokenValidator + serviceToken TokenValidator + serviceIssuer string +} + +// NewRoutingValidator returns a validator that supports both interactive OAuth +// provider tokens and non-interactive service tokens. +func NewRoutingValidator(primary TokenValidator, serviceToken TokenValidator, serviceIssuer string) TokenValidator { + return &RoutingValidator{ + primary: primary, + serviceToken: serviceToken, + serviceIssuer: serviceIssuer, + } +} + +// ValidateToken validates service-token issuer JWTs with the service-token +// validator and all other tokens with the primary validator. +func (v *RoutingValidator) ValidateToken(ctx context.Context, tokenString string) (*User, error) { + if peekJWTIssuer(tokenString) == v.serviceIssuer { + return v.serviceToken.ValidateToken(ctx, tokenString) + } + return v.primary.ValidateToken(ctx, tokenString) +} + +// Initialize is not used; child validators are initialized before routing. +func (v *RoutingValidator) Initialize(cfg *Config) error { + return nil +} + +// Initialize sets up the asymmetric service-token validator. +func (v *ServiceTokenValidator) Initialize(cfg *Config) error { + v.issuer = cfg.ServiceTokenIssuer + v.audience = cfg.ServiceTokenAudience + v.subjectPrefix = cfg.ServiceTokenSubjectPrefix + + if v.issuer == "" { + return fmt.Errorf("service token issuer is required") + } + if v.audience == "" { + return fmt.Errorf("service token audience is required") + } + if v.subjectPrefix == "" { + return fmt.Errorf("service token subject prefix is required") + } + + publicKey, err := parseServiceTokenPublicKey(cfg.ServiceTokenPublicKeyPEM) + if err != nil { + return err + } + v.publicKey = publicKey + return nil +} + +// ValidateToken validates an EdDSA or RS256 service JWT using the configured public key. +func (v *ServiceTokenValidator) ValidateToken(ctx context.Context, tokenString string) (*User, error) { + tokenString = strings.TrimPrefix(tokenString, "Bearer ") + + claims := &struct { + PreferredUsername string `json:"preferred_username"` + Email string `json:"email"` + jwt.RegisteredClaims + }{} + + token, err := jwt.ParseWithClaims( + tokenString, + claims, + func(token *jwt.Token) (interface{}, error) { + switch token.Method.Alg() { + case jwt.SigningMethodEdDSA.Alg(): + if key, ok := v.publicKey.(ed25519.PublicKey); ok { + return key, nil + } + case jwt.SigningMethodRS256.Alg(): + if key, ok := v.publicKey.(*rsa.PublicKey); ok { + return key, nil + } + } + return nil, fmt.Errorf("unexpected signing method: %s", token.Method.Alg()) + }, + jwt.WithIssuer(v.issuer), + jwt.WithAudience(v.audience), + jwt.WithExpirationRequired(), + jwt.WithIssuedAt(), + ) + if err != nil { + return nil, fmt.Errorf("failed to parse and validate service token: %w", err) + } + if !token.Valid { + return nil, fmt.Errorf("invalid service token") + } + if claims.Subject == "" { + return nil, fmt.Errorf("missing subject in service token") + } + if !strings.HasPrefix(claims.Subject, v.subjectPrefix) { + return nil, fmt.Errorf("service token subject must start with %q", v.subjectPrefix) + } + + username := claims.PreferredUsername + if username == "" { + username = claims.Subject + } + + return &User{ + Subject: claims.Subject, + Username: username, + Email: claims.Email, + }, nil +} + +func parseServiceTokenPublicKey(publicKeyPEM string) (any, error) { + if strings.TrimSpace(publicKeyPEM) == "" { + return nil, fmt.Errorf("service token public key PEM is required") + } + + block, _ := pem.Decode([]byte(publicKeyPEM)) + if block == nil { + return nil, fmt.Errorf("failed to decode service token public key PEM") + } + + if block.Type == "RSA PUBLIC KEY" { + key, err := x509.ParsePKCS1PublicKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("failed to parse RSA service token public key: %w", err) + } + return key, nil + } + + key, err := x509.ParsePKIXPublicKey(block.Bytes) + if err != nil { + return nil, fmt.Errorf("failed to parse service token public key: %w", err) + } + + switch key := key.(type) { + case ed25519.PublicKey: + return key, nil + case *rsa.PublicKey: + return key, nil + default: + return nil, fmt.Errorf("unsupported service token public key type %T", key) + } +} + +func peekJWTIssuer(tokenString string) string { + tokenString = strings.TrimPrefix(tokenString, "Bearer ") + + parser := jwt.NewParser() + claims := jwt.MapClaims{} + if _, _, err := parser.ParseUnverified(tokenString, claims); err != nil { + return "" + } + return getStringClaim(claims, "iss") +} + // Initialize sets up the HMAC validator with JWT secret and audience func (v *HMACValidator) Initialize(cfg *Config) error { v.secretOnce.Do(func() { diff --git a/service_token_test.go b/service_token_test.go new file mode 100644 index 0000000..82d4480 --- /dev/null +++ b/service_token_test.go @@ -0,0 +1,359 @@ +package oauth + +import ( + "context" + "crypto/ed25519" + "crypto/rand" + "crypto/rsa" + "crypto/x509" + "encoding/base64" + "encoding/pem" + "strings" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +const ( + testPrimaryAudience = "api://primary" + testServiceIssuer = "agent-auth-service" + testServiceAudience = "api://example-mcp-server" +) + +func TestServiceTokenValidation(t *testing.T) { + publicKeyPEM, privateKey := generateEd25519KeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + token := signServiceToken(t, privateKey, jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "preferred_username": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }) + + user, err := server.ValidateTokenCached(context.Background(), token) + if err != nil { + t.Fatalf("ValidateTokenCached() error = %v", err) + } + if user.Subject != "svc-example-agent" { + t.Fatalf("Subject = %q, want svc-example-agent", user.Subject) + } + if user.Username != "svc-example-agent" { + t.Fatalf("Username = %q, want svc-example-agent", user.Username) + } +} + +func TestServiceTokenValidationRS256(t *testing.T) { + publicKeyPEM, privateKey := generateRSAKeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }) + tokenString, err := token.SignedString(privateKey) + if err != nil { + t.Fatalf("SignedString() error = %v", err) + } + + user, err := server.ValidateTokenCached(context.Background(), tokenString) + if err != nil { + t.Fatalf("ValidateTokenCached() error = %v", err) + } + if user.Subject != "svc-example-agent" { + t.Fatalf("Subject = %q, want svc-example-agent", user.Subject) + } +} + +func TestServiceTokenCacheDoesNotOutliveExpiration(t *testing.T) { + publicKeyPEM, privateKey := generateEd25519KeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + exp := time.Now().Add(2 * time.Second).Unix() + token := signServiceToken(t, privateKey, jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": exp, + }) + + if _, err := server.ValidateTokenCached(context.Background(), token); err != nil { + t.Fatalf("first ValidateTokenCached() error = %v", err) + } + + if wait := time.Until(time.Unix(exp, 0).Add(100 * time.Millisecond)); wait > 0 { + time.Sleep(wait) + } + + if _, err := server.ValidateTokenCached(context.Background(), token); err == nil { + t.Fatal("second ValidateTokenCached() error = nil, want expired token error") + } +} + +func TestServiceTokenRejectsInvalidTokens(t *testing.T) { + publicKeyPEM, privateKey := generateEd25519KeyPair(t) + _, wrongPrivateKey := generateEd25519KeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + tests := []struct { + name string + claims jwt.MapClaims + key ed25519.PrivateKey + }{ + { + name: "wrong signature", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: wrongPrivateKey, + }, + { + name: "wrong issuer", + claims: jwt.MapClaims{ + "iss": "other-service", + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "wrong audience", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": "api://other", + "sub": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "expired", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": time.Now().Add(-time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "missing expiry", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + }, + key: privateKey, + }, + { + name: "invalid subject prefix", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "future issued at", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "iat": time.Now().Add(time.Hour).Unix(), + "exp": time.Now().Add(2 * time.Hour).Unix(), + }, + key: privateKey, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + token := signServiceToken(t, tt.key, tt.claims) + if _, err := server.ValidateTokenCached(context.Background(), token); err == nil { + t.Fatal("ValidateTokenCached() error = nil, want error") + } + }) + } +} + +func TestServiceTokenRejectsSymmetricAlgorithm(t *testing.T) { + publicKeyPEM, _ := generateEd25519KeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-example-agent", + "exp": time.Now().Add(time.Hour).Unix(), + }) + tokenString, err := token.SignedString([]byte("shared-secret")) + if err != nil { + t.Fatalf("SignedString() error = %v", err) + } + + if _, err := server.ValidateTokenCached(context.Background(), tokenString); err == nil { + t.Fatal("ValidateTokenCached() error = nil, want error") + } +} + +func TestServiceTokenRoutingFallsBackToPrimaryValidator(t *testing.T) { + publicKeyPEM, _ := generateEd25519KeyPair(t) + server := newServiceTokenTestServer(t, publicKeyPEM) + + token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ + "iss": "human-issuer", + "aud": testPrimaryAudience, + "sub": "user-123", + "preferred_username": "user@example.com", + "exp": time.Now().Add(time.Hour).Unix(), + }) + tokenString, err := token.SignedString([]byte("primary-secret")) + if err != nil { + t.Fatalf("SignedString() error = %v", err) + } + + user, err := server.ValidateTokenCached(context.Background(), tokenString) + if err != nil { + t.Fatalf("ValidateTokenCached() error = %v", err) + } + if user.Subject != "user-123" { + t.Fatalf("Subject = %q, want user-123", user.Subject) + } +} + +func TestServiceTokenConfigValidation(t *testing.T) { + cfg := &Config{ + Provider: "hmac", + Audience: testPrimaryAudience, + JWTSecret: []byte("primary-secret"), + ServiceTokenEnabled: true, + ServiceTokenIssuer: testServiceIssuer, + ServiceTokenAudience: testServiceAudience, + } + + err := cfg.Validate() + if err == nil { + t.Fatal("Validate() error = nil, want error") + } + if !strings.Contains(err.Error(), "service token public key PEM is required") { + t.Fatalf("Validate() error = %v, want missing public key error", err) + } +} + +func TestServiceTokenIssuerMustBeDistinctFromOAuthIssuer(t *testing.T) { + publicKeyPEM, _ := generateEd25519KeyPair(t) + cfg := &Config{ + Provider: "okta", + Issuer: testServiceIssuer, + Audience: testPrimaryAudience, + ServiceTokenEnabled: true, + ServiceTokenIssuer: testServiceIssuer, + ServiceTokenAudience: testServiceAudience, + ServiceTokenPublicKeyPEM: publicKeyPEM, + } + + err := cfg.Validate() + if err == nil { + t.Fatal("Validate() error = nil, want issuer collision error") + } + if !strings.Contains(err.Error(), "service token issuer must be distinct") { + t.Fatalf("Validate() error = %v, want issuer collision error", err) + } +} + +func TestFromEnvServiceTokenPublicKeyBase64(t *testing.T) { + publicKeyPEM, _ := generateEd25519KeyPair(t) + + t.Setenv("OAUTH_PROVIDER", "hmac") + t.Setenv("OIDC_AUDIENCE", testPrimaryAudience) + t.Setenv("JWT_SECRET", "primary-secret") + t.Setenv("SERVICE_TOKEN_ENABLED", "true") + t.Setenv("SERVICE_TOKEN_ISSUER", testServiceIssuer) + t.Setenv("SERVICE_TOKEN_AUDIENCE", testServiceAudience) + t.Setenv("SERVICE_TOKEN_PUBLIC_KEY_PEM_B64", base64.StdEncoding.EncodeToString([]byte(publicKeyPEM))) + + cfg, err := FromEnv() + if err != nil { + t.Fatalf("FromEnv() error = %v", err) + } + if cfg.ServiceTokenPublicKeyPEM != publicKeyPEM { + t.Fatal("ServiceTokenPublicKeyPEM was not decoded from SERVICE_TOKEN_PUBLIC_KEY_PEM_B64") + } + if cfg.ServiceTokenSubjectPrefix != "svc-" { + t.Fatalf("ServiceTokenSubjectPrefix = %q, want svc-", cfg.ServiceTokenSubjectPrefix) + } +} + +func generateEd25519KeyPair(t *testing.T) (string, ed25519.PrivateKey) { + t.Helper() + + publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatalf("GenerateKey() error = %v", err) + } + + der, err := x509.MarshalPKIXPublicKey(publicKey) + if err != nil { + t.Fatalf("MarshalPKIXPublicKey() error = %v", err) + } + + block := &pem.Block{Type: "PUBLIC KEY", Bytes: der} + return string(pem.EncodeToMemory(block)), privateKey +} + +func generateRSAKeyPair(t *testing.T) (string, *rsa.PrivateKey) { + t.Helper() + + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + if err != nil { + t.Fatalf("GenerateKey() error = %v", err) + } + + der, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey) + if err != nil { + t.Fatalf("MarshalPKIXPublicKey() error = %v", err) + } + + block := &pem.Block{Type: "PUBLIC KEY", Bytes: der} + return string(pem.EncodeToMemory(block)), privateKey +} + +func signServiceToken(t *testing.T, privateKey ed25519.PrivateKey, claims jwt.MapClaims) string { + t.Helper() + + token := jwt.NewWithClaims(jwt.SigningMethodEdDSA, claims) + tokenString, err := token.SignedString(privateKey) + if err != nil { + t.Fatalf("SignedString() error = %v", err) + } + return tokenString +} + +func newServiceTokenTestServer(t *testing.T, publicKeyPEM string) *Server { + t.Helper() + + server, err := NewServer(&Config{ + Provider: "hmac", + Audience: testPrimaryAudience, + JWTSecret: []byte("primary-secret"), + ServiceTokenEnabled: true, + ServiceTokenIssuer: testServiceIssuer, + ServiceTokenAudience: testServiceAudience, + ServiceTokenPublicKeyPEM: publicKeyPEM, + ServiceTokenSubjectPrefix: "svc-", + }) + if err != nil { + t.Fatalf("NewServer() error = %v", err) + } + return server +}