Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
192 changes: 119 additions & 73 deletions go/adk/pkg/sts/plugin.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,12 +2,17 @@ package sts

import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"strings"
"sync"
"time"

"github.com/a2aproject/a2a-go/a2asrv"
"github.com/go-logr/logr"
"github.com/golang-jwt/jwt/v5"
"github.com/kagent-dev/kagent/go/adk/pkg/constants"
"github.com/kagent-dev/kagent/go/adk/pkg/models"
"google.golang.org/adk/v2/agent"
adkplugin "google.golang.org/adk/v2/plugin"
Expand All @@ -33,8 +38,8 @@ func (e *TokenCacheEntry) HasExpired(bufferSeconds int64) bool {
// a header provider used by MCP tool transports.
type TokenPropagationPlugin struct {
integration *STSIntegration
tokenCache map[string]*TokenCacheEntry // keyed by session ID
actorTokenCache *TokenCacheEntry // used only for dynamic fetchActorToken providers
tokenCache map[cacheKey]*TokenCacheEntry
actorTokenCache *TokenCacheEntry // used only for dynamic fetchActorToken providers
mu sync.RWMutex
logger logr.Logger
bufferSeconds int64
Expand All @@ -45,18 +50,88 @@ type TokenPropagationPlugin struct {
func NewTokenPropagationPlugin(integration *STSIntegration, logger logr.Logger) *TokenPropagationPlugin {
return &TokenPropagationPlugin{
integration: integration,
tokenCache: make(map[string]*TokenCacheEntry),
tokenCache: make(map[cacheKey]*TokenCacheEntry),
logger: logger.WithName("sts-plugin"),
bufferSeconds: 5,
}
}

// getCachedToken retrieves a valid cached token for the session.
func (p *TokenPropagationPlugin) getCachedToken(sessionID string) (*TokenCacheEntry, bool) {
// parseUnverifiedClaims parses a JWT's claims WITHOUT signature or time
// validation. It is used only for cache partitioning and TTL, never for a
// security decision; tokens are validated server-side during STS exchange.
func parseUnverifiedClaims(token string) (jwt.MapClaims, bool) {
if token == "" {
return nil, false
}
claims := jwt.MapClaims{}
if _, _, err := jwt.NewParser(jwt.WithoutClaimsValidation()).ParseUnverified(token, claims); err != nil {
return nil, false
}
return claims, true
}

// subjectKey derives a stable per-principal cache discriminator from a bearer
// token: the issuer-scoped "sub" claim when present, otherwise a hash of the
// raw token so opaque or sub-less tokens still partition per principal. "sub"
// is only unique within an issuer, so it is combined with "iss" to avoid two
// principals from different issuers colliding onto one cache entry.
func subjectKey(token string) string {
if token == "" {
return ""
}
if claims, ok := parseUnverifiedClaims(token); ok {
if sub, _ := claims["sub"].(string); sub != "" {
iss, _ := claims["iss"].(string)
return iss + "\x00" + sub
}
}
sum := sha256.Sum256([]byte(token))
return "h:" + hex.EncodeToString(sum[:])
}

// actingBearer recovers the caller's raw bearer token for this request. It
// prefers the value executor.withBearerToken stored (models.BearerTokenKey) and
// falls back to the A2A CallContext Authorization header. The fallback keeps the
// per-subject cache key reliable at the MCP transport layer: the CallContext is
// the same source the round-tripper's propagateToken path reads, so it reaches
// the caller even when BearerTokenKey is not threaded to the MCP request context.
func actingBearer(ctx context.Context) string {
if token, ok := ctx.Value(models.BearerTokenKey).(string); ok && token != "" {
return token
}
callCtx, ok := a2asrv.CallContextFrom(ctx)
if !ok {
return ""
}
meta := callCtx.RequestMeta()
if meta == nil {
return ""
}
vals, ok := meta.Get(constants.AuthorizationHeader)
if !ok || len(vals) == 0 {
return ""
}
parts := strings.Fields(strings.TrimSpace(vals[0]))
if len(parts) >= 2 && strings.EqualFold(parts[0], "Bearer") {
return parts[1]
}
return ""
}

// cacheKey scopes a cache entry to both the session and the acting subject so a
// session that carries messages from multiple subjects keeps one exchanged
// token per subject rather than collapsing to whichever arrived first.
type cacheKey struct {
sessionID string
subject string
}

// getCachedToken retrieves a valid cached token for the session and subject.
func (p *TokenPropagationPlugin) getCachedToken(sessionID, subject string) (*TokenCacheEntry, bool) {
p.mu.RLock()
defer p.mu.RUnlock()

entry, ok := p.tokenCache[sessionID]
entry, ok := p.tokenCache[cacheKey{sessionID: sessionID, subject: subject}]
if !ok {
return nil, false
}
Expand All @@ -68,12 +143,12 @@ func (p *TokenPropagationPlugin) getCachedToken(sessionID string) (*TokenCacheEn
return entry, true
}

// setCachedToken caches a token for the session.
func (p *TokenPropagationPlugin) setCachedToken(sessionID string, token string, expiry int64) {
// setCachedToken caches a token for the session and subject.
func (p *TokenPropagationPlugin) setCachedToken(sessionID, subject, token string, expiry int64) {
p.mu.Lock()
defer p.mu.Unlock()

p.tokenCache[sessionID] = &TokenCacheEntry{
p.tokenCache[cacheKey{sessionID: sessionID, subject: subject}] = &TokenCacheEntry{
Token: token,
Expiry: expiry,
}
Expand Down Expand Up @@ -133,8 +208,20 @@ func (p *TokenPropagationPlugin) BeforeRunCallback(ctx agent.InvocationContext)
return nil, nil
}

// Check if we already have a valid cached token for this session.
if entry, ok := p.getCachedToken(sessionID); ok {
// Recover the acting bearer before the cache lookup: the cache is keyed by the
// acting subject, and a session shared by multiple subjects would otherwise
// reuse the first caller's token for every later caller.
bearerToken := actingBearer(ctx)

if bearerToken == "" {
p.logger.V(1).Info("No bearer token in context, skipping token propagation", "sessionID", sessionID)
return nil, nil
}

subject := subjectKey(bearerToken)

// Check if we already have a valid cached token for this session and subject.
if entry, ok := p.getCachedToken(sessionID, subject); ok {
p.logger.V(1).Info("Using cached STS token", "sessionID", sessionID)
if entry.Expiry > 0 {
p.logger.V(1).Info("Token expiry remaining",
Expand All @@ -143,19 +230,6 @@ func (p *TokenPropagationPlugin) BeforeRunCallback(ctx agent.InvocationContext)
return nil, nil
}

// Extract bearer token from context. executor.go stores it with models.BearerTokenKey.
bearerToken := ""
if v := ctx.Value(models.BearerTokenKey); v != nil {
if token, ok := v.(string); ok {
bearerToken = token
}
}

if bearerToken == "" {
p.logger.V(1).Info("No bearer token in context, skipping token propagation", "sessionID", sessionID)
return nil, nil
}

// Get subject token
subjectToken := bearerToken
if p.integration != nil {
Expand Down Expand Up @@ -198,12 +272,12 @@ func (p *TokenPropagationPlugin) BeforeRunCallback(ctx agent.InvocationContext)
// Fall back to JWT exp claim for cache TTL.
expiry = extractJWTExpiry(exchangedToken)
}
p.setCachedToken(sessionID, exchangedToken, expiry)
p.setCachedToken(sessionID, subject, exchangedToken, expiry)
p.logger.Info("Successfully exchanged and cached STS token", "sessionID", sessionID)
} else {
// No STS integration — cache the raw subject token for header injection.
expiry := extractJWTExpiry(subjectToken)
p.setCachedToken(sessionID, subjectToken, expiry)
p.setCachedToken(sessionID, subject, subjectToken, expiry)
p.logger.V(1).Info("Cached subject token (no STS exchange)", "sessionID", sessionID)
}

Expand All @@ -212,27 +286,16 @@ func (p *TokenPropagationPlugin) BeforeRunCallback(ctx agent.InvocationContext)

// AfterRunCallback is called after the ADK run finishes.
// It cleans up expired tokens from the cache.
func (p *TokenPropagationPlugin) AfterRunCallback(ctx agent.InvocationContext) {
sessionID := ""
if session := ctx.Session(); session != nil {
sessionID = session.ID()
}
if sessionID == "" {
return
}

func (p *TokenPropagationPlugin) AfterRunCallback(_ agent.InvocationContext) {
p.mu.Lock()
defer p.mu.Unlock()

// Remove expired subject token.
if entry, ok := p.tokenCache[sessionID]; ok {
for key, entry := range p.tokenCache {
if entry.HasExpired(p.bufferSeconds) {
p.logger.V(1).Info("Removing expired subject token from cache", "sessionID", sessionID)
delete(p.tokenCache, sessionID)
delete(p.tokenCache, key)
}
}
if p.actorTokenCache != nil && p.actorTokenCache.HasExpired(p.bufferSeconds) {
p.logger.V(1).Info("Removing expired actor token from cache")
p.actorTokenCache = nil
}
}
Expand All @@ -250,9 +313,14 @@ func (p *TokenPropagationPlugin) HeaderProvider(ctx context.Context) map[string]
return nil
}

entry, ok := p.getCachedToken(sessionID)
// Recover the acting subject from this request's own bearer, so the injected
// token matches the caller of this request rather than whichever subject
// first seeded the session.
subject := subjectKey(actingBearer(ctx))

entry, ok := p.getCachedToken(sessionID, subject)
if !ok {
p.logger.V(1).Info("No cached STS token for session, MCP request will use existing headers", "sessionID", sessionID)
p.logger.V(1).Info("No cached STS token for session/subject, MCP request will use existing headers", "sessionID", sessionID)
return nil
}

Expand All @@ -274,22 +342,12 @@ func sessionIDFromContext(ctx context.Context) string {
return sessionCtx.SessionID()
}

// GetTokenForSession retrieves the cached token for a specific session.
// Returns empty string if no valid token is cached.
func (p *TokenPropagationPlugin) GetTokenForSession(sessionID string) string {
entry, ok := p.getCachedToken(sessionID)
if !ok {
return ""
}
return entry.Token
}

// ClearCache clears all cached tokens.
func (p *TokenPropagationPlugin) ClearCache() {
p.mu.Lock()
defer p.mu.Unlock()

p.tokenCache = make(map[string]*TokenCacheEntry)
p.tokenCache = make(map[cacheKey]*TokenCacheEntry)
p.actorTokenCache = nil
p.logger.Info("Cleared STS token cache")
}
Expand All @@ -303,29 +361,17 @@ func (p *TokenPropagationPlugin) ADKPlugin() (*adkplugin.Plugin, error) {
})
}

// extractJWTExpiry extracts the 'exp' claim from a JWT token without verifying its signature.
// This is ONLY used for cache TTL management, not for security decisions.
// Token validation happens server-side during STS exchange.
// extractJWTExpiry extracts the 'exp' claim from a JWT token without verifying
// its signature. This is ONLY used for cache TTL management, not for security
// decisions. Token validation happens server-side during STS exchange.
func extractJWTExpiry(token string) int64 {
if token == "" {
claims, ok := parseUnverifiedClaims(token)
if !ok {
return 0
}

claims := jwt.MapClaims{}
if _, _, err := jwt.NewParser(jwt.WithoutClaimsValidation()).ParseUnverified(token, claims); err != nil {
exp, err := claims.GetExpirationTime()
if err != nil || exp == nil {
return 0
}

if exp, ok := claims["exp"]; ok {
switch v := exp.(type) {
case float64:
return int64(v)
case int64:
return v
case int:
return int64(v)
}
}

return 0
return exp.Unix()
}
Loading
Loading