From 2ee88de9a0b9426fe1963c96de8ac61b55de8476 Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Wed, 20 May 2026 14:56:40 +0100 Subject: [PATCH 1/2] feat(secureservice): add admission primitives Introduce provider-neutral admission config and verifier types for future federated network admission support. The new config is disabled by default and does not change handshake behavior. --- net/secureservice/admission.go | 75 +++++++++++++++++++++++++++ net/secureservice/admission_test.go | 78 +++++++++++++++++++++++++++++ net/secureservice/config.go | 5 +- 3 files changed, 156 insertions(+), 2 deletions(-) create mode 100644 net/secureservice/admission.go create mode 100644 net/secureservice/admission_test.go diff --git a/net/secureservice/admission.go b/net/secureservice/admission.go new file mode 100644 index 00000000..a788f746 --- /dev/null +++ b/net/secureservice/admission.go @@ -0,0 +1,75 @@ +package secureservice + +import "context" + +const ( + DefaultAdmissionIdentityClaim = "anytype_identity" + DefaultAdmissionNetworkClaim = "network_id" + DefaultAdmissionSubjectClaim = "sub" + DefaultAdmissionClockSkewSec = 60 +) + +type AdmissionConfig struct { + Enabled bool `yaml:"enabled"` + Required bool `yaml:"required"` + Issuer string `yaml:"issuer"` + Audience string `yaml:"audience"` + JWKSURL string `yaml:"jwksUrl"` + RequiredClaims map[string]any `yaml:"requiredClaims"` + IdentityClaim string `yaml:"identityClaim"` + NetworkClaim string `yaml:"networkClaim"` + SubjectClaim string `yaml:"subjectClaim"` + ClockSkewSec int `yaml:"clockSkewSec"` +} + +func (c AdmissionConfig) WithDefaults() AdmissionConfig { + if c.IdentityClaim == "" { + c.IdentityClaim = DefaultAdmissionIdentityClaim + } + if c.NetworkClaim == "" { + c.NetworkClaim = DefaultAdmissionNetworkClaim + } + if c.SubjectClaim == "" { + c.SubjectClaim = DefaultAdmissionSubjectClaim + } + if c.ClockSkewSec == 0 { + c.ClockSkewSec = DefaultAdmissionClockSkewSec + } + return c +} + +type AdmissionRequest struct { + Token string + Identity []byte + NetworkID string + PeerID string + ClientVersion string +} + +type AdmissionClaims struct { + Subject string + Issuer string + Audience []string + NetworkID string + Identity []byte + Claims map[string]any +} + +type AdmissionDecision struct { + Allowed bool + Reason string + Claims AdmissionClaims +} + +type AdmissionVerifier interface { + VerifyAdmission(ctx context.Context, req AdmissionRequest) (AdmissionDecision, error) +} + +type NoopAdmissionVerifier struct{} + +func (NoopAdmissionVerifier) VerifyAdmission(ctx context.Context, req AdmissionRequest) (AdmissionDecision, error) { + return AdmissionDecision{ + Allowed: true, + Reason: "admission disabled", + }, nil +} diff --git a/net/secureservice/admission_test.go b/net/secureservice/admission_test.go new file mode 100644 index 00000000..8a59ca05 --- /dev/null +++ b/net/secureservice/admission_test.go @@ -0,0 +1,78 @@ +package secureservice + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gopkg.in/yaml.v3" +) + +func TestAdmissionConfig_WithDefaults(t *testing.T) { + conf := AdmissionConfig{}.WithDefaults() + + assert.False(t, conf.Enabled) + assert.False(t, conf.Required) + assert.Equal(t, DefaultAdmissionIdentityClaim, conf.IdentityClaim) + assert.Equal(t, DefaultAdmissionNetworkClaim, conf.NetworkClaim) + assert.Equal(t, DefaultAdmissionSubjectClaim, conf.SubjectClaim) + assert.Equal(t, DefaultAdmissionClockSkewSec, conf.ClockSkewSec) +} + +func TestAdmissionConfig_WithDefaultsPreservesConfiguredValues(t *testing.T) { + conf := AdmissionConfig{ + Enabled: true, + Required: true, + IdentityClaim: "identity", + NetworkClaim: "network", + SubjectClaim: "subject", + ClockSkewSec: 120, + }.WithDefaults() + + assert.True(t, conf.Enabled) + assert.True(t, conf.Required) + assert.Equal(t, "identity", conf.IdentityClaim) + assert.Equal(t, "network", conf.NetworkClaim) + assert.Equal(t, "subject", conf.SubjectClaim) + assert.Equal(t, 120, conf.ClockSkewSec) +} + +func TestAdmissionConfig_YAML(t *testing.T) { + var conf Config + err := yaml.Unmarshal([]byte(` +admission: + enabled: true + required: true + issuer: https://issuer.example + audience: any-sync + jwksUrl: https://issuer.example/.well-known/jwks.json + requiredClaims: + group: docs-users + identityClaim: anytype_id + networkClaim: sync_network + subjectClaim: user + clockSkewSec: 30 +`), &conf) + require.NoError(t, err) + + assert.True(t, conf.Admission.Enabled) + assert.True(t, conf.Admission.Required) + assert.Equal(t, "https://issuer.example", conf.Admission.Issuer) + assert.Equal(t, "any-sync", conf.Admission.Audience) + assert.Equal(t, "https://issuer.example/.well-known/jwks.json", conf.Admission.JWKSURL) + assert.Equal(t, map[string]any{"group": "docs-users"}, conf.Admission.RequiredClaims) + assert.Equal(t, "anytype_id", conf.Admission.IdentityClaim) + assert.Equal(t, "sync_network", conf.Admission.NetworkClaim) + assert.Equal(t, "user", conf.Admission.SubjectClaim) + assert.Equal(t, 30, conf.Admission.ClockSkewSec) +} + +func TestNoopAdmissionVerifier_AllowsRequests(t *testing.T) { + verifier := NoopAdmissionVerifier{} + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{}) + + require.NoError(t, err) + assert.True(t, decision.Allowed) + assert.Equal(t, "admission disabled", decision.Reason) +} diff --git a/net/secureservice/config.go b/net/secureservice/config.go index 7d5eb3ae..44a841ce 100644 --- a/net/secureservice/config.go +++ b/net/secureservice/config.go @@ -13,8 +13,9 @@ type configGetter interface { } type Config struct { - RequireClientAuth bool `yaml:"requireClientAuth"` - CompatibleVersions []uint32 `yaml:"compatibleVersions"` + RequireClientAuth bool `yaml:"requireClientAuth"` + CompatibleVersions []uint32 `yaml:"compatibleVersions"` + Admission AdmissionConfig `yaml:"admission"` } // CtxAllowAccountCheck upgrades the context to allow identity check on handshake From 11be3faef84dafcbbbfee9e839795ca890f3a42d Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Wed, 20 May 2026 15:07:28 +0100 Subject: [PATCH 2/2] feat(secureservice): add JWT admission verifier Add a static JWKS-backed AdmissionVerifier implementation for provider-neutral federated admission. The verifier validates token signature, issuer, audience, expiry, network id, Anytype identity binding, subject, and required claims without wiring it into the handshake yet. --- net/secureservice/admission_jwt.go | 495 ++++++++++++++++++++++++ net/secureservice/admission_jwt_test.go | 227 +++++++++++ 2 files changed, 722 insertions(+) create mode 100644 net/secureservice/admission_jwt.go create mode 100644 net/secureservice/admission_jwt_test.go diff --git a/net/secureservice/admission_jwt.go b/net/secureservice/admission_jwt.go new file mode 100644 index 00000000..2d39bd71 --- /dev/null +++ b/net/secureservice/admission_jwt.go @@ -0,0 +1,495 @@ +package secureservice + +import ( + "bytes" + "context" + stdcrypto "crypto" + "crypto/ecdsa" + "crypto/ed25519" + "crypto/elliptic" + "crypto/rsa" + "crypto/sha256" + "crypto/sha512" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "math/big" + "strings" + "time" + + anycrypto "github.com/anyproto/any-sync/util/crypto" +) + +var ( + // ErrAdmissionDenied is returned when an admission request is rejected. + ErrAdmissionDenied = errors.New("admission denied") + // ErrAdmissionInvalidConfig is returned when JWT admission verifier configuration is incomplete or invalid. + ErrAdmissionInvalidConfig = errors.New("invalid admission config") + // ErrAdmissionInvalidToken is returned when an admission token is malformed or fails validation. + ErrAdmissionInvalidToken = errors.New("invalid admission token") +) + +type jwtAdmissionVerifier struct { + conf AdmissionConfig + keys map[string]jwtPublicKey + now func() time.Time +} + +type jwtPublicKey struct { + algorithm string + key any +} + +type jwtHeader struct { + Algorithm string `json:"alg"` + KeyID string `json:"kid"` +} + +type jsonWebKeySet struct { + Keys []jsonWebKey `json:"keys"` +} + +type jsonWebKey struct { + KeyType string `json:"kty"` + KeyID string `json:"kid"` + Use string `json:"use"` + Algorithm string `json:"alg"` + N string `json:"n"` + E string `json:"e"` + Curve string `json:"crv"` + X string `json:"x"` + Y string `json:"y"` +} + +// NewJWTAdmissionVerifier creates an AdmissionVerifier that validates JWTs against a static JWKS document. +func NewJWTAdmissionVerifier(conf AdmissionConfig, jwks []byte) (AdmissionVerifier, error) { + return newJWTAdmissionVerifier(conf, jwks, time.Now) +} + +func newJWTAdmissionVerifier(conf AdmissionConfig, jwks []byte, now func() time.Time) (*jwtAdmissionVerifier, error) { + conf = conf.WithDefaults() + if conf.Issuer == "" || conf.Audience == "" || conf.IdentityClaim == "" || conf.NetworkClaim == "" || conf.SubjectClaim == "" { + return nil, ErrAdmissionInvalidConfig + } + keys, err := parseJWKS(jwks) + if err != nil { + return nil, err + } + if len(keys) == 0 { + return nil, ErrAdmissionInvalidConfig + } + if now == nil { + now = time.Now + } + return &jwtAdmissionVerifier{ + conf: conf, + keys: keys, + now: now, + }, nil +} + +func (v *jwtAdmissionVerifier) VerifyAdmission(ctx context.Context, req AdmissionRequest) (AdmissionDecision, error) { + if err := ctx.Err(); err != nil { + return AdmissionDecision{Allowed: false, Reason: err.Error()}, err + } + if req.Token == "" || len(req.Identity) == 0 || req.NetworkID == "" { + return v.deny("missing admission token, identity, or network id", ErrAdmissionInvalidToken) + } + + header, claims, signingInput, signature, err := parseJWT(req.Token) + if err != nil { + return v.deny("malformed admission token", err) + } + key, ok := v.keys[header.KeyID] + if !ok { + return v.deny("unknown signing key", ErrAdmissionInvalidToken) + } + if err = verifyJWTSignature(header.Algorithm, key, signingInput, signature); err != nil { + return v.deny("invalid admission token signature", err) + } + + decisionClaims, err := v.validateClaims(req, claims) + if err != nil { + return v.deny(err.Error(), err) + } + return AdmissionDecision{ + Allowed: true, + Reason: "admission allowed", + Claims: decisionClaims, + }, nil +} + +func (v *jwtAdmissionVerifier) deny(reason string, err error) (AdmissionDecision, error) { + return AdmissionDecision{Allowed: false, Reason: reason}, errors.Join(ErrAdmissionDenied, err) +} + +func parseJWT(token string) (jwtHeader, map[string]any, []byte, []byte, error) { + parts := strings.Split(token, ".") + if len(parts) != 3 { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + headerBytes, err := base64.RawURLEncoding.DecodeString(parts[0]) + if err != nil { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + payloadBytes, err := base64.RawURLEncoding.DecodeString(parts[1]) + if err != nil { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + signature, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + + var header jwtHeader + if err = decodeJSON(headerBytes, &header); err != nil { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + if header.Algorithm == "" || header.KeyID == "" { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + + claims := map[string]any{} + if err = decodeJSON(payloadBytes, &claims); err != nil { + return jwtHeader{}, nil, nil, nil, ErrAdmissionInvalidToken + } + return header, claims, []byte(parts[0] + "." + parts[1]), signature, nil +} + +func decodeJSON(data []byte, v any) error { + dec := json.NewDecoder(bytes.NewReader(data)) + dec.UseNumber() + return dec.Decode(v) +} + +func parseJWKS(jwks []byte) (map[string]jwtPublicKey, error) { + var set jsonWebKeySet + if err := decodeJSON(jwks, &set); err != nil { + return nil, ErrAdmissionInvalidConfig + } + keys := make(map[string]jwtPublicKey, len(set.Keys)) + for _, key := range set.Keys { + if key.KeyID == "" || key.Use == "enc" { + continue + } + parsed, err := parseJWK(key) + if err != nil { + return nil, err + } + keys[key.KeyID] = parsed + } + return keys, nil +} + +func parseJWK(key jsonWebKey) (jwtPublicKey, error) { + switch key.KeyType { + case "RSA": + return parseRSAJWK(key) + case "EC": + return parseECJWK(key) + case "OKP": + return parseOKPJWK(key) + default: + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } +} + +func parseRSAJWK(key jsonWebKey) (jwtPublicKey, error) { + n, err := base64.RawURLEncoding.DecodeString(key.N) + if err != nil { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + e, err := base64.RawURLEncoding.DecodeString(key.E) + if err != nil { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + exponent := int(new(big.Int).SetBytes(e).Int64()) + if exponent == 0 || len(n) == 0 { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + return jwtPublicKey{ + algorithm: key.Algorithm, + key: &rsa.PublicKey{ + N: new(big.Int).SetBytes(n), + E: exponent, + }, + }, nil +} + +func parseECJWK(key jsonWebKey) (jwtPublicKey, error) { + x, err := base64.RawURLEncoding.DecodeString(key.X) + if err != nil { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + y, err := base64.RawURLEncoding.DecodeString(key.Y) + if err != nil { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + curve := curveForJWK(key.Curve) + xCoord := new(big.Int).SetBytes(x) + yCoord := new(big.Int).SetBytes(y) + if curve == nil || len(x) == 0 || len(y) == 0 || !curve.IsOnCurve(xCoord, yCoord) { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + return jwtPublicKey{ + algorithm: key.Algorithm, + key: &ecdsa.PublicKey{ + Curve: curve, + X: xCoord, + Y: yCoord, + }, + }, nil +} + +func curveForJWK(curve string) elliptic.Curve { + switch curve { + case "P-256": + return elliptic.P256() + case "P-384": + return elliptic.P384() + case "P-521": + return elliptic.P521() + default: + return nil + } +} + +func parseOKPJWK(key jsonWebKey) (jwtPublicKey, error) { + if key.Curve != "Ed25519" { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + x, err := base64.RawURLEncoding.DecodeString(key.X) + if err != nil || len(x) != ed25519.PublicKeySize { + return jwtPublicKey{}, ErrAdmissionInvalidConfig + } + return jwtPublicKey{ + algorithm: key.Algorithm, + key: ed25519.PublicKey(x), + }, nil +} + +func verifyJWTSignature(algorithm string, key jwtPublicKey, signingInput []byte, signature []byte) error { + if key.algorithm != "" && key.algorithm != algorithm { + return ErrAdmissionInvalidToken + } + switch algorithm { + case "RS256": + digest := sha256.Sum256(signingInput) + return verifyRSA(key.key, stdcrypto.SHA256, digest[:], signature) + case "RS384": + digest := sha512.Sum384(signingInput) + return verifyRSA(key.key, stdcrypto.SHA384, digest[:], signature) + case "RS512": + digest := sha512.Sum512(signingInput) + return verifyRSA(key.key, stdcrypto.SHA512, digest[:], signature) + case "ES256": + digest := sha256.Sum256(signingInput) + return verifyECDSA(key.key, digest[:], signature, 32) + case "ES384": + digest := sha512.Sum384(signingInput) + return verifyECDSA(key.key, digest[:], signature, 48) + case "ES512": + digest := sha512.Sum512(signingInput) + return verifyECDSA(key.key, digest[:], signature, 66) + case "EdDSA": + pub, ok := key.key.(ed25519.PublicKey) + if !ok || !ed25519.Verify(pub, signingInput, signature) { + return ErrAdmissionInvalidToken + } + return nil + default: + return ErrAdmissionInvalidToken + } +} + +func verifyRSA(key any, hash stdcrypto.Hash, digest []byte, signature []byte) error { + pub, ok := key.(*rsa.PublicKey) + if !ok { + return ErrAdmissionInvalidToken + } + if err := rsa.VerifyPKCS1v15(pub, hash, digest, signature); err != nil { + return ErrAdmissionInvalidToken + } + return nil +} + +func verifyECDSA(key any, digest []byte, signature []byte, keyBytes int) error { + pub, ok := key.(*ecdsa.PublicKey) + if !ok || len(signature) != keyBytes*2 { + return ErrAdmissionInvalidToken + } + r := new(big.Int).SetBytes(signature[:keyBytes]) + s := new(big.Int).SetBytes(signature[keyBytes:]) + if !ecdsa.Verify(pub, digest, r, s) { + return ErrAdmissionInvalidToken + } + return nil +} + +func (v *jwtAdmissionVerifier) validateClaims(req AdmissionRequest, claims map[string]any) (AdmissionClaims, error) { + now := v.now() + if err := validateStringClaim(claims, "iss", v.conf.Issuer); err != nil { + return AdmissionClaims{}, err + } + audience, err := validateAudienceClaim(claims, v.conf.Audience) + if err != nil { + return AdmissionClaims{}, err + } + if err = validateExpiration(claims, now, time.Duration(v.conf.ClockSkewSec)*time.Second); err != nil { + return AdmissionClaims{}, err + } + if err = validateStringClaim(claims, v.conf.NetworkClaim, req.NetworkID); err != nil { + return AdmissionClaims{}, err + } + accountID, err := accountIDFromIdentity(req.Identity) + if err != nil { + return AdmissionClaims{}, ErrAdmissionInvalidToken + } + if err = validateStringClaim(claims, v.conf.IdentityClaim, accountID); err != nil { + return AdmissionClaims{}, err + } + subject, err := stringClaim(claims, v.conf.SubjectClaim) + if err != nil || subject == "" { + return AdmissionClaims{}, ErrAdmissionInvalidToken + } + for name, expected := range v.conf.RequiredClaims { + actual, ok := claims[name] + if !ok || !claimMatches(actual, expected) { + return AdmissionClaims{}, fmt.Errorf("required claim %q is missing or invalid: %w", name, ErrAdmissionInvalidToken) + } + } + return AdmissionClaims{ + Subject: subject, + Issuer: v.conf.Issuer, + Audience: audience, + NetworkID: req.NetworkID, + Identity: req.Identity, + Claims: claims, + }, nil +} + +func accountIDFromIdentity(identity []byte) (string, error) { + pub, err := anycrypto.UnmarshalEd25519PublicKeyProto(identity) + if err != nil { + return "", err + } + return pub.Account(), nil +} + +func validateStringClaim(claims map[string]any, name string, expected string) error { + actual, err := stringClaim(claims, name) + if err != nil || actual != expected { + return ErrAdmissionInvalidToken + } + return nil +} + +func stringClaim(claims map[string]any, name string) (string, error) { + value, ok := claims[name].(string) + if !ok { + return "", ErrAdmissionInvalidToken + } + return value, nil +} + +func validateAudienceClaim(claims map[string]any, expected string) ([]string, error) { + audience := stringListClaim(claims["aud"]) + for _, value := range audience { + if value == expected { + return audience, nil + } + } + return nil, ErrAdmissionInvalidToken +} + +func stringListClaim(value any) []string { + switch value := value.(type) { + case string: + return []string{value} + case []any: + result := make([]string, 0, len(value)) + for _, item := range value { + if str, ok := item.(string); ok { + result = append(result, str) + } + } + return result + default: + return nil + } +} + +func validateExpiration(claims map[string]any, now time.Time, skew time.Duration) error { + exp, ok, err := numericDateClaim(claims, "exp") + if err != nil || !ok || now.After(exp.Add(skew)) { + return ErrAdmissionInvalidToken + } + if nbf, ok, err := numericDateClaim(claims, "nbf"); err != nil || ok && now.Add(skew).Before(nbf) { + return ErrAdmissionInvalidToken + } + if iat, ok, err := numericDateClaim(claims, "iat"); err != nil || ok && now.Add(skew).Before(iat) { + return ErrAdmissionInvalidToken + } + return nil +} + +func numericDateClaim(claims map[string]any, name string) (time.Time, bool, error) { + value, ok := claims[name] + if !ok { + return time.Time{}, false, nil + } + seconds, err := int64ClaimValue(value) + if err != nil { + return time.Time{}, false, err + } + return time.Unix(seconds, 0), true, nil +} + +func int64ClaimValue(value any) (int64, error) { + switch value := value.(type) { + case json.Number: + return value.Int64() + case float64: + return int64(value), nil + case int64: + return value, nil + case int: + return int64(value), nil + default: + return 0, ErrAdmissionInvalidToken + } +} + +func claimMatches(actual any, expected any) bool { + if values, ok := actual.([]any); ok { + for _, value := range values { + if claimMatches(value, expected) { + return true + } + } + return false + } + switch expected := expected.(type) { + case string: + actual, ok := actual.(string) + return ok && actual == expected + case bool: + actual, ok := actual.(bool) + return ok && actual == expected + case int: + return claimNumberMatches(actual, int64(expected)) + case int64: + return claimNumberMatches(actual, expected) + case json.Number: + expectedInt, err := expected.Int64() + return err == nil && claimNumberMatches(actual, expectedInt) + default: + return actual == expected + } +} + +func claimNumberMatches(actual any, expected int64) bool { + actualInt, err := int64ClaimValue(actual) + return err == nil && actualInt == expected +} diff --git a/net/secureservice/admission_jwt_test.go b/net/secureservice/admission_jwt_test.go new file mode 100644 index 00000000..882d4cbb --- /dev/null +++ b/net/secureservice/admission_jwt_test.go @@ -0,0 +1,227 @@ +package secureservice + +import ( + "context" + stdcrypto "crypto" + "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "errors" + "math/big" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestJWTAdmissionVerifier_VerifyAdmission(t *testing.T) { + now := time.Unix(1700000000, 0) + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + account := newTestAccData(t) + identity, err := account.SignKey.GetPublic().Marshall() + require.NoError(t, err) + + verifier := newTestJWTAdmissionVerifier(t, privateKey, now) + token := signAdmissionToken(t, privateKey, map[string]any{ + "iss": "https://issuer.example", + "aud": []string{"any-sync"}, + "sub": "user-1", + "network_id": "network-1", + "anytype_identity": account.SignKey.GetPublic().Account(), + "groups": []string{"docs-users"}, + "exp": now.Add(time.Hour).Unix(), + "nbf": now.Add(-time.Minute).Unix(), + "iat": now.Add(-time.Minute).Unix(), + }) + + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{ + Token: token, + Identity: identity, + NetworkID: "network-1", + }) + + require.NoError(t, err) + assert.True(t, decision.Allowed) + assert.Equal(t, "admission allowed", decision.Reason) + assert.Equal(t, "user-1", decision.Claims.Subject) + assert.Equal(t, "network-1", decision.Claims.NetworkID) + assert.Equal(t, identity, decision.Claims.Identity) +} + +func TestJWTAdmissionVerifier_DeniesWrongAudience(t *testing.T) { + now := time.Unix(1700000000, 0) + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + account := newTestAccData(t) + identity, err := account.SignKey.GetPublic().Marshall() + require.NoError(t, err) + + verifier := newTestJWTAdmissionVerifier(t, privateKey, now) + token := signAdmissionToken(t, privateKey, map[string]any{ + "iss": "https://issuer.example", + "aud": "other-audience", + "sub": "user-1", + "network_id": "network-1", + "anytype_identity": account.SignKey.GetPublic().Account(), + "groups": []string{"docs-users"}, + "exp": now.Add(time.Hour).Unix(), + }) + + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{ + Token: token, + Identity: identity, + NetworkID: "network-1", + }) + + require.Error(t, err) + assert.True(t, errors.Is(err, ErrAdmissionDenied)) + assert.False(t, decision.Allowed) +} + +func TestJWTAdmissionVerifier_DeniesWrongIdentity(t *testing.T) { + now := time.Unix(1700000000, 0) + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + account := newTestAccData(t) + otherAccount := newTestAccData(t) + identity, err := account.SignKey.GetPublic().Marshall() + require.NoError(t, err) + + verifier := newTestJWTAdmissionVerifier(t, privateKey, now) + token := signAdmissionToken(t, privateKey, map[string]any{ + "iss": "https://issuer.example", + "aud": "any-sync", + "sub": "user-1", + "network_id": "network-1", + "anytype_identity": otherAccount.SignKey.GetPublic().Account(), + "groups": []string{"docs-users"}, + "exp": now.Add(time.Hour).Unix(), + }) + + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{ + Token: token, + Identity: identity, + NetworkID: "network-1", + }) + + require.Error(t, err) + assert.True(t, errors.Is(err, ErrAdmissionDenied)) + assert.False(t, decision.Allowed) +} + +func TestJWTAdmissionVerifier_DeniesTamperedToken(t *testing.T) { + now := time.Unix(1700000000, 0) + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + account := newTestAccData(t) + identity, err := account.SignKey.GetPublic().Marshall() + require.NoError(t, err) + + verifier := newTestJWTAdmissionVerifier(t, privateKey, now) + token := signAdmissionToken(t, privateKey, map[string]any{ + "iss": "https://issuer.example", + "aud": "any-sync", + "sub": "user-1", + "network_id": "network-1", + "anytype_identity": account.SignKey.GetPublic().Account(), + "groups": []string{"docs-users"}, + "exp": now.Add(time.Hour).Unix(), + }) + tokenParts := strings.Split(token, ".") + require.Len(t, tokenParts, 3) + tokenParts[2] = base64.RawURLEncoding.EncodeToString([]byte("bad-signature")) + + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{ + Token: strings.Join(tokenParts, "."), + Identity: identity, + NetworkID: "network-1", + }) + + require.Error(t, err) + assert.True(t, errors.Is(err, ErrAdmissionDenied)) + assert.False(t, decision.Allowed) +} + +func TestJWTAdmissionVerifier_DeniesExpiredToken(t *testing.T) { + now := time.Unix(1700000000, 0) + privateKey, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + account := newTestAccData(t) + identity, err := account.SignKey.GetPublic().Marshall() + require.NoError(t, err) + + verifier := newTestJWTAdmissionVerifier(t, privateKey, now) + token := signAdmissionToken(t, privateKey, map[string]any{ + "iss": "https://issuer.example", + "aud": "any-sync", + "sub": "user-1", + "network_id": "network-1", + "anytype_identity": account.SignKey.GetPublic().Account(), + "groups": []string{"docs-users"}, + "exp": now.Add(-2 * time.Minute).Unix(), + }) + + decision, err := verifier.VerifyAdmission(context.Background(), AdmissionRequest{ + Token: token, + Identity: identity, + NetworkID: "network-1", + }) + + require.Error(t, err) + assert.True(t, errors.Is(err, ErrAdmissionDenied)) + assert.False(t, decision.Allowed) +} + +func newTestJWTAdmissionVerifier(t *testing.T, privateKey *rsa.PrivateKey, now time.Time) *jwtAdmissionVerifier { + verifier, err := newJWTAdmissionVerifier(AdmissionConfig{ + Issuer: "https://issuer.example", + Audience: "any-sync", + RequiredClaims: map[string]any{"groups": "docs-users"}, + }, rsaJWKS(t, &privateKey.PublicKey, "test-key"), func() time.Time { return now }) + require.NoError(t, err) + return verifier +} + +func signAdmissionToken(t *testing.T, privateKey *rsa.PrivateKey, claims map[string]any) string { + header := map[string]any{ + "alg": "RS256", + "kid": "test-key", + } + encodedHeader := encodeJWTPart(t, header) + encodedClaims := encodeJWTPart(t, claims) + signingInput := encodedHeader + "." + encodedClaims + digest := sha256.Sum256([]byte(signingInput)) + signature, err := rsa.SignPKCS1v15(rand.Reader, privateKey, stdcrypto.SHA256, digest[:]) + require.NoError(t, err) + return signingInput + "." + base64.RawURLEncoding.EncodeToString(signature) +} + +func rsaJWKS(t *testing.T, publicKey *rsa.PublicKey, keyID string) []byte { + return mustJSON(t, map[string]any{ + "keys": []map[string]any{ + { + "kty": "RSA", + "kid": keyID, + "use": "sig", + "alg": "RS256", + "n": base64.RawURLEncoding.EncodeToString(publicKey.N.Bytes()), + "e": base64.RawURLEncoding.EncodeToString(big.NewInt(int64(publicKey.E)).Bytes()), + }, + }, + }) +} + +func encodeJWTPart(t *testing.T, value any) string { + return base64.RawURLEncoding.EncodeToString(mustJSON(t, value)) +} + +func mustJSON(t *testing.T, value any) []byte { + data, err := json.Marshal(value) + require.NoError(t, err) + return data +}