diff --git a/internal/atproto/auth/dpop.go b/internal/atproto/auth/dpop.go new file mode 100644 index 0000000..e8263fa --- /dev/null +++ b/internal/atproto/auth/dpop.go @@ -0,0 +1,484 @@ +package auth + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "fmt" + "math/big" + "strings" + "sync" + "time" + + indigoCrypto "github.com/bluesky-social/indigo/atproto/atcrypto" + "github.com/golang-jwt/jwt/v5" +) + +// NonceCache provides replay protection for DPoP proofs by tracking seen jti values. +// This prevents an attacker from reusing a captured DPoP proof within the validity window. +// Per RFC 9449 Section 11.1, servers SHOULD prevent replay attacks. +type NonceCache struct { + seen map[string]time.Time // jti -> expiration time + stopCh chan struct{} + maxAge time.Duration // How long to keep entries + cleanup time.Duration // How often to clean up expired entries + mu sync.RWMutex +} + +// NewNonceCache creates a new nonce cache for DPoP replay protection. +// maxAge should match or exceed DPoPVerifier.MaxProofAge. +func NewNonceCache(maxAge time.Duration) *NonceCache { + nc := &NonceCache{ + seen: make(map[string]time.Time), + maxAge: maxAge, + cleanup: maxAge / 2, // Clean up at half the max age + stopCh: make(chan struct{}), + } + + // Start background cleanup goroutine + go nc.cleanupLoop() + + return nc +} + +// CheckAndStore checks if a jti has been seen before and stores it if not. +// Returns true if the jti is fresh (not a replay), false if it's a replay. +func (nc *NonceCache) CheckAndStore(jti string) bool { + nc.mu.Lock() + defer nc.mu.Unlock() + + now := time.Now() + expiry := now.Add(nc.maxAge) + + // Check if already seen + if existingExpiry, seen := nc.seen[jti]; seen { + // Still valid (not expired) - this is a replay + if existingExpiry.After(now) { + return false + } + // Expired entry - allow reuse and update expiry + } + + // Store the new jti + nc.seen[jti] = expiry + return true +} + +// cleanupLoop periodically removes expired entries from the cache +func (nc *NonceCache) cleanupLoop() { + ticker := time.NewTicker(nc.cleanup) + defer ticker.Stop() + + for { + select { + case <-ticker.C: + nc.cleanupExpired() + case <-nc.stopCh: + return + } + } +} + +// cleanupExpired removes expired entries from the cache +func (nc *NonceCache) cleanupExpired() { + nc.mu.Lock() + defer nc.mu.Unlock() + + now := time.Now() + for jti, expiry := range nc.seen { + if expiry.Before(now) { + delete(nc.seen, jti) + } + } +} + +// Stop stops the cleanup goroutine. Call this when done with the cache. +func (nc *NonceCache) Stop() { + close(nc.stopCh) +} + +// Size returns the number of entries in the cache (for testing/monitoring) +func (nc *NonceCache) Size() int { + nc.mu.RLock() + defer nc.mu.RUnlock() + return len(nc.seen) +} + +// DPoPClaims represents the claims in a DPoP proof JWT (RFC 9449) +type DPoPClaims struct { + jwt.RegisteredClaims + + // HTTP method of the request (e.g., "GET", "POST") + HTTPMethod string `json:"htm"` + + // HTTP URI of the request (without query and fragment parts) + HTTPURI string `json:"htu"` + + // Access token hash (optional, for token binding) + AccessTokenHash string `json:"ath,omitempty"` +} + +// DPoPProof represents a parsed and verified DPoP proof +type DPoPProof struct { + RawPublicJWK map[string]interface{} + Claims *DPoPClaims + PublicKey interface{} // *ecdsa.PublicKey or similar + Thumbprint string // JWK thumbprint (base64url) +} + +// DPoPVerifier verifies DPoP proofs for OAuth token binding +type DPoPVerifier struct { + // Optional: custom nonce validation function (for server-issued nonces) + ValidateNonce func(nonce string) bool + + // NonceCache for replay protection (optional but recommended) + // If nil, jti replay protection is disabled + NonceCache *NonceCache + + // Maximum allowed clock skew for timestamp validation + MaxClockSkew time.Duration + + // Maximum age of DPoP proof (prevents replay with old proofs) + MaxProofAge time.Duration +} + +// NewDPoPVerifier creates a DPoP verifier with sensible defaults including replay protection +func NewDPoPVerifier() *DPoPVerifier { + maxProofAge := 5 * time.Minute + return &DPoPVerifier{ + MaxClockSkew: 30 * time.Second, + MaxProofAge: maxProofAge, + NonceCache: NewNonceCache(maxProofAge), + } +} + +// NewDPoPVerifierWithoutReplayProtection creates a DPoP verifier without replay protection. +// This should only be used in testing or when replay protection is handled externally. +func NewDPoPVerifierWithoutReplayProtection() *DPoPVerifier { + return &DPoPVerifier{ + MaxClockSkew: 30 * time.Second, + MaxProofAge: 5 * time.Minute, + NonceCache: nil, // No replay protection + } +} + +// Stop stops background goroutines. Call this when shutting down. +func (v *DPoPVerifier) Stop() { + if v.NonceCache != nil { + v.NonceCache.Stop() + } +} + +// VerifyDPoPProof verifies a DPoP proof JWT and returns the parsed proof +func (v *DPoPVerifier) VerifyDPoPProof(dpopProof, httpMethod, httpURI string) (*DPoPProof, error) { + // Parse the DPoP JWT without verification first to extract the header + parser := jwt.NewParser(jwt.WithoutClaimsValidation()) + token, _, err := parser.ParseUnverified(dpopProof, &DPoPClaims{}) + if err != nil { + return nil, fmt.Errorf("failed to parse DPoP proof: %w", err) + } + + // Extract and validate the header + header, ok := token.Header["typ"].(string) + if !ok || header != "dpop+jwt" { + return nil, fmt.Errorf("invalid DPoP proof: typ must be 'dpop+jwt', got '%s'", header) + } + + alg, ok := token.Header["alg"].(string) + if !ok { + return nil, fmt.Errorf("invalid DPoP proof: missing alg header") + } + + // Extract the JWK from the header + jwkRaw, ok := token.Header["jwk"] + if !ok { + return nil, fmt.Errorf("invalid DPoP proof: missing jwk header") + } + + jwkMap, ok := jwkRaw.(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("invalid DPoP proof: jwk must be an object") + } + + // Parse the public key from JWK + publicKey, err := parseJWKToPublicKey(jwkMap) + if err != nil { + return nil, fmt.Errorf("invalid DPoP proof JWK: %w", err) + } + + // Calculate the JWK thumbprint + thumbprint, err := CalculateJWKThumbprint(jwkMap) + if err != nil { + return nil, fmt.Errorf("failed to calculate JWK thumbprint: %w", err) + } + + // Now verify the signature + verifiedToken, err := jwt.ParseWithClaims(dpopProof, &DPoPClaims{}, func(token *jwt.Token) (interface{}, error) { + // Verify the signing method matches what we expect + switch alg { + case "ES256": + if _, ok := token.Method.(*jwt.SigningMethodECDSA); !ok { + return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + } + case "ES384", "ES512": + if _, ok := token.Method.(*jwt.SigningMethodECDSA); !ok { + return nil, fmt.Errorf("unexpected signing method: %v", token.Header["alg"]) + } + case "RS256", "RS384", "RS512", "PS256", "PS384", "PS512": + // RSA methods - we primarily support ES256 for atproto + return nil, fmt.Errorf("RSA algorithms not yet supported for DPoP: %s", alg) + default: + return nil, fmt.Errorf("unsupported DPoP algorithm: %s", alg) + } + return publicKey, nil + }) + if err != nil { + return nil, fmt.Errorf("DPoP proof signature verification failed: %w", err) + } + + claims, ok := verifiedToken.Claims.(*DPoPClaims) + if !ok { + return nil, fmt.Errorf("invalid DPoP claims type") + } + + // Validate the claims + if err := v.validateDPoPClaims(claims, httpMethod, httpURI); err != nil { + return nil, err + } + + return &DPoPProof{ + Claims: claims, + PublicKey: publicKey, + Thumbprint: thumbprint, + RawPublicJWK: jwkMap, + }, nil +} + +// validateDPoPClaims validates the DPoP proof claims +func (v *DPoPVerifier) validateDPoPClaims(claims *DPoPClaims, expectedMethod, expectedURI string) error { + // Validate jti (unique identifier) is present + if claims.ID == "" { + return fmt.Errorf("DPoP proof missing jti claim") + } + + // Validate htm (HTTP method) + if !strings.EqualFold(claims.HTTPMethod, expectedMethod) { + return fmt.Errorf("DPoP proof htm mismatch: expected %s, got %s", expectedMethod, claims.HTTPMethod) + } + + // Validate htu (HTTP URI) - compare without query/fragment + expectedURIBase := stripQueryFragment(expectedURI) + claimURIBase := stripQueryFragment(claims.HTTPURI) + if expectedURIBase != claimURIBase { + return fmt.Errorf("DPoP proof htu mismatch: expected %s, got %s", expectedURIBase, claimURIBase) + } + + // Validate iat (issued at) is present and recent + if claims.IssuedAt == nil { + return fmt.Errorf("DPoP proof missing iat claim") + } + + now := time.Now() + iat := claims.IssuedAt.Time + + // Check clock skew (not too far in the future) + if iat.After(now.Add(v.MaxClockSkew)) { + return fmt.Errorf("DPoP proof iat is in the future") + } + + // Check proof age (not too old) + if now.Sub(iat) > v.MaxProofAge { + return fmt.Errorf("DPoP proof is too old (issued %v ago, max %v)", now.Sub(iat), v.MaxProofAge) + } + + // SECURITY: Check for replay attack using jti + // Per RFC 9449 Section 11.1, servers SHOULD prevent replay attacks + if v.NonceCache != nil { + if !v.NonceCache.CheckAndStore(claims.ID) { + return fmt.Errorf("DPoP proof replay detected: jti %s already used", claims.ID) + } + } + + return nil +} + +// VerifyTokenBinding verifies that the DPoP proof binds to the access token +// by comparing the proof's thumbprint to the token's cnf.jkt claim +func (v *DPoPVerifier) VerifyTokenBinding(proof *DPoPProof, expectedThumbprint string) error { + if proof.Thumbprint != expectedThumbprint { + return fmt.Errorf("DPoP proof thumbprint mismatch: token expects %s, proof has %s", + expectedThumbprint, proof.Thumbprint) + } + return nil +} + +// CalculateJWKThumbprint calculates the JWK thumbprint per RFC 7638 +// The thumbprint is the base64url-encoded SHA-256 hash of the canonical JWK representation +func CalculateJWKThumbprint(jwk map[string]interface{}) (string, error) { + kty, ok := jwk["kty"].(string) + if !ok { + return "", fmt.Errorf("JWK missing kty") + } + + // Build the canonical JWK representation based on key type + // Per RFC 7638, only specific members are included, in lexicographic order + var canonical map[string]string + + switch kty { + case "EC": + crv, ok := jwk["crv"].(string) + if !ok { + return "", fmt.Errorf("EC JWK missing crv") + } + x, ok := jwk["x"].(string) + if !ok { + return "", fmt.Errorf("EC JWK missing x") + } + y, ok := jwk["y"].(string) + if !ok { + return "", fmt.Errorf("EC JWK missing y") + } + // Lexicographic order: crv, kty, x, y + canonical = map[string]string{ + "crv": crv, + "kty": kty, + "x": x, + "y": y, + } + case "RSA": + e, ok := jwk["e"].(string) + if !ok { + return "", fmt.Errorf("RSA JWK missing e") + } + n, ok := jwk["n"].(string) + if !ok { + return "", fmt.Errorf("RSA JWK missing n") + } + // Lexicographic order: e, kty, n + canonical = map[string]string{ + "e": e, + "kty": kty, + "n": n, + } + case "OKP": + crv, ok := jwk["crv"].(string) + if !ok { + return "", fmt.Errorf("OKP JWK missing crv") + } + x, ok := jwk["x"].(string) + if !ok { + return "", fmt.Errorf("OKP JWK missing x") + } + // Lexicographic order: crv, kty, x + canonical = map[string]string{ + "crv": crv, + "kty": kty, + "x": x, + } + default: + return "", fmt.Errorf("unsupported JWK key type: %s", kty) + } + + // Serialize to JSON (Go's json.Marshal produces lexicographically ordered keys for map[string]string) + canonicalJSON, err := json.Marshal(canonical) + if err != nil { + return "", fmt.Errorf("failed to serialize canonical JWK: %w", err) + } + + // SHA-256 hash + hash := sha256.Sum256(canonicalJSON) + + // Base64url encode (no padding) + thumbprint := base64.RawURLEncoding.EncodeToString(hash[:]) + + return thumbprint, nil +} + +// parseJWKToPublicKey parses a JWK map to a Go public key +func parseJWKToPublicKey(jwkMap map[string]interface{}) (interface{}, error) { + // Convert map to JSON bytes for indigo's parser + jwkBytes, err := json.Marshal(jwkMap) + if err != nil { + return nil, fmt.Errorf("failed to serialize JWK: %w", err) + } + + // Try to parse with indigo's crypto package + pubKey, err := indigoCrypto.ParsePublicJWKBytes(jwkBytes) + if err != nil { + return nil, fmt.Errorf("failed to parse JWK: %w", err) + } + + // Convert indigo's PublicKey to Go's ecdsa.PublicKey + jwk, err := pubKey.JWK() + if err != nil { + return nil, fmt.Errorf("failed to get JWK from public key: %w", err) + } + + // Use our existing conversion function + return atcryptoJWKToECDSAFromIndigoJWK(jwk) +} + +// atcryptoJWKToECDSAFromIndigoJWK converts an indigo JWK to Go ecdsa.PublicKey +func atcryptoJWKToECDSAFromIndigoJWK(jwk *indigoCrypto.JWK) (*ecdsa.PublicKey, error) { + if jwk.KeyType != "EC" { + return nil, fmt.Errorf("unsupported JWK key type: %s (expected EC)", jwk.KeyType) + } + + xBytes, err := base64.RawURLEncoding.DecodeString(jwk.X) + if err != nil { + return nil, fmt.Errorf("invalid JWK X coordinate: %w", err) + } + yBytes, err := base64.RawURLEncoding.DecodeString(jwk.Y) + if err != nil { + return nil, fmt.Errorf("invalid JWK Y coordinate: %w", err) + } + + var curve ecdsa.PublicKey + switch jwk.Curve { + case "P-256": + curve.Curve = ecdsaP256Curve() + case "P-384": + curve.Curve = ecdsaP384Curve() + case "P-521": + curve.Curve = ecdsaP521Curve() + default: + return nil, fmt.Errorf("unsupported curve: %s", jwk.Curve) + } + + curve.X = new(big.Int).SetBytes(xBytes) + curve.Y = new(big.Int).SetBytes(yBytes) + + return &curve, nil +} + +// Helper functions for elliptic curves +func ecdsaP256Curve() elliptic.Curve { return elliptic.P256() } +func ecdsaP384Curve() elliptic.Curve { return elliptic.P384() } +func ecdsaP521Curve() elliptic.Curve { return elliptic.P521() } + +// stripQueryFragment removes query and fragment from a URI +func stripQueryFragment(uri string) string { + if idx := strings.Index(uri, "?"); idx != -1 { + uri = uri[:idx] + } + if idx := strings.Index(uri, "#"); idx != -1 { + uri = uri[:idx] + } + return uri +} + +// ExtractCnfJkt extracts the cnf.jkt (confirmation key thumbprint) from JWT claims +func ExtractCnfJkt(claims *Claims) (string, error) { + if claims.Confirmation == nil { + return "", fmt.Errorf("token missing cnf claim (no DPoP binding)") + } + + jkt, ok := claims.Confirmation["jkt"].(string) + if !ok || jkt == "" { + return "", fmt.Errorf("token cnf claim missing jkt (DPoP key thumbprint)") + } + + return jkt, nil +} diff --git a/internal/atproto/auth/dpop_test.go b/internal/atproto/auth/dpop_test.go new file mode 100644 index 0000000..b647923 --- /dev/null +++ b/internal/atproto/auth/dpop_test.go @@ -0,0 +1,921 @@ +package auth + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "strings" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" + "github.com/google/uuid" +) + +// === Test Helpers === + +// testECKey holds a test ES256 key pair +type testECKey struct { + privateKey *ecdsa.PrivateKey + publicKey *ecdsa.PublicKey + jwk map[string]interface{} + thumbprint string +} + +// generateTestES256Key generates a test ES256 key pair and JWK +func generateTestES256Key(t *testing.T) *testECKey { + t.Helper() + + privateKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + t.Fatalf("Failed to generate test key: %v", err) + } + + // Encode public key coordinates as base64url + xBytes := privateKey.PublicKey.X.Bytes() + yBytes := privateKey.PublicKey.Y.Bytes() + + // P-256 coordinates must be 32 bytes (pad if needed) + xBytes = padTo32Bytes(xBytes) + yBytes = padTo32Bytes(yBytes) + + x := base64.RawURLEncoding.EncodeToString(xBytes) + y := base64.RawURLEncoding.EncodeToString(yBytes) + + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": x, + "y": y, + } + + // Calculate thumbprint + thumbprint, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("Failed to calculate thumbprint: %v", err) + } + + return &testECKey{ + privateKey: privateKey, + publicKey: &privateKey.PublicKey, + jwk: jwk, + thumbprint: thumbprint, + } +} + +// padTo32Bytes pads a byte slice to 32 bytes (required for P-256 coordinates) +func padTo32Bytes(b []byte) []byte { + if len(b) >= 32 { + return b + } + padded := make([]byte, 32) + copy(padded[32-len(b):], b) + return padded +} + +// createDPoPProof creates a DPoP proof JWT for testing +func createDPoPProof(t *testing.T, key *testECKey, method, uri string, iat time.Time, jti string) string { + t.Helper() + + claims := &DPoPClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + ID: jti, + IssuedAt: jwt.NewNumericDate(iat), + }, + HTTPMethod: method, + HTTPURI: uri, + } + + token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) + token.Header["typ"] = "dpop+jwt" + token.Header["jwk"] = key.jwk + + tokenString, err := token.SignedString(key.privateKey) + if err != nil { + t.Fatalf("Failed to create DPoP proof: %v", err) + } + + return tokenString +} + +// === JWK Thumbprint Tests (RFC 7638) === + +func TestCalculateJWKThumbprint_EC_P256(t *testing.T) { + // Test with known values from RFC 7638 Appendix A (adapted for P-256) + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "WKn-ZIGevcwGIyyrzFoZNBdaq9_TsqzGl96oc0CWuis", + "y": "y77t-RvAHRKTsSGdIYUfweuOvwrvDD-Q3Hv5J0fSKbE", + } + + thumbprint, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("CalculateJWKThumbprint failed: %v", err) + } + + if thumbprint == "" { + t.Error("Expected non-empty thumbprint") + } + + // Verify it's valid base64url + _, err = base64.RawURLEncoding.DecodeString(thumbprint) + if err != nil { + t.Errorf("Thumbprint is not valid base64url: %v", err) + } + + // Verify length (SHA-256 produces 32 bytes = 43 base64url chars) + if len(thumbprint) != 43 { + t.Errorf("Expected thumbprint length 43, got %d", len(thumbprint)) + } +} + +func TestCalculateJWKThumbprint_Deterministic(t *testing.T) { + // Same key should produce same thumbprint + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "test-x-coordinate", + "y": "test-y-coordinate", + } + + thumbprint1, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("First CalculateJWKThumbprint failed: %v", err) + } + + thumbprint2, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("Second CalculateJWKThumbprint failed: %v", err) + } + + if thumbprint1 != thumbprint2 { + t.Errorf("Thumbprints are not deterministic: %s != %s", thumbprint1, thumbprint2) + } +} + +func TestCalculateJWKThumbprint_DifferentKeys(t *testing.T) { + // Different keys should produce different thumbprints + jwk1 := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "coordinate-x-1", + "y": "coordinate-y-1", + } + + jwk2 := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "coordinate-x-2", + "y": "coordinate-y-2", + } + + thumbprint1, err := CalculateJWKThumbprint(jwk1) + if err != nil { + t.Fatalf("First CalculateJWKThumbprint failed: %v", err) + } + + thumbprint2, err := CalculateJWKThumbprint(jwk2) + if err != nil { + t.Fatalf("Second CalculateJWKThumbprint failed: %v", err) + } + + if thumbprint1 == thumbprint2 { + t.Error("Different keys produced same thumbprint (collision)") + } +} + +func TestCalculateJWKThumbprint_MissingKty(t *testing.T) { + jwk := map[string]interface{}{ + "crv": "P-256", + "x": "test-x", + "y": "test-y", + } + + _, err := CalculateJWKThumbprint(jwk) + if err == nil { + t.Error("Expected error for missing kty, got nil") + } + if err != nil && !contains(err.Error(), "missing kty") { + t.Errorf("Expected error about missing kty, got: %v", err) + } +} + +func TestCalculateJWKThumbprint_EC_MissingCrv(t *testing.T) { + jwk := map[string]interface{}{ + "kty": "EC", + "x": "test-x", + "y": "test-y", + } + + _, err := CalculateJWKThumbprint(jwk) + if err == nil { + t.Error("Expected error for missing crv, got nil") + } + if err != nil && !contains(err.Error(), "missing crv") { + t.Errorf("Expected error about missing crv, got: %v", err) + } +} + +func TestCalculateJWKThumbprint_EC_MissingX(t *testing.T) { + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "y": "test-y", + } + + _, err := CalculateJWKThumbprint(jwk) + if err == nil { + t.Error("Expected error for missing x, got nil") + } + if err != nil && !contains(err.Error(), "missing x") { + t.Errorf("Expected error about missing x, got: %v", err) + } +} + +func TestCalculateJWKThumbprint_EC_MissingY(t *testing.T) { + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "test-x", + } + + _, err := CalculateJWKThumbprint(jwk) + if err == nil { + t.Error("Expected error for missing y, got nil") + } + if err != nil && !contains(err.Error(), "missing y") { + t.Errorf("Expected error about missing y, got: %v", err) + } +} + +func TestCalculateJWKThumbprint_RSA(t *testing.T) { + // Test RSA key thumbprint calculation + jwk := map[string]interface{}{ + "kty": "RSA", + "e": "AQAB", + "n": "test-modulus", + } + + thumbprint, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("CalculateJWKThumbprint failed for RSA: %v", err) + } + + if thumbprint == "" { + t.Error("Expected non-empty thumbprint for RSA key") + } +} + +func TestCalculateJWKThumbprint_OKP(t *testing.T) { + // Test OKP (Octet Key Pair) thumbprint calculation + jwk := map[string]interface{}{ + "kty": "OKP", + "crv": "Ed25519", + "x": "test-x-coordinate", + } + + thumbprint, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("CalculateJWKThumbprint failed for OKP: %v", err) + } + + if thumbprint == "" { + t.Error("Expected non-empty thumbprint for OKP key") + } +} + +func TestCalculateJWKThumbprint_UnsupportedKeyType(t *testing.T) { + jwk := map[string]interface{}{ + "kty": "UNKNOWN", + } + + _, err := CalculateJWKThumbprint(jwk) + if err == nil { + t.Error("Expected error for unsupported key type, got nil") + } + if err != nil && !contains(err.Error(), "unsupported JWK key type") { + t.Errorf("Expected error about unsupported key type, got: %v", err) + } +} + +func TestCalculateJWKThumbprint_CanonicalJSON(t *testing.T) { + // RFC 7638 requires lexicographic ordering of keys in canonical JSON + // This test verifies that the canonical JSON is correctly ordered + + jwk := map[string]interface{}{ + "kty": "EC", + "crv": "P-256", + "x": "x-coord", + "y": "y-coord", + } + + // The canonical JSON should be: {"crv":"P-256","kty":"EC","x":"x-coord","y":"y-coord"} + // (lexicographically ordered: crv, kty, x, y) + + canonical := map[string]string{ + "crv": "P-256", + "kty": "EC", + "x": "x-coord", + "y": "y-coord", + } + + canonicalJSON, err := json.Marshal(canonical) + if err != nil { + t.Fatalf("Failed to marshal canonical JSON: %v", err) + } + + expectedHash := sha256.Sum256(canonicalJSON) + expectedThumbprint := base64.RawURLEncoding.EncodeToString(expectedHash[:]) + + actualThumbprint, err := CalculateJWKThumbprint(jwk) + if err != nil { + t.Fatalf("CalculateJWKThumbprint failed: %v", err) + } + + if actualThumbprint != expectedThumbprint { + t.Errorf("Thumbprint doesn't match expected canonical JSON hash\nExpected: %s\nGot: %s", + expectedThumbprint, actualThumbprint) + } +} + +// === DPoP Proof Verification Tests === + +func TestVerifyDPoPProof_Valid(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + result, err := verifier.VerifyDPoPProof(proof, method, uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed for valid proof: %v", err) + } + + if result == nil { + t.Fatal("Expected non-nil proof result") + } + + if result.Claims.HTTPMethod != method { + t.Errorf("Expected method %s, got %s", method, result.Claims.HTTPMethod) + } + + if result.Claims.HTTPURI != uri { + t.Errorf("Expected URI %s, got %s", uri, result.Claims.HTTPURI) + } + + if result.Claims.ID != jti { + t.Errorf("Expected jti %s, got %s", jti, result.Claims.ID) + } + + if result.Thumbprint != key.thumbprint { + t.Errorf("Expected thumbprint %s, got %s", key.thumbprint, result.Thumbprint) + } +} + +func TestVerifyDPoPProof_InvalidSignature(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + wrongKey := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + // Create proof with one key + proof := createDPoPProof(t, key, method, uri, iat, jti) + + // Parse and modify to use wrong key's JWK in header (signature won't match) + parts := splitJWT(proof) + header := parseJWTHeader(t, parts[0]) + header["jwk"] = wrongKey.jwk + modifiedHeader := encodeJSON(t, header) + tamperedProof := modifiedHeader + "." + parts[1] + "." + parts[2] + + _, err := verifier.VerifyDPoPProof(tamperedProof, method, uri) + if err == nil { + t.Error("Expected error for invalid signature, got nil") + } + if err != nil && !contains(err.Error(), "signature verification failed") { + t.Errorf("Expected signature verification error, got: %v", err) + } +} + +func TestVerifyDPoPProof_WrongHTTPMethod(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + wrongMethod := "GET" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, wrongMethod, uri) + if err == nil { + t.Error("Expected error for HTTP method mismatch, got nil") + } + if err != nil && !contains(err.Error(), "htm mismatch") { + t.Errorf("Expected htm mismatch error, got: %v", err) + } +} + +func TestVerifyDPoPProof_WrongURI(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + wrongURI := "https://api.example.com/different" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, method, wrongURI) + if err == nil { + t.Error("Expected error for URI mismatch, got nil") + } + if err != nil && !contains(err.Error(), "htu mismatch") { + t.Errorf("Expected htu mismatch error, got: %v", err) + } +} + +func TestVerifyDPoPProof_URIWithQuery(t *testing.T) { + // URI comparison should strip query and fragment + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + baseURI := "https://api.example.com/resource" + uriWithQuery := baseURI + "?param=value" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, baseURI, iat, jti) + + // Should succeed because query is stripped + _, err := verifier.VerifyDPoPProof(proof, method, uriWithQuery) + if err != nil { + t.Fatalf("VerifyDPoPProof failed for URI with query: %v", err) + } +} + +func TestVerifyDPoPProof_URIWithFragment(t *testing.T) { + // URI comparison should strip query and fragment + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + baseURI := "https://api.example.com/resource" + uriWithFragment := baseURI + "#section" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, baseURI, iat, jti) + + // Should succeed because fragment is stripped + _, err := verifier.VerifyDPoPProof(proof, method, uriWithFragment) + if err != nil { + t.Fatalf("VerifyDPoPProof failed for URI with fragment: %v", err) + } +} + +func TestVerifyDPoPProof_ExpiredProof(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + // Proof issued 10 minutes ago (exceeds default MaxProofAge of 5 minutes) + iat := time.Now().Add(-10 * time.Minute) + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for expired proof, got nil") + } + if err != nil && !contains(err.Error(), "too old") { + t.Errorf("Expected 'too old' error, got: %v", err) + } +} + +func TestVerifyDPoPProof_FutureProof(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + // Proof issued 1 minute in the future (exceeds MaxClockSkew) + iat := time.Now().Add(1 * time.Minute) + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for future proof, got nil") + } + if err != nil && !contains(err.Error(), "in the future") { + t.Errorf("Expected 'in the future' error, got: %v", err) + } +} + +func TestVerifyDPoPProof_WithinClockSkew(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + // Proof issued 15 seconds in the future (within MaxClockSkew of 30s) + iat := time.Now().Add(15 * time.Second) + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, method, uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed for proof within clock skew: %v", err) + } +} + +func TestVerifyDPoPProof_MissingJti(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + + claims := &DPoPClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + // No ID (jti) + IssuedAt: jwt.NewNumericDate(iat), + }, + HTTPMethod: method, + HTTPURI: uri, + } + + token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) + token.Header["typ"] = "dpop+jwt" + token.Header["jwk"] = key.jwk + + proof, err := token.SignedString(key.privateKey) + if err != nil { + t.Fatalf("Failed to create test proof: %v", err) + } + + _, err = verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for missing jti, got nil") + } + if err != nil && !contains(err.Error(), "missing jti") { + t.Errorf("Expected missing jti error, got: %v", err) + } +} + +func TestVerifyDPoPProof_MissingTypHeader(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + claims := &DPoPClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + ID: jti, + IssuedAt: jwt.NewNumericDate(iat), + }, + HTTPMethod: method, + HTTPURI: uri, + } + + token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) + // Don't set typ header + token.Header["jwk"] = key.jwk + + proof, err := token.SignedString(key.privateKey) + if err != nil { + t.Fatalf("Failed to create test proof: %v", err) + } + + _, err = verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for missing typ header, got nil") + } + if err != nil && !contains(err.Error(), "typ must be 'dpop+jwt'") { + t.Errorf("Expected typ header error, got: %v", err) + } +} + +func TestVerifyDPoPProof_WrongTypHeader(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + claims := &DPoPClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + ID: jti, + IssuedAt: jwt.NewNumericDate(iat), + }, + HTTPMethod: method, + HTTPURI: uri, + } + + token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) + token.Header["typ"] = "JWT" // Wrong typ + token.Header["jwk"] = key.jwk + + proof, err := token.SignedString(key.privateKey) + if err != nil { + t.Fatalf("Failed to create test proof: %v", err) + } + + _, err = verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for wrong typ header, got nil") + } + if err != nil && !contains(err.Error(), "typ must be 'dpop+jwt'") { + t.Errorf("Expected typ header error, got: %v", err) + } +} + +func TestVerifyDPoPProof_MissingJWK(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + claims := &DPoPClaims{ + RegisteredClaims: jwt.RegisteredClaims{ + ID: jti, + IssuedAt: jwt.NewNumericDate(iat), + }, + HTTPMethod: method, + HTTPURI: uri, + } + + token := jwt.NewWithClaims(jwt.SigningMethodES256, claims) + token.Header["typ"] = "dpop+jwt" + // Don't include JWK + + proof, err := token.SignedString(key.privateKey) + if err != nil { + t.Fatalf("Failed to create test proof: %v", err) + } + + _, err = verifier.VerifyDPoPProof(proof, method, uri) + if err == nil { + t.Error("Expected error for missing jwk header, got nil") + } + if err != nil && !contains(err.Error(), "missing jwk") { + t.Errorf("Expected missing jwk error, got: %v", err) + } +} + +func TestVerifyDPoPProof_CustomTimeSettings(t *testing.T) { + verifier := &DPoPVerifier{ + MaxClockSkew: 1 * time.Minute, + MaxProofAge: 10 * time.Minute, + } + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + // Proof issued 50 seconds in the future (within custom MaxClockSkew) + iat := time.Now().Add(50 * time.Second) + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + _, err := verifier.VerifyDPoPProof(proof, method, uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed with custom time settings: %v", err) + } +} + +func TestVerifyDPoPProof_HTTPMethodCaseInsensitive(t *testing.T) { + // HTTP method comparison should be case-insensitive per spec + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "post" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + // Verify with uppercase method + _, err := verifier.VerifyDPoPProof(proof, "POST", uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed for case-insensitive method: %v", err) + } +} + +// === Token Binding Verification Tests === + +func TestVerifyTokenBinding_Matching(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + result, err := verifier.VerifyDPoPProof(proof, method, uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed: %v", err) + } + + // Verify token binding with matching thumbprint + err = verifier.VerifyTokenBinding(result, key.thumbprint) + if err != nil { + t.Fatalf("VerifyTokenBinding failed for matching thumbprint: %v", err) + } +} + +func TestVerifyTokenBinding_Mismatch(t *testing.T) { + verifier := NewDPoPVerifier() + key := generateTestES256Key(t) + wrongKey := generateTestES256Key(t) + + method := "POST" + uri := "https://api.example.com/resource" + iat := time.Now() + jti := uuid.New().String() + + proof := createDPoPProof(t, key, method, uri, iat, jti) + + result, err := verifier.VerifyDPoPProof(proof, method, uri) + if err != nil { + t.Fatalf("VerifyDPoPProof failed: %v", err) + } + + // Verify token binding with wrong thumbprint + err = verifier.VerifyTokenBinding(result, wrongKey.thumbprint) + if err == nil { + t.Error("Expected error for thumbprint mismatch, got nil") + } + if err != nil && !contains(err.Error(), "thumbprint mismatch") { + t.Errorf("Expected thumbprint mismatch error, got: %v", err) + } +} + +// === ExtractCnfJkt Tests === + +func TestExtractCnfJkt_Valid(t *testing.T) { + expectedJkt := "test-thumbprint-123" + claims := &Claims{ + Confirmation: map[string]interface{}{ + "jkt": expectedJkt, + }, + } + + jkt, err := ExtractCnfJkt(claims) + if err != nil { + t.Fatalf("ExtractCnfJkt failed for valid claims: %v", err) + } + + if jkt != expectedJkt { + t.Errorf("Expected jkt %s, got %s", expectedJkt, jkt) + } +} + +func TestExtractCnfJkt_MissingCnf(t *testing.T) { + claims := &Claims{ + // No Confirmation + } + + _, err := ExtractCnfJkt(claims) + if err == nil { + t.Error("Expected error for missing cnf, got nil") + } + if err != nil && !contains(err.Error(), "missing cnf claim") { + t.Errorf("Expected missing cnf error, got: %v", err) + } +} + +func TestExtractCnfJkt_NilCnf(t *testing.T) { + claims := &Claims{ + Confirmation: nil, + } + + _, err := ExtractCnfJkt(claims) + if err == nil { + t.Error("Expected error for nil cnf, got nil") + } + if err != nil && !contains(err.Error(), "missing cnf claim") { + t.Errorf("Expected missing cnf error, got: %v", err) + } +} + +func TestExtractCnfJkt_MissingJkt(t *testing.T) { + claims := &Claims{ + Confirmation: map[string]interface{}{ + "other": "value", + }, + } + + _, err := ExtractCnfJkt(claims) + if err == nil { + t.Error("Expected error for missing jkt, got nil") + } + if err != nil && !contains(err.Error(), "missing jkt") { + t.Errorf("Expected missing jkt error, got: %v", err) + } +} + +func TestExtractCnfJkt_EmptyJkt(t *testing.T) { + claims := &Claims{ + Confirmation: map[string]interface{}{ + "jkt": "", + }, + } + + _, err := ExtractCnfJkt(claims) + if err == nil { + t.Error("Expected error for empty jkt, got nil") + } + if err != nil && !contains(err.Error(), "missing jkt") { + t.Errorf("Expected missing jkt error, got: %v", err) + } +} + +func TestExtractCnfJkt_WrongType(t *testing.T) { + claims := &Claims{ + Confirmation: map[string]interface{}{ + "jkt": 123, // Not a string + }, + } + + _, err := ExtractCnfJkt(claims) + if err == nil { + t.Error("Expected error for wrong type jkt, got nil") + } + if err != nil && !contains(err.Error(), "missing jkt") { + t.Errorf("Expected missing jkt error, got: %v", err) + } +} + +// === Helper Functions for Tests === + +// splitJWT splits a JWT into its three parts +func splitJWT(token string) []string { + return []string{ + token[:strings.IndexByte(token, '.')], + token[strings.IndexByte(token, '.')+1 : strings.LastIndexByte(token, '.')], + token[strings.LastIndexByte(token, '.')+1:], + } +} + +// parseJWTHeader parses a base64url-encoded JWT header +func parseJWTHeader(t *testing.T, encoded string) map[string]interface{} { + t.Helper() + decoded, err := base64.RawURLEncoding.DecodeString(encoded) + if err != nil { + t.Fatalf("Failed to decode header: %v", err) + } + + var header map[string]interface{} + if err := json.Unmarshal(decoded, &header); err != nil { + t.Fatalf("Failed to unmarshal header: %v", err) + } + + return header +} + +// encodeJSON encodes a value to base64url-encoded JSON +func encodeJSON(t *testing.T, v interface{}) string { + t.Helper() + data, err := json.Marshal(v) + if err != nil { + t.Fatalf("Failed to marshal JSON: %v", err) + } + return base64.RawURLEncoding.EncodeToString(data) +}