From 6fa2356c0cb19b2b89dadc4a45564527e131e45b Mon Sep 17 00:00:00 2001 From: yekkhan-liftoff Date: Fri, 19 Jun 2026 20:23:24 +0800 Subject: [PATCH 1/4] feat(auth): add service token validation --- README.md | 9 ++ config.go | 91 +++++++++++-- docs/CONFIGURATION.md | 35 +++++ docs/SERVICE-TOKENS.md | 74 ++++++++++ provider/provider.go | 181 +++++++++++++++++++++++- service_token_test.go | 302 +++++++++++++++++++++++++++++++++++++++++ 6 files changed, 679 insertions(+), 13 deletions(-) create mode 100644 docs/SERVICE-TOKENS.md create mode 100644 service_token_test.go 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 ef67b30..fdf8fb1 100644 --- a/config.go +++ b/config.go @@ -1,6 +1,7 @@ package oauth import ( + "encoding/base64" "fmt" "github.com/Vungle/oauth-mcp-proxy/provider" @@ -25,6 +26,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.). @@ -73,6 +81,21 @@ 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") + } + } + // Validate proxy mode requirements if c.Mode == "proxy" { if c.ClientID == "" { @@ -119,11 +142,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 @@ -140,6 +167,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 } @@ -217,6 +252,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 @@ -280,8 +325,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", "")). @@ -290,6 +338,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 969120d..5724c62 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: "phoebe-service", + ServiceTokenAudience: "api://phoebe-mcp", + 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..1f14ec6 --- /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: "phoebe-service", + ServiceTokenAudience: "api://phoebe-mcp", + 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=phoebe-service +SERVICE_TOKEN_AUDIENCE=api://phoebe-mcp +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/provider/provider.go b/provider/provider.go index 39c5847..15e4f55 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,169 @@ 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(), + ) + 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..c454825 --- /dev/null +++ b/service_token_test.go @@ -0,0 +1,302 @@ +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 = "phoebe-service" + testServiceAudience = "api://phoebe-mcp" +) + +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-calypso", + "preferred_username": "svc-calypso", + "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-calypso" { + t.Fatalf("Subject = %q, want svc-calypso", user.Subject) + } + if user.Username != "svc-calypso" { + t.Fatalf("Username = %q, want svc-calypso", 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-calypso", + "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-calypso" { + t.Fatalf("Subject = %q, want svc-calypso", user.Subject) + } +} + +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-calypso", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: wrongPrivateKey, + }, + { + name: "wrong issuer", + claims: jwt.MapClaims{ + "iss": "other-service", + "aud": testServiceAudience, + "sub": "svc-calypso", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "wrong audience", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": "api://other", + "sub": "svc-calypso", + "exp": time.Now().Add(time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "expired", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-calypso", + "exp": time.Now().Add(-time.Hour).Unix(), + }, + key: privateKey, + }, + { + name: "missing expiry", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "svc-calypso", + }, + key: privateKey, + }, + { + name: "invalid subject prefix", + claims: jwt.MapClaims{ + "iss": testServiceIssuer, + "aud": testServiceAudience, + "sub": "calypso", + "exp": time.Now().Add(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-calypso", + "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 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 +} From 6edcc70314e18d8271375b135e3a267dce783fdc Mon Sep 17 00:00:00 2001 From: yekkhan-liftoff Date: Fri, 19 Jun 2026 20:53:27 +0800 Subject: [PATCH 2/4] docs(auth): use generic service token examples --- docs/CONFIGURATION.md | 4 ++-- docs/SERVICE-TOKENS.md | 8 ++++---- service_token_test.go | 36 ++++++++++++++++++------------------ 3 files changed, 24 insertions(+), 24 deletions(-) diff --git a/docs/CONFIGURATION.md b/docs/CONFIGURATION.md index 5724c62..8bacaaf 100644 --- a/docs/CONFIGURATION.md +++ b/docs/CONFIGURATION.md @@ -142,8 +142,8 @@ cfg := &oauth.Config{ Audience: "https://company.okta.com", ServiceTokenEnabled: true, - ServiceTokenIssuer: "phoebe-service", - ServiceTokenAudience: "api://phoebe-mcp", + ServiceTokenIssuer: "agent-auth-service", + ServiceTokenAudience: "api://example-mcp-server", ServiceTokenPublicKeyPEM: publicKeyPEM, ServiceTokenSubjectPrefix: "svc-", } diff --git a/docs/SERVICE-TOKENS.md b/docs/SERVICE-TOKENS.md index 1f14ec6..45fd802 100644 --- a/docs/SERVICE-TOKENS.md +++ b/docs/SERVICE-TOKENS.md @@ -29,8 +29,8 @@ oauthServer, oauthOption, err := mark3labs.WithOAuth(mux, &oauth.Config{ Audience: "https://your-company.okta.com", ServiceTokenEnabled: true, - ServiceTokenIssuer: "phoebe-service", - ServiceTokenAudience: "api://phoebe-mcp", + ServiceTokenIssuer: "agent-auth-service", + ServiceTokenAudience: "api://example-mcp-server", ServiceTokenPublicKeyPEM: publicKeyPEM, ServiceTokenSubjectPrefix: "svc-", }) @@ -44,8 +44,8 @@ oauthServer, oauthOption, err := mark3labs.WithOAuth(mux, &oauth.Config{ ```text SERVICE_TOKEN_ENABLED=true -SERVICE_TOKEN_ISSUER=phoebe-service -SERVICE_TOKEN_AUDIENCE=api://phoebe-mcp +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- diff --git a/service_token_test.go b/service_token_test.go index c454825..c6d92d8 100644 --- a/service_token_test.go +++ b/service_token_test.go @@ -17,8 +17,8 @@ import ( const ( testPrimaryAudience = "api://primary" - testServiceIssuer = "phoebe-service" - testServiceAudience = "api://phoebe-mcp" + testServiceIssuer = "agent-auth-service" + testServiceAudience = "api://example-mcp-server" ) func TestServiceTokenValidation(t *testing.T) { @@ -28,8 +28,8 @@ func TestServiceTokenValidation(t *testing.T) { token := signServiceToken(t, privateKey, jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", - "preferred_username": "svc-calypso", + "sub": "svc-example-agent", + "preferred_username": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }) @@ -37,11 +37,11 @@ func TestServiceTokenValidation(t *testing.T) { if err != nil { t.Fatalf("ValidateTokenCached() error = %v", err) } - if user.Subject != "svc-calypso" { - t.Fatalf("Subject = %q, want svc-calypso", user.Subject) + if user.Subject != "svc-example-agent" { + t.Fatalf("Subject = %q, want svc-example-agent", user.Subject) } - if user.Username != "svc-calypso" { - t.Fatalf("Username = %q, want svc-calypso", user.Username) + if user.Username != "svc-example-agent" { + t.Fatalf("Username = %q, want svc-example-agent", user.Username) } } @@ -52,7 +52,7 @@ func TestServiceTokenValidationRS256(t *testing.T) { token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }) tokenString, err := token.SignedString(privateKey) @@ -64,8 +64,8 @@ func TestServiceTokenValidationRS256(t *testing.T) { if err != nil { t.Fatalf("ValidateTokenCached() error = %v", err) } - if user.Subject != "svc-calypso" { - t.Fatalf("Subject = %q, want svc-calypso", user.Subject) + if user.Subject != "svc-example-agent" { + t.Fatalf("Subject = %q, want svc-example-agent", user.Subject) } } @@ -84,7 +84,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }, key: wrongPrivateKey, @@ -94,7 +94,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": "other-service", "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }, key: privateKey, @@ -104,7 +104,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": testServiceIssuer, "aud": "api://other", - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }, key: privateKey, @@ -114,7 +114,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(-time.Hour).Unix(), }, key: privateKey, @@ -124,7 +124,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", }, key: privateKey, }, @@ -133,7 +133,7 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { claims: jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "calypso", + "sub": "example-agent", "exp": time.Now().Add(time.Hour).Unix(), }, key: privateKey, @@ -157,7 +157,7 @@ func TestServiceTokenRejectsSymmetricAlgorithm(t *testing.T) { token := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{ "iss": testServiceIssuer, "aud": testServiceAudience, - "sub": "svc-calypso", + "sub": "svc-example-agent", "exp": time.Now().Add(time.Hour).Unix(), }) tokenString, err := token.SignedString([]byte("shared-secret")) From f318d793f1a2454d04f561c2e8af42af7668182e Mon Sep 17 00:00:00 2001 From: yekkhan-liftoff Date: Mon, 22 Jun 2026 18:04:47 +0800 Subject: [PATCH 3/4] fix(auth): cap token cache by jwt expiry --- config.go | 3 +++ middleware.go | 9 ++++--- oauth.go | 34 +++++++++++++++++++++++--- provider/provider.go | 1 + service_token_test.go | 57 +++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 97 insertions(+), 7 deletions(-) diff --git a/config.go b/config.go index 086d45d..cb205ec 100644 --- a/config.go +++ b/config.go @@ -97,6 +97,9 @@ func (c *Config) Validate() error { 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 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..3eaf67e 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.Time.After(now) { + return time.Time{}, false + } + if claims.ExpiresAt.Time.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 15e4f55..40b9504 100644 --- a/provider/provider.go +++ b/provider/provider.go @@ -159,6 +159,7 @@ func (v *ServiceTokenValidator) ValidateToken(ctx context.Context, tokenString s 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) diff --git a/service_token_test.go b/service_token_test.go index c6d92d8..82d4480 100644 --- a/service_token_test.go +++ b/service_token_test.go @@ -69,6 +69,31 @@ func TestServiceTokenValidationRS256(t *testing.T) { } } +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) @@ -138,6 +163,17 @@ func TestServiceTokenRejectsInvalidTokens(t *testing.T) { }, 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 { @@ -214,6 +250,27 @@ func TestServiceTokenConfigValidation(t *testing.T) { } } +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) From 71732545cc1786a7b5c459b516cfd4c37b93cb92 Mon Sep 17 00:00:00 2001 From: yekkhan-liftoff Date: Mon, 22 Jun 2026 18:25:03 +0800 Subject: [PATCH 4/4] fix(auth): satisfy staticcheck for token expiry --- oauth.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/oauth.go b/oauth.go index 3eaf67e..9dbc4c3 100644 --- a/oauth.go +++ b/oauth.go @@ -151,10 +151,10 @@ func cacheExpiresAtForToken(tokenString string, now time.Time) (time.Time, bool) if claims.ExpiresAt == nil { return expiresAt, true } - if !claims.ExpiresAt.Time.After(now) { + if !claims.ExpiresAt.After(now) { return time.Time{}, false } - if claims.ExpiresAt.Time.Before(expiresAt) { + if claims.ExpiresAt.Before(expiresAt) { return claims.ExpiresAt.Time, true }