From 02094e72abdf8214de468e5c5490641bf242a74b Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 18:23:21 +0530 Subject: [PATCH 1/2] feat(internal): jwtsign + public-request auth exemptions --- contrib/server/enforce_auth_test.go | 58 ++++ go.mod | 1 + go.sum | 2 + internal/jwtsign/jwtsign.go | 276 ++++++++++++++++++ internal/jwtsign/jwtsign_test.go | 308 +++++++++++++++++++ server/aws/authgate.go | 18 ++ server/aws/aws.go | 2 +- server/aws/publicauth.go | 275 +++++++++++++++++ server/aws/publicauth_test.go | 438 ++++++++++++++++++++++++++++ server/server.go | 29 +- 10 files changed, 1398 insertions(+), 9 deletions(-) create mode 100644 internal/jwtsign/jwtsign.go create mode 100644 internal/jwtsign/jwtsign_test.go create mode 100644 server/aws/publicauth.go create mode 100644 server/aws/publicauth_test.go diff --git a/contrib/server/enforce_auth_test.go b/contrib/server/enforce_auth_test.go index 7737f10a8..df13a0268 100644 --- a/contrib/server/enforce_auth_test.go +++ b/contrib/server/enforce_auth_test.go @@ -2,6 +2,8 @@ package main import ( "context" + "io" + "net/http" "strings" "testing" @@ -33,6 +35,32 @@ func TestEnforceAuthRejectsInvalidSignature(t *testing.T) { } }) + t.Run("public operations skip the gate", func(t *testing.T) { + cfg := testConfig(t, allEnginesOff()) + cfg.EnforceAuth = true + + awsURL, stop := startAWS(t, cfg, mustOptions(t, &cfg)) + defer stop() + + cases := []struct { + name, target, host, path string + wantGate bool + }{ + {name: "InitiateAuth", target: "AWSCognitoIdentityProviderService.InitiateAuth", path: "/"}, + {name: "CreateUserPool", target: "AWSCognitoIdentityProviderService.CreateUserPool", path: "/", wantGate: true}, + {name: "execute-api", host: "abc123.execute-api.us-east-1.amazonaws.com", path: "/prod/pets"}, + } + + for _, tc := range cases { + status, body := unsignedPost(t, awsURL+tc.path, tc.host, tc.target) + gated := status == http.StatusForbidden && strings.Contains(body, "MissingAuthenticationToken") + + if gated != tc.wantGate { + t.Fatalf("%s: status %d body %s, gate rejection = %v, want %v", tc.name, status, body, gated, tc.wantGate) + } + } + }) + t.Run("not enforced passes", func(t *testing.T) { cfg := testConfig(t, allEnginesOff()) cfg.EnforceAuth = false @@ -47,3 +75,33 @@ func TestEnforceAuthRejectsInvalidSignature(t *testing.T) { } }) } + +// unsignedPost sends a POST with no SigV4 Authorization header, optionally +// overriding Host and setting X-Amz-Target, and returns the status and body. +func unsignedPost(t *testing.T, url, host, target string) (int, string) { + t.Helper() + + req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, url, strings.NewReader("{}")) + if err != nil { + t.Fatalf("new request: %v", err) + } + + if host != "" { + req.Host = host + } + + if target != "" { + req.Header.Set("X-Amz-Target", target) + req.Header.Set("Content-Type", "application/x-amz-json-1.1") + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("do: %v", err) + } + defer resp.Body.Close() + + b, _ := io.ReadAll(resp.Body) + + return resp.StatusCode, string(b) +} diff --git a/go.mod b/go.mod index f72168235..273b190d4 100644 --- a/go.mod +++ b/go.mod @@ -78,6 +78,7 @@ require ( github.com/aws/aws-sdk-go-v2/service/cloudwatch v1.56.2 github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0 github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0 + github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0 github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0 github.com/aws/aws-sdk-go-v2/service/configservice v1.68.4 github.com/aws/aws-sdk-go-v2/service/costexplorer v1.63.6 diff --git a/go.sum b/go.sum index db65c75d3..7434a7237 100644 --- a/go.sum +++ b/go.sum @@ -210,6 +210,8 @@ github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0 h1:5W/KOwsnZrdi7RD97 github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0/go.mod h1:h1Iw2nkdpmAUJaa89RvX3cg/HGLgdSkCWpMNgKvBSHA= github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0 h1:zssdCGwtqMsIrkYFCWu4yhdHHbp+OxbHoFSWp0QDsZo= github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0/go.mod h1:b8QnRsunCLqoz5Elpypig/tbI23E2N0betuvS6wwbsE= +github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0 h1:Wp96KzoBhY4T8nIU/PK/Lf6mCD9nelkXuY6SvvKvdgY= +github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0/go.mod h1:Nqfwixdr/+xgqWGHVkOcVytNVy3VCxFCzWr+8XxEn10= github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0 h1:3Vje2gVkUDNSksJ8NXLcLCSg5m/YtsTqSNfDupy3qeI= github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0/go.mod h1:ygltZT++6Wn2uG4+tqE0NW1MkdEtb5W2O/CFc0xJX/g= github.com/aws/aws-sdk-go-v2/service/configservice v1.68.4 h1:L37DoV4JOQFqucVUeYaJrrhIc1995EPra+8UtGRCsgo= diff --git a/internal/jwtsign/jwtsign.go b/internal/jwtsign/jwtsign.go new file mode 100644 index 000000000..8d4dd933a --- /dev/null +++ b/internal/jwtsign/jwtsign.go @@ -0,0 +1,276 @@ +// Package jwtsign issues and verifies RS256 JSON Web Tokens and publishes the +// matching JSON Web Key Set. It is the shared signer for the emulated identity +// services (Cognito user pools first) that must hand out tokens a real client +// library can verify against the service's JWKS endpoint. +// +// Each issuer owns one or more Keys. A Key's ID (the JWT "kid") is its RFC 7638 +// JWK thumbprint, so it is stable across persistence and never collides between +// keys. Rotation is adding a new Key and signing with it while the old one stays +// in the verification set (and the JWKS) until its tokens have expired. +package jwtsign + +import ( + "crypto" + "crypto/rand" + "crypto/rsa" + "crypto/sha256" + "crypto/x509" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "math/big" + "strings" + "time" + + "github.com/stackshy/cloudemu/v2/config" +) + +// algRS256 is the only signing algorithm this package issues or accepts. +const algRS256 = "RS256" + +const ( + keyBits = 2048 + jwtSegments = 3 +) + +// Errors returned by Verify. ErrExpired is kept apart from ErrInvalidToken so a +// service can answer with its distinct "token has expired" message. +var ( + ErrInvalidToken = errors.New("jwtsign: invalid token") + ErrUnknownKey = errors.New("jwtsign: unknown signing key") + ErrExpired = errors.New("jwtsign: token has expired") + ErrNotYetValid = errors.New("jwtsign: token is not yet valid") + + errNotRSA = errors.New("jwtsign: parse key: not an RSA key") +) + +// Key is an RS256 signing key and the kid it is published under. +type Key struct { + ID string + Private *rsa.PrivateKey +} + +// NewRSAKey generates a 2048-bit RSA key whose ID is its JWK thumbprint. +func NewRSAKey() (*Key, error) { + priv, err := rsa.GenerateKey(rand.Reader, keyBits) + if err != nil { + return nil, fmt.Errorf("jwtsign: generate key: %w", err) + } + + return &Key{ID: thumbprint(&priv.PublicKey), Private: priv}, nil +} + +// MarshalPKCS8 encodes the private key as PKCS#8 DER, the form snapshots store. +func MarshalPKCS8(k *Key) ([]byte, error) { + der, err := x509.MarshalPKCS8PrivateKey(k.Private) + if err != nil { + return nil, fmt.Errorf("jwtsign: marshal key: %w", err) + } + + return der, nil +} + +// ParsePKCS8 decodes a PKCS#8 DER RSA private key. The kid is recomputed from +// the public key, so it matches the one the key had when it was marshaled. +func ParsePKCS8(der []byte) (*Key, error) { + parsed, err := x509.ParsePKCS8PrivateKey(der) + if err != nil { + return nil, fmt.Errorf("jwtsign: parse key: %w", err) + } + + priv, ok := parsed.(*rsa.PrivateKey) + if !ok { + return nil, fmt.Errorf("%w: got %T", errNotRSA, parsed) + } + + return &Key{ID: thumbprint(&priv.PublicKey), Private: priv}, nil +} + +type header struct { + Alg string `json:"alg"` + Kid string `json:"kid"` +} + +// Sign returns the compact RS256 JWT for claims, with k.ID as the header kid. +func Sign(k *Key, claims map[string]any) (string, error) { + h, err := json.Marshal(header{Alg: algRS256, Kid: k.ID}) + if err != nil { + return "", fmt.Errorf("jwtsign: header: %w", err) + } + + p, err := json.Marshal(claims) + if err != nil { + return "", fmt.Errorf("jwtsign: claims: %w", err) + } + + signingInput := b64(h) + "." + b64(p) + sum := sha256.Sum256([]byte(signingInput)) + + sig, err := rsa.SignPKCS1v15(rand.Reader, k.Private, crypto.SHA256, sum[:]) + if err != nil { + return "", fmt.Errorf("jwtsign: sign: %w", err) + } + + return signingInput + "." + b64(sig), nil +} + +// Verify checks token against keys and the clock and returns its claims. It +// requires alg RS256, a kid present in keys, a valid signature, and an exp +// claim. It rejects a token at or after exp, before nbf, or with iat in the +// future. Numeric claims come back as json.Number. +func Verify(token string, keys []*Key, clock config.Clock) (map[string]any, error) { + parts := strings.Split(token, ".") + if len(parts) != jwtSegments { + return nil, fmt.Errorf("%w: want %d segments", ErrInvalidToken, jwtSegments) + } + + var h header + if err := decodeJSON(parts[0], &h); err != nil { + return nil, err + } + + if h.Alg != algRS256 { + return nil, fmt.Errorf("%w: alg %q", ErrInvalidToken, h.Alg) + } + + key := findKey(keys, h.Kid) + if key == nil { + return nil, fmt.Errorf("%w: kid %q", ErrUnknownKey, h.Kid) + } + + sig, err := base64.RawURLEncoding.DecodeString(parts[2]) + if err != nil { + return nil, fmt.Errorf("%w: signature encoding", ErrInvalidToken) + } + + sum := sha256.Sum256([]byte(parts[0] + "." + parts[1])) + if err := rsa.VerifyPKCS1v15(&key.Private.PublicKey, crypto.SHA256, sum[:], sig); err != nil { + return nil, fmt.Errorf("%w: signature", ErrInvalidToken) + } + + var claims map[string]any + if err := decodeJSON(parts[1], &claims); err != nil { + return nil, err + } + + if err := checkTimes(claims, clock.Now()); err != nil { + return nil, err + } + + return claims, nil +} + +func checkTimes(claims map[string]any, now time.Time) error { + exp, ok := numericDate(claims, "exp") + if !ok { + return fmt.Errorf("%w: missing or non-numeric exp", ErrInvalidToken) + } + + unix := now.Unix() + + if unix >= exp { + return ErrExpired + } + + if nbf, ok := numericDate(claims, "nbf"); ok && unix < nbf { + return ErrNotYetValid + } + + if iat, ok := numericDate(claims, "iat"); ok && unix < iat { + return ErrNotYetValid + } + + return nil +} + +func numericDate(claims map[string]any, name string) (int64, bool) { + n, ok := claims[name].(json.Number) + if !ok { + return 0, false + } + + v, err := n.Int64() + if err != nil { + f, ferr := n.Float64() + if ferr != nil { + return 0, false + } + + v = int64(f) + } + + return v, true +} + +func findKey(keys []*Key, kid string) *Key { + for _, k := range keys { + if k.ID == kid { + return k + } + } + + return nil +} + +// JWK is one public key in a JSON Web Key Set, in the field set Cognito's +// jwks.json publishes. +type JWK struct { + Alg string `json:"alg"` + E string `json:"e"` + Kid string `json:"kid"` + Kty string `json:"kty"` + N string `json:"n"` + Use string `json:"use"` +} + +// JWKSet is a JSON Web Key Set document ({"keys": [...]}). +type JWKSet struct { + Keys []JWK `json:"keys"` +} + +// JWKS returns the public half of keys, in order, as a JSON Web Key Set. +func JWKS(keys ...*Key) JWKSet { + set := JWKSet{Keys: make([]JWK, 0, len(keys))} + + for _, k := range keys { + pub := &k.Private.PublicKey + set.Keys = append(set.Keys, JWK{ + Alg: algRS256, + E: b64(big.NewInt(int64(pub.E)).Bytes()), + Kid: k.ID, + Kty: "RSA", + N: b64(pub.N.Bytes()), + Use: "sig", + }) + } + + return set +} + +// thumbprint is the RFC 7638 JWK thumbprint of an RSA public key: the base64url +// SHA-256 of the canonical {"e","kty","n"} member set. +func thumbprint(pub *rsa.PublicKey) string { + canonical := `{"e":"` + b64(big.NewInt(int64(pub.E)).Bytes()) + `","kty":"RSA","n":"` + b64(pub.N.Bytes()) + `"}` + sum := sha256.Sum256([]byte(canonical)) + + return b64(sum[:]) +} + +func decodeJSON(seg string, v any) error { + raw, err := base64.RawURLEncoding.DecodeString(seg) + if err != nil { + return fmt.Errorf("%w: segment encoding", ErrInvalidToken) + } + + dec := json.NewDecoder(strings.NewReader(string(raw))) + dec.UseNumber() + + if err := dec.Decode(v); err != nil { + return fmt.Errorf("%w: segment json", ErrInvalidToken) + } + + return nil +} + +func b64(b []byte) string { return base64.RawURLEncoding.EncodeToString(b) } diff --git a/internal/jwtsign/jwtsign_test.go b/internal/jwtsign/jwtsign_test.go new file mode 100644 index 000000000..866ec991d --- /dev/null +++ b/internal/jwtsign/jwtsign_test.go @@ -0,0 +1,308 @@ +package jwtsign + +import ( + "crypto/rsa" + "encoding/base64" + "encoding/json" + "errors" + "math/big" + "strings" + "testing" + "time" + + "github.com/stackshy/cloudemu/v2/config" +) + +var epoch = time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC) //nolint:gochecknoglobals // fixed test instant + +func mustKey(t *testing.T) *Key { + t.Helper() + + k, err := NewRSAKey() + if err != nil { + t.Fatalf("NewRSAKey: %v", err) + } + + return k +} + +func claimsAt(now time.Time, ttl time.Duration) map[string]any { + return map[string]any{ + "sub": "user-1", + "token_use": "access", + "iat": now.Unix(), + "nbf": now.Unix(), + "exp": now.Add(ttl).Unix(), + } +} + +func TestSignVerifyRoundTrip(t *testing.T) { + k := mustKey(t) + clock := config.NewFakeClock(epoch) + + tok, err := Sign(k, claimsAt(epoch, time.Hour)) + if err != nil { + t.Fatalf("Sign: %v", err) + } + + if n := strings.Count(tok, "."); n != 2 { + t.Fatalf("token has %d dots, want 2", n) + } + + got, err := Verify(tok, []*Key{k}, clock) + if err != nil { + t.Fatalf("Verify: %v", err) + } + + if got["sub"] != "user-1" || got["token_use"] != "access" { + t.Fatalf("claims = %v", got) + } + + hdr := decodeSegment(t, strings.Split(tok, ".")[0]) + if hdr["alg"] != "RS256" || hdr["kid"] != k.ID { + t.Fatalf("header = %v, want alg RS256 kid %s", hdr, k.ID) + } +} + +func TestVerifyRejectsTamperedPayload(t *testing.T) { + k := mustKey(t) + + tok, err := Sign(k, claimsAt(epoch, time.Hour)) + if err != nil { + t.Fatalf("Sign: %v", err) + } + + parts := strings.Split(tok, ".") + forged := claimsAt(epoch, time.Hour) + forged["sub"] = "admin" + b, _ := json.Marshal(forged) + parts[1] = base64.RawURLEncoding.EncodeToString(b) + + _, err = Verify(strings.Join(parts, "."), []*Key{k}, config.NewFakeClock(epoch)) + if !errors.Is(err, ErrInvalidToken) { + t.Fatalf("tampered token: err = %v, want ErrInvalidToken", err) + } +} + +func TestVerifyRejectsWrongKeyWithSameKID(t *testing.T) { + signer := mustKey(t) + other := mustKey(t) + other.ID = signer.ID + + tok, _ := Sign(signer, claimsAt(epoch, time.Hour)) + + if _, err := Verify(tok, []*Key{other}, config.NewFakeClock(epoch)); !errors.Is(err, ErrInvalidToken) { + t.Fatalf("err = %v, want ErrInvalidToken", err) + } +} + +func TestVerifyRejectsAlgNoneAndHS256(t *testing.T) { + k := mustKey(t) + payload := segmentOf(t, claimsAt(epoch, time.Hour)) + + for _, alg := range []string{"none", "HS256", ""} { + hdr := segmentOf(t, map[string]any{"alg": alg, "kid": k.ID}) + tok := hdr + "." + payload + "." + + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrInvalidToken) { + t.Fatalf("alg %q: err = %v, want ErrInvalidToken", alg, err) + } + } +} + +func TestVerifyRejectsMalformed(t *testing.T) { + k := mustKey(t) + + for _, tok := range []string{"", "a.b", "a.b.c.d", "!!.!!.!!", "e30.e30.e30"} { + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrInvalidToken) && + !errors.Is(err, ErrUnknownKey) { + t.Fatalf("token %q: err = %v, want rejection", tok, err) + } + } +} + +func TestVerifyExpiry(t *testing.T) { + k := mustKey(t) + clock := config.NewFakeClock(epoch) + + tok, _ := Sign(k, claimsAt(epoch, time.Hour)) + + clock.Advance(time.Hour - time.Second) + + if _, err := Verify(tok, []*Key{k}, clock); err != nil { + t.Fatalf("one second before exp: %v", err) + } + + clock.Advance(time.Second) + + if _, err := Verify(tok, []*Key{k}, clock); !errors.Is(err, ErrExpired) { + t.Fatalf("at exp: err = %v, want ErrExpired", err) + } +} + +func TestVerifyNotBeforeAndIssuedInFuture(t *testing.T) { + k := mustKey(t) + future := epoch.Add(time.Minute) + + nbf := claimsAt(epoch, time.Hour) + nbf["nbf"] = future.Unix() + tok, _ := Sign(k, nbf) + + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrNotYetValid) { + t.Fatalf("nbf in future: err = %v, want ErrNotYetValid", err) + } + + iat := claimsAt(epoch, time.Hour) + iat["iat"] = future.Unix() + delete(iat, "nbf") + tok, _ = Sign(k, iat) + + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrNotYetValid) { + t.Fatalf("iat in future: err = %v, want ErrNotYetValid", err) + } +} + +func TestVerifyRequiresExp(t *testing.T) { + k := mustKey(t) + c := claimsAt(epoch, time.Hour) + delete(c, "exp") + tok, _ := Sign(k, c) + + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrInvalidToken) { + t.Fatalf("missing exp: err = %v, want ErrInvalidToken", err) + } +} + +func TestVerifyUnknownKID(t *testing.T) { + signer := mustKey(t) + other := mustKey(t) + tok, _ := Sign(signer, claimsAt(epoch, time.Hour)) + + if _, err := Verify(tok, []*Key{other}, config.NewFakeClock(epoch)); !errors.Is(err, ErrUnknownKey) { + t.Fatalf("err = %v, want ErrUnknownKey", err) + } +} + +// TestKIDRotation: after a new key is added, tokens signed with the old key +// still verify while the old key stays in the set, new tokens carry the new +// kid, and dropping the old key retires its tokens. +func TestKIDRotation(t *testing.T) { + oldKey := mustKey(t) + newKey := mustKey(t) + clock := config.NewFakeClock(epoch) + + if oldKey.ID == newKey.ID { + t.Fatalf("two keys share kid %q", oldKey.ID) + } + + oldTok, _ := Sign(oldKey, claimsAt(epoch, time.Hour)) + newTok, _ := Sign(newKey, claimsAt(epoch, time.Hour)) + + both := []*Key{newKey, oldKey} + for _, tok := range []string{oldTok, newTok} { + if _, err := Verify(tok, both, clock); err != nil { + t.Fatalf("verify during rotation: %v", err) + } + } + + if _, err := Verify(oldTok, []*Key{newKey}, clock); !errors.Is(err, ErrUnknownKey) { + t.Fatalf("retired key: err = %v, want ErrUnknownKey", err) + } +} + +func TestJWKSPublishesPublicKeys(t *testing.T) { + k1, k2 := mustKey(t), mustKey(t) + + set := JWKS(k1, k2) + + raw, err := json.Marshal(set) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + if strings.Contains(string(raw), `"d"`) || strings.Contains(string(raw), `"p"`) { + t.Fatalf("JWKS leaks private material: %s", raw) + } + + var doc struct { + Keys []map[string]string `json:"keys"` + } + if err := json.Unmarshal(raw, &doc); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if len(doc.Keys) != 2 { + t.Fatalf("got %d keys, want 2", len(doc.Keys)) + } + + for i, k := range []*Key{k1, k2} { + j := doc.Keys[i] + if j["kid"] != k.ID || j["kty"] != "RSA" || j["alg"] != "RS256" || j["use"] != "sig" || j["e"] != "AQAB" { + t.Fatalf("jwk %d = %v", i, j) + } + + n, err := base64.RawURLEncoding.DecodeString(j["n"]) + if err != nil { + t.Fatalf("n not base64url: %v", err) + } + + pub := &rsa.PublicKey{N: new(big.Int).SetBytes(n), E: k.Private.E} + if !pub.Equal(&k.Private.PublicKey) { + t.Fatalf("jwk %d modulus does not match the key", i) + } + } +} + +func TestPKCS8RoundTripKeepsKID(t *testing.T) { + k := mustKey(t) + + der, err := MarshalPKCS8(k) + if err != nil { + t.Fatalf("MarshalPKCS8: %v", err) + } + + back, err := ParsePKCS8(der) + if err != nil { + t.Fatalf("ParsePKCS8: %v", err) + } + + if back.ID != k.ID || !back.Private.Equal(k.Private) { + t.Fatalf("round trip changed the key (kid %q -> %q)", k.ID, back.ID) + } + + tok, _ := Sign(k, claimsAt(epoch, time.Hour)) + if _, err := Verify(tok, []*Key{back}, config.NewFakeClock(epoch)); err != nil { + t.Fatalf("token from original does not verify with restored key: %v", err) + } + + if _, err := ParsePKCS8([]byte("junk")); err == nil { + t.Fatalf("ParsePKCS8(junk) succeeded") + } +} + +func decodeSegment(t *testing.T, seg string) map[string]any { + t.Helper() + + b, err := base64.RawURLEncoding.DecodeString(seg) + if err != nil { + t.Fatalf("decode: %v", err) + } + + var m map[string]any + if err := json.Unmarshal(b, &m); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + return m +} + +func segmentOf(t *testing.T, v any) string { + t.Helper() + + b, err := json.Marshal(v) + if err != nil { + t.Fatalf("marshal: %v", err) + } + + return base64.RawURLEncoding.EncodeToString(b) +} diff --git a/server/aws/authgate.go b/server/aws/authgate.go index 3e4292031..150f3ef1f 100644 --- a/server/aws/authgate.go +++ b/server/aws/authgate.go @@ -7,6 +7,7 @@ import ( "strings" "github.com/stackshy/cloudemu/v2/config" + "github.com/stackshy/cloudemu/v2/server" "github.com/stackshy/cloudemu/v2/server/authctx" stssrv "github.com/stackshy/cloudemu/v2/server/aws/sts" "github.com/stackshy/cloudemu/v2/server/wire" @@ -28,8 +29,12 @@ const tempCredentialPrefix = "ASIA" // the STS session store (temporary ASIA credentials), verifies the signature, // and either attaches the resolved principal to the request context (proceed) // or writes a 403 AWS error (stop). clock drives timestamp-expiry evaluation. +// match is the dispatcher's handler lookup. It binds the public-operation +// exemption (see exemptPublic) to the handler that will actually serve the +// request. func newAuthGate( iamDriver iamdriver.IAM, accountID string, sessions *stssrv.SessionStore, clock config.Clock, + match func(*http.Request) server.Handler, ) func(http.ResponseWriter, *http.Request) (*http.Request, bool) { resolver, _ := iamDriver.(iamdriver.AccessKeyResolver) @@ -41,6 +46,19 @@ func newAuthGate( body := drainBody(r) restore := func() { r.Body = io.NopCloser(bytes.NewReader(body)) } + // Operations AWS serves without SigV4 (noAuth) skip authentication and + // authorization. The handler lookup may read the body, so restore it + // before and after. + restore() + + public := exemptPublic(r, body, match) + + restore() + + if public { + return r, true + } + akid := sigv4.AccessKeyID(r) if akid == "" { restore() diff --git a/server/aws/aws.go b/server/aws/aws.go index 30cb1e884..b9a6bd637 100644 --- a/server/aws/aws.go +++ b/server/aws/aws.go @@ -1239,7 +1239,7 @@ func New(d Drivers) *server.Server { // on, and adds no request-path change beyond a context value. var authGate func(http.ResponseWriter, *http.Request) (*http.Request, bool) if d.EnforceAuth { - authGate = newAuthGate(d.IAM, d.AccountID, stsSessions, authClock) + authGate = newAuthGate(d.IAM, d.AccountID, stsSessions, authClock, srv.Match) } srv.SetPreDispatch(composePreDispatch(newRegionStamp(), authGate)) diff --git a/server/aws/publicauth.go b/server/aws/publicauth.go new file mode 100644 index 000000000..b86cafb70 --- /dev/null +++ b/server/aws/publicauth.go @@ -0,0 +1,275 @@ +package aws + +import ( + "net" + "net/http" + "net/url" + "regexp" + "strings" + + "github.com/stackshy/cloudemu/v2/server" + apigatewaysrv "github.com/stackshy/cloudemu/v2/server/aws/apigateway" + appsyncsrv "github.com/stackshy/cloudemu/v2/server/aws/appsync" + cognitosrv "github.com/stackshy/cloudemu/v2/server/aws/cognito" + stssrv "github.com/stackshy/cloudemu/v2/server/aws/sts" +) + +// Some AWS operations are called without SigV4 credentials: a user signs in to +// Cognito before it has any AWS credentials, a web-identity token is traded for +// credentials, and API Gateway / AppSync endpoints are hit by browsers. The +// operation tables below are the operations whose Smithy model carries +// "auth": ["smithy.api#noAuth"] (botocore: "authtype": "none"), taken from the +// botocore service models (awscli 2.31.19, botocore/data//*/service-2.json). +// TestPublicOpsMatchSDKModels re-derives the same sets from the generated +// aws-sdk-go-v2 auth resolvers and fails when a table drifts from the model. +// Of the other services with noAuth operations, sso and sso-oidc are not served. + +// cognitoIDPPublicOps are the cognito-idp (AWSCognitoIdentityProviderService) +// noAuth operations. They authenticate with an access token, a session, or a +// client secret hash inside the request, never with SigV4. +// +//nolint:gochecknoglobals // static protocol lookup table +var cognitoIDPPublicOps = map[string]struct{}{ + "AssociateSoftwareToken": {}, + "ChangePassword": {}, + "CompleteWebAuthnRegistration": {}, + "ConfirmDevice": {}, + "ConfirmForgotPassword": {}, + "ConfirmSignUp": {}, + "DeleteUser": {}, + "DeleteUserAttributes": {}, + "DeleteWebAuthnCredential": {}, + "ForgetDevice": {}, + "ForgotPassword": {}, + "GetDevice": {}, + "GetTokensFromRefreshToken": {}, + "GetUser": {}, + "GetUserAttributeVerificationCode": {}, + "GetUserAuthFactors": {}, + "GlobalSignOut": {}, + "InitiateAuth": {}, + "ListDevices": {}, + "ListWebAuthnCredentials": {}, + "ResendConfirmationCode": {}, + "RespondToAuthChallenge": {}, + "RevokeToken": {}, + "SetUserMFAPreference": {}, + "SetUserSettings": {}, + "SignUp": {}, + "StartWebAuthnRegistration": {}, + "UpdateAuthEventFeedback": {}, + "UpdateDeviceStatus": {}, + "UpdateUserAttributes": {}, + "VerifySoftwareToken": {}, + "VerifyUserAttribute": {}, +} + +// cognitoIdentityPublicOps are the cognito-identity (AWSCognitoIdentityService) +// noAuth operations, which authenticate with the identity-pool logins map. +// +//nolint:gochecknoglobals // static protocol lookup table +var cognitoIdentityPublicOps = map[string]struct{}{ + "GetCredentialsForIdentity": {}, + "GetId": {}, + "GetOpenIdToken": {}, + "UnlinkIdentity": {}, +} + +// stsPublicActions are the STS noAuth actions: the caller presents a web +// identity token or a SAML assertion instead of AWS credentials. +// +//nolint:gochecknoglobals // static protocol lookup table +var stsPublicActions = map[string]struct{}{ + "AssumeRoleWithSAML": {}, + "AssumeRoleWithWebIdentity": {}, +} + +const ( + cognitoIDPTargetPrefix = "AWSCognitoIdentityProviderService." + cognitoIdentityTargetPrefix = "AWSCognitoIdentityService." + wellKnownSegment = "/.well-known/" + hostedUIPathPrefix = "/_cognito/" + restAPIsPrefix = "/restapis/" + userRequestSegment = "_user_request_" + userRequestSegmentIndex = 2 // {apiId}/{stage}/_user_request_/... + userRequestSplitParts = 4 // the three leading segments plus the rest + graphQLPath = "/graphql" + urlEncodedForm = "application/x-www-form-urlencoded" +) + +// userPoolIDPattern is the Cognito user pool id shape ({region}_{suffix}). The +// underscore makes it an illegal S3 bucket name, so a .well-known path under it +// never names a bucket. +var userPoolIDPattern = regexp.MustCompile(`^[a-z]{2}(-[a-z]+)+-\d_[0-9A-Za-z]+$`) + +// publicRoute is one family of requests AWS serves without SigV4. match decides +// from the request shape alone. owner reports whether the handler that would +// actually serve the request is the one that owns this public surface; nil +// means no served handler owns it yet. +type publicRoute struct { + match func(r *http.Request, body []byte) bool + owner func(server.Handler) bool +} + +// publicRoutes lists every public surface. A route whose owner handler does not +// claim the request yet (the Cognito JWKS and hosted UI, the AppSync GraphQL +// endpoint) stays gated. It opens once that handler starts claiming it. +// +//nolint:gochecknoglobals // static route table +var publicRoutes = []publicRoute{ + {match: targetIn(cognitoIDPTargetPrefix, cognitoIDPPublicOps), owner: ownedBy[*cognitosrv.Handler]}, + {match: targetIn(cognitoIdentityTargetPrefix, cognitoIdentityPublicOps)}, + {match: isUserPoolWellKnown, owner: ownedBy[*cognitosrv.Handler]}, + {match: isHostedUI, owner: ownedBy[*cognitosrv.Handler]}, + {match: isPublicSTSAction, owner: ownedBy[*stssrv.Handler]}, + {match: isExecuteAPI, owner: ownedBy[*apigatewaysrv.Handler]}, + {match: isUnsignedGraphQL, owner: ownedBy[*appsyncsrv.Handler]}, +} + +// exemptPublic reports whether r may skip the SigV4 gate. The request must have +// the shape of a public operation, and the handler that would serve it (per +// match, the dispatcher's own first-match lookup) must be that operation's +// owner. Binding to the dispatch target stops a request from borrowing a public +// marker, such as an execute-api Host or a public Action, while being routed to +// a different service. When no handler would serve the request it is exempt too, +// since the dispatcher then answers 501 and touches no state. +// +// Matches may parse the body, so the caller restores it afterwards. Exempt +// requests skip authorization as well: IAM does not govern noAuth operations. +func exemptPublic(r *http.Request, body []byte, match func(*http.Request) server.Handler) bool { + rt := publicRouteFor(r, body) + if rt == nil { + return false + } + + h := match(r) + if h == nil { + return true + } + + return rt.owner != nil && rt.owner(h) +} + +// publicRouteFor returns the public route r's shape matches, or nil. +func publicRouteFor(r *http.Request, body []byte) *publicRoute { + for i := range publicRoutes { + if publicRoutes[i].match(r, body) { + return &publicRoutes[i] + } + } + + return nil +} + +func ownedBy[T server.Handler](h server.Handler) bool { + _, ok := h.(T) + return ok +} + +// targetIn matches a JSON-RPC request whose X-Amz-Target is prefix+op for an +// op in ops. +func targetIn(prefix string, ops map[string]struct{}) func(*http.Request, []byte) bool { + return func(r *http.Request, _ []byte) bool { + op, ok := strings.CutPrefix(r.Header.Get("X-Amz-Target"), prefix) + if !ok { + return false + } + + _, public := ops[op] + + return public + } +} + +// isPublicSTSAction matches a query-protocol request whose Action (form body +// first, then the query string, the order net/http's ParseForm gives) is a +// public STS action. +func isPublicSTSAction(r *http.Request, body []byte) bool { + if r.Header.Get("X-Amz-Target") != "" { + return false + } + + action := "" + + if strings.HasPrefix(r.Header.Get("Content-Type"), urlEncodedForm) { + if form, err := url.ParseQuery(string(body)); err == nil { + action = form.Get("Action") + } + } + + if action == "" { + action = r.URL.Query().Get("Action") + } + + _, public := stsPublicActions[action] + + return public +} + +// isUserPoolWellKnown matches GET /{userPoolId}/.well-known/... (jwks.json and +// openid-configuration). +func isUserPoolWellKnown(r *http.Request, _ []byte) bool { + if r.Method != http.MethodGet && r.Method != http.MethodHead { + return false + } + + pool, rest, ok := strings.Cut(strings.TrimPrefix(r.URL.Path, "/"), "/") + if !ok { + return false + } + + return strings.HasPrefix("/"+rest, wellKnownSegment) && userPoolIDPattern.MatchString(pool) +} + +// isHostedUI matches the Cognito hosted domain ({domain}.auth.{region}.amazoncognito.com, +// or {domain}.auth.localhost locally) and its path fallback /_cognito/{domain}/... +func isHostedUI(r *http.Request, _ []byte) bool { + if strings.HasPrefix(r.URL.Path, hostedUIPathPrefix) { + return true + } + + host := hostOnly(r.Host) + + return strings.Contains(host, ".auth.") && + (strings.HasSuffix(host, ".amazoncognito.com") || strings.HasSuffix(host, ".auth.localhost")) +} + +// isExecuteAPI matches an API Gateway invocation: an execute-api host, or the +// path form /restapis/{apiId}/{stage}/_user_request_/... +func isExecuteAPI(r *http.Request, _ []byte) bool { + if strings.Contains(r.Host, ".execute-api.") { + return true + } + + rest, ok := strings.CutPrefix(r.URL.Path, restAPIsPrefix) + if !ok { + return false + } + + segs := strings.SplitN(rest, "/", userRequestSplitParts) + + return len(segs) > userRequestSegmentIndex && segs[userRequestSegmentIndex] == userRequestSegment +} + +// isUnsignedGraphQL matches an AppSync GraphQL call that carries no SigV4 +// Authorization header (API key, Cognito or OIDC auth modes). A SigV4-signed +// call uses the IAM auth mode and goes through the gate. +func isUnsignedGraphQL(r *http.Request, _ []byte) bool { + if r.Header.Get("Authorization") != "" { + return false + } + + if strings.Contains(r.Host, ".appsync-api.") { + return true + } + + return r.URL.Path == graphQLPath || strings.HasPrefix(r.URL.Path, graphQLPath+"/") +} + +func hostOnly(hostport string) string { + if h, _, err := net.SplitHostPort(hostport); err == nil { + return h + } + + return hostport +} diff --git a/server/aws/publicauth_test.go b/server/aws/publicauth_test.go new file mode 100644 index 000000000..d772c546a --- /dev/null +++ b/server/aws/publicauth_test.go @@ -0,0 +1,438 @@ +package aws + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "reflect" + "strings" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" + "github.com/aws/aws-sdk-go-v2/service/cognitoidentity" + "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider" + awssts "github.com/aws/aws-sdk-go-v2/service/sts" + smithyauth "github.com/aws/smithy-go/auth" + + cloudemu "github.com/stackshy/cloudemu/v2" + awsprovider "github.com/stackshy/cloudemu/v2/providers/aws" + iamdriver "github.com/stackshy/cloudemu/v2/services/iam/driver" +) + +// anonymousOps asks an SDK client's default auth scheme resolver, which is +// generated from the service's Smithy model, which of the client's operations +// resolve to the anonymous scheme (smithy.api#noAuth). Operations are the +// client's exported methods other than Options. +func anonymousOps(t *testing.T, client any, resolve func(op string) []*smithyauth.Option) map[string]bool { + t.Helper() + + out := map[string]bool{} + typ := reflect.TypeOf(client) + + for i := range typ.NumMethod() { + op := typ.Method(i).Name + if op == "Options" { + continue + } + + opts := resolve(op) + out[op] = len(opts) > 0 && opts[0].SchemeID == smithyauth.SchemeIDAnonymous + } + + return out +} + +// newerModelNoAuth lists operations the botocore models the tables were taken +// from mark noAuth while the aws-sdk-go-v2 version pinned in go.mod still signs +// them. Drop an entry once the pinned SDK catches up. +// +//nolint:gochecknoglobals // test fixture +var newerModelNoAuth = map[string]struct{}{ + "cognito-idp GetTokensFromRefreshToken": {}, +} + +// checkAgainstModel asserts the exemption table agrees with the SDK model: every +// operation the model marks noAuth is exempt, and every exempt operation the SDK +// knows is noAuth (bar newerModelNoAuth). Exempt operations newer than the +// pinned SDK are tolerated. +func checkAgainstModel(t *testing.T, service string, model map[string]bool, exempt map[string]struct{}) { + t.Helper() + + anon := 0 + + for op, isAnon := range model { + _, ok := exempt[op] + if isAnon { + anon++ + } + + if _, newer := newerModelNoAuth[service+" "+op]; newer && !isAnon && ok { + continue + } + + if isAnon != ok { + t.Errorf("%s %s: model noAuth=%v, exempt=%v", service, op, isAnon, ok) + } + } + + if anon == 0 { + t.Fatalf("%s: SDK model reports no anonymous operations; resolver probe is broken", service) + } +} + +func TestPublicOpsMatchSDKModels(t *testing.T) { + ctx := context.Background() + + idp := cognitoidentityprovider.New(cognitoidentityprovider.Options{Region: "us-east-1"}) + checkAgainstModel(t, "cognito-idp", anonymousOps(t, idp, func(op string) []*smithyauth.Option { + o, _ := idp.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, + &cognitoidentityprovider.AuthResolverParameters{Operation: op, Region: "us-east-1"}) + return o + }), cognitoIDPPublicOps) + + ident := cognitoidentity.New(cognitoidentity.Options{Region: "us-east-1"}) + checkAgainstModel(t, "cognito-identity", anonymousOps(t, ident, func(op string) []*smithyauth.Option { + o, _ := ident.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, + &cognitoidentity.AuthResolverParameters{Operation: op, Region: "us-east-1"}) + return o + }), cognitoIdentityPublicOps) + + sts := awssts.New(awssts.Options{Region: "us-east-1"}) + checkAgainstModel(t, "sts", anonymousOps(t, sts, func(op string) []*smithyauth.Option { + o, _ := sts.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, + &awssts.AuthResolverParameters{Operation: op, Region: "us-east-1"}) + return o + }), stsPublicActions) +} + +// enforcedServer starts the full AWS wire server with --enforce-auth on. +func enforcedServer(t *testing.T) (*httptest.Server, *awsprovider.Provider) { + t.Helper() + + cloud := cloudemu.NewAWS() + d := DriversFrom(cloud) + d.EnforceAuth = true + + ts := httptest.NewServer(New(d)) + t.Cleanup(ts.Close) + + return ts, cloud +} + +type rawReq struct { + method, path, host, body string + header map[string]string +} + +func doRaw(t *testing.T, ts *httptest.Server, rq rawReq) (int, string) { + t.Helper() + + req, err := http.NewRequestWithContext(context.Background(), rq.method, ts.URL+rq.path, strings.NewReader(rq.body)) + if err != nil { + t.Fatalf("new request: %v", err) + } + + if rq.host != "" { + req.Host = rq.host + } + + for k, v := range rq.header { + req.Header.Set(k, v) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("do: %v", err) + } + defer resp.Body.Close() + + b, _ := io.ReadAll(resp.Body) + + return resp.StatusCode, string(b) +} + +const ( + formCT = "application/x-www-form-urlencoded" + amzJSON11 = "application/x-amz-json-1.1" + missingTok = "MissingAuthenticationToken" + idpTarget = "AWSCognitoIdentityProviderService." + identTgt = "AWSCognitoIdentityService." + lambdaPath = "/2015-03-31/functions" +) + +func jsonRPC(target, body string) rawReq { + return rawReq{ + method: http.MethodPost, path: "/", body: body, + header: map[string]string{"X-Amz-Target": target, "Content-Type": amzJSON11}, + } +} + +func queryForm(body string) rawReq { + return rawReq{method: http.MethodPost, path: "/", body: body, header: map[string]string{"Content-Type": formCT}} +} + +// TestEnforcedGateAdmitsUnsignedPublicOps sends unsigned requests for operations +// AWS serves without SigV4 and asserts each reaches its handler instead of the +// gate's 403 MissingAuthenticationToken. +func TestEnforcedGateAdmitsUnsignedPublicOps(t *testing.T) { + ts, _ := enforcedServer(t) + + cases := []struct { + name string + req rawReq + want int + }{ + {"sts AssumeRoleWithWebIdentity", queryForm("Action=AssumeRoleWithWebIdentity&Version=2011-06-15" + + "&RoleArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fweb&RoleSessionName=s&WebIdentityToken=tok"), http.StatusOK}, + {"sts AssumeRoleWithSAML", queryForm("Action=AssumeRoleWithSAML&Version=2011-06-15" + + "&RoleArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fsaml" + + "&PrincipalArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Asaml-provider%2Fidp&SAMLAssertion=eA%3D%3D"), http.StatusOK}, + // Cognito user-pool public ops are not routed yet, so the handler answers + // UnknownOperationException (400): the point is the gate let them through. + {"cognito-idp InitiateAuth", jsonRPC(idpTarget+"InitiateAuth", `{}`), http.StatusBadRequest}, + {"cognito-idp SignUp", jsonRPC(idpTarget+"SignUp", `{}`), http.StatusBadRequest}, + {"cognito-idp RespondToAuthChallenge", jsonRPC(idpTarget+"RespondToAuthChallenge", `{}`), http.StatusBadRequest}, + // No identity-pool handler is served yet: the request falls to the + // dispatcher's 501, not the gate's 403. + {"cognito-identity GetId", jsonRPC(identTgt+"GetId", `{}`), http.StatusNotImplemented}, + // An unknown API reaches the API Gateway data plane, which answers like + // real API Gateway: 403 {"message":"Missing Authentication Token"}. That + // body differs from the gate's MissingAuthenticationToken error. + {"execute-api host", rawReq{method: http.MethodGet, path: "/prod/pets", host: "abc123.execute-api.us-east-1.amazonaws.com"}, + http.StatusForbidden}, + {"execute-api path", rawReq{method: http.MethodGet, path: "/restapis/abc123/prod/_user_request_/pets"}, http.StatusForbidden}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + status, body := doRaw(t, ts, tc.req) + if strings.Contains(body, missingTok) || status != tc.want { + t.Fatalf("status %d (want %d), body %s", status, tc.want, body) + } + }) + } +} + +// TestEnforcedGateRejectsUnsignedPrivateOps asserts everything that is not a +// public operation of the handler that actually serves it is still rejected +// when unsigned, including requests that borrow a public marker (an +// execute-api Host, a public Action) but would dispatch elsewhere. +func TestEnforcedGateRejectsUnsignedPrivateOps(t *testing.T) { + ts, _ := enforcedServer(t) + + cases := []struct { + name string + req rawReq + }{ + {"ec2 DescribeInstances", queryForm("Action=DescribeInstances&Version=2016-11-15")}, + {"sts GetCallerIdentity", queryForm("Action=GetCallerIdentity&Version=2011-06-15")}, + {"sts AssumeRole", queryForm("Action=AssumeRole&Version=2011-06-15&RoleArn=x&RoleSessionName=s")}, + {"cognito-idp CreateUserPool", jsonRPC(idpTarget+"CreateUserPool", `{"PoolName":"p"}`)}, + {"cognito-idp AdminInitiateAuth", jsonRPC(idpTarget+"AdminInitiateAuth", `{}`)}, + {"cognito-identity CreateIdentityPool", jsonRPC(identTgt+"CreateIdentityPool", `{}`)}, + {"dynamodb ListTables", jsonRPC("DynamoDB_20120810.ListTables", `{}`)}, + {"s3 ListBuckets", rawReq{method: http.MethodGet, path: "/"}}, + {"execute-api host on a lambda path", rawReq{method: http.MethodGet, path: lambdaPath, + host: "abc123.execute-api.us-east-1.amazonaws.com"}}, + {"public Action on a lambda path", rawReq{method: http.MethodGet, path: lambdaPath + "?Action=AssumeRoleWithWebIdentity"}}, + // Surfaces that are exempt in real AWS but not served yet fall to S3 here, + // so they must stay gated until their own handler claims them. + {"jwks before cognito serves it", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json"}}, + {"appsync graphql before appsync serves it", rawReq{method: http.MethodPost, path: "/graphql", body: `{}`}}, + {"hosted ui before cognito serves it", rawReq{method: http.MethodGet, path: "/oauth2/userInfo", + host: "mydomain.auth.us-east-1.amazoncognito.com"}}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + status, body := doRaw(t, ts, tc.req) + if status != http.StatusForbidden || !strings.Contains(body, missingTok) { + t.Fatalf("status %d, body %s; want 403 %s", status, body, missingTok) + } + }) + } +} + +func TestPublicRequestShapes(t *testing.T) { + mk := func(method, target, host, path, auth string) *http.Request { + r := httptest.NewRequest(method, "http://localhost"+path, nil) + if host != "" { + r.Host = host + } + + if target != "" { + r.Header.Set("X-Amz-Target", target) + } + + if auth != "" { + r.Header.Set("Authorization", auth) + } + + return r + } + + idp := targetIn(cognitoIDPTargetPrefix, cognitoIDPPublicOps) + ident := targetIn(cognitoIdentityTargetPrefix, cognitoIdentityPublicOps) + + cases := []struct { + name string + match func(*http.Request, []byte) bool + r *http.Request + want bool + }{ + {"idp public target", idp, mk(http.MethodPost, idpTarget+"GetUser", "", "/", ""), true}, + {"idp private target", idp, mk(http.MethodPost, idpTarget+"AdminGetUser", "", "/", ""), false}, + {"idp op under the identity prefix", idp, mk(http.MethodPost, identTgt+"InitiateAuth", "", "/", ""), false}, + {"identity public target", ident, mk(http.MethodPost, identTgt+"GetCredentialsForIdentity", "", "/", ""), true}, + {"identity private target", ident, mk(http.MethodPost, identTgt+"DescribeIdentity", "", "/", ""), false}, + {"jwks", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/eu-west-2_Ab12/.well-known/jwks.json", ""), true}, + {"openid-configuration", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/us-east-1_x/.well-known/openid-configuration", ""), true}, + {"well-known on a bucket-shaped id", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/mybucket/.well-known/jwks.json", ""), false}, + {"well-known POST", isUserPoolWellKnown, mk(http.MethodPost, "", "", "/us-east-1_x/.well-known/jwks.json", ""), false}, + {"hosted domain host", isHostedUI, mk(http.MethodGet, "", "d.auth.us-east-1.amazoncognito.com", "/oauth2/token", ""), true}, + {"hosted domain localhost", isHostedUI, mk(http.MethodPost, "", "d.auth.localhost:4566", "/oauth2/token", ""), true}, + {"hosted path fallback", isHostedUI, mk(http.MethodPost, "", "", "/_cognito/d/oauth2/token", ""), true}, + {"other amazoncognito host", isHostedUI, mk(http.MethodGet, "", "cognito-idp.us-east-1.amazoncognito.com", "/", ""), false}, + {"execute-api host", isExecuteAPI, mk(http.MethodGet, "", "a.execute-api.us-east-1.amazonaws.com", "/p", ""), true}, + {"execute-api path", isExecuteAPI, mk(http.MethodGet, "", "", "/restapis/a/s/_user_request_/p", ""), true}, + {"restapis control plane", isExecuteAPI, mk(http.MethodGet, "", "", "/restapis/a/resources/r", ""), false}, + {"appsync host unsigned", isUnsignedGraphQL, mk(http.MethodPost, "", "a.appsync-api.us-east-1.amazonaws.com", "/graphql", ""), true}, + {"appsync path unsigned", isUnsignedGraphQL, mk(http.MethodPost, "", "", "/graphql", ""), true}, + {"appsync signed (IAM auth mode)", isUnsignedGraphQL, mk(http.MethodPost, "", "", "/graphql", "AWS4-HMAC-SHA256 Credential=x"), false}, + {"graphql-prefixed bucket", isUnsignedGraphQL, mk(http.MethodGet, "", "", "/graphqlbucket/k", ""), false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := tc.match(tc.r, nil); got != tc.want { + t.Fatalf("match = %v, want %v", got, tc.want) + } + }) + } + + if rt := publicRouteFor(mk(http.MethodGet, "", "", "/", ""), nil); rt != nil { + t.Fatalf("plain GET / matched a public route") + } + + form := func() *http.Request { + r := httptest.NewRequest(http.MethodPost, "http://localhost/", nil) + r.Header.Set("Content-Type", formCT) + + return r + } + + if !isPublicSTSAction(form(), []byte("Action=AssumeRoleWithWebIdentity")) { + t.Fatalf("form Action AssumeRoleWithWebIdentity not recognized as public STS") + } + + if isPublicSTSAction(form(), []byte("Action=GetCallerIdentity")) { + t.Fatalf("GetCallerIdentity recognized as public") + } + + if !isPublicSTSAction(httptest.NewRequest(http.MethodGet, "http://localhost/?Action=AssumeRoleWithSAML", nil), nil) { + t.Fatalf("query-string Action AssumeRoleWithSAML not recognized as public STS") + } +} + +// TestAuthzSkipsPublicOps: a signed caller whose IAM policy allows only +// DynamoDB is still served a public Cognito operation (IAM never governs it), +// while a private Cognito operation is denied by the authorization gate. +func TestAuthzSkipsPublicOps(t *testing.T) { + ts, cloud := enforcedServer(t) + ctx := context.Background() + + if _, err := cloud.IAM.CreateUser(ctx, iamdriver.UserConfig{Name: "dynonly"}); err != nil { + t.Fatalf("CreateUser: %v", err) + } + + doc := `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"dynamodb:*","Resource":"*"}]}` + + pol, err := cloud.IAM.CreatePolicy(ctx, iamdriver.PolicyConfig{Name: "dynonly", PolicyDocument: doc}) + if err != nil { + t.Fatalf("CreatePolicy: %v", err) + } + + if err := cloud.IAM.AttachUserPolicy(ctx, "dynonly", pol.ARN); err != nil { + t.Fatalf("AttachUserPolicy: %v", err) + } + + ak, err := cloud.IAM.CreateAccessKey(ctx, iamdriver.AccessKeyConfig{UserName: "dynonly"}) + if err != nil { + t.Fatalf("CreateAccessKey: %v", err) + } + + creds := aws.Credentials{AccessKeyID: ak.AccessKeyID, SecretAccessKey: ak.SecretAccessKey} + + send := func(op string) (int, string) { + body := `{}` + + req, _ := http.NewRequestWithContext(ctx, http.MethodPost, ts.URL+"/", strings.NewReader(body)) + req.Header.Set("X-Amz-Target", idpTarget+op) + req.Header.Set("Content-Type", amzJSON11) + + sum := sha256.Sum256([]byte(body)) + if err := v4.NewSigner().SignHTTP(ctx, creds, req, hex.EncodeToString(sum[:]), "cognito-idp", "us-east-1", time.Now()); err != nil { + t.Fatalf("sign: %v", err) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("do: %v", err) + } + defer resp.Body.Close() + + b, _ := io.ReadAll(resp.Body) + + var e struct { + Type string `json:"__type"` + } + + _ = json.Unmarshal(b, &e) + + return resp.StatusCode, e.Type + } + + if status, typ := send("InitiateAuth"); status == http.StatusForbidden || typ == "AccessDeniedException" { + t.Fatalf("public InitiateAuth denied: %d %s", status, typ) + } + + if status, typ := send("ListUserPools"); status != http.StatusForbidden || typ != "AccessDeniedException" { + t.Fatalf("private ListUserPools: %d %s, want 403 AccessDeniedException", status, typ) + } +} + +// TestSDKAnonymousSTSCallPassesEnforcedGate drives the real STS client, which +// sends AssumeRoleWithWebIdentity unsigned because its model marks it noAuth. +func TestSDKAnonymousSTSCallPassesEnforcedGate(t *testing.T) { + ts, _ := enforcedServer(t) + + client := awssts.New(awssts.Options{ + Region: "us-east-1", + BaseEndpoint: aws.String(ts.URL), + Credentials: aws.AnonymousCredentials{}, + }) + + out, err := client.AssumeRoleWithWebIdentity(context.Background(), &awssts.AssumeRoleWithWebIdentityInput{ + RoleArn: aws.String("arn:aws:iam::123456789012:role/web"), + RoleSessionName: aws.String("s"), + WebIdentityToken: aws.String("header.payload.sig"), + }) + if err != nil { + t.Fatalf("AssumeRoleWithWebIdentity unsigned under --enforce-auth: %v", err) + } + + if out.Credentials == nil || aws.ToString(out.Credentials.AccessKeyId) == "" { + t.Fatalf("no credentials returned") + } + + if _, err := client.GetCallerIdentity(context.Background(), &awssts.GetCallerIdentityInput{}); err == nil || + !strings.Contains(err.Error(), missingTok) { + t.Fatalf("unsigned GetCallerIdentity: err = %v, want %s", err, missingTok) + } +} diff --git a/server/server.go b/server/server.go index 388079b3b..8465d1717 100644 --- a/server/server.go +++ b/server/server.go @@ -71,17 +71,30 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { } } - for _, h := range s.handlers { - if h.Matches(r) { - h.ServeHTTP(w, r) - - if s.observer != nil { - s.observer(r) - } + if h := s.Match(r); h != nil { + h.ServeHTTP(w, r) - return + if s.observer != nil { + s.observer(r) } + + return } http.Error(w, "no handler registered for this request", http.StatusNotImplemented) } + +// Match returns the handler that would serve r (the first registered handler +// whose Matches returns true), or nil when none would. A pre-dispatch hook uses +// it to bind a decision to the handler that actually runs. Matches may read the +// request body (form parsing), so a caller that needs the body afterwards must +// buffer and restore it. +func (s *Server) Match(r *http.Request) Handler { + for _, h := range s.handlers { + if h.Matches(r) { + return h + } + } + + return nil +} From 5d40947ca7a91cdf7b09abf55b5485c21c86fe35 Mon Sep 17 00:00:00 2001 From: Nitin Kumar Date: Sun, 27 Sep 2026 21:29:37 +0530 Subject: [PATCH 2/2] fix(auth): bind public exemptions to the serving handler and authorize STS sessions --- go.mod | 1 - go.sum | 2 - internal/jwtsign/jwtsign.go | 10 +- internal/jwtsign/jwtsign_test.go | 13 +- server/aws/apigateway/handler.go | 28 +- server/aws/authbypass_test.go | 130 +++++++++ server/aws/authgate.go | 51 ++-- server/aws/authgate_temp_test.go | 14 +- server/aws/authzgate.go | 10 +- server/aws/cognito/public.go | 67 +++++ server/aws/cognito/public_test.go | 100 +++++++ server/aws/publicauth.go | 272 ++---------------- server/aws/publicauth_test.go | 442 +++++++++++++----------------- server/aws/sts/operations.go | 48 +++- server/aws/sts/sessions.go | 25 +- server/aws/sts/sessions_test.go | 4 +- server/server.go | 10 + 17 files changed, 665 insertions(+), 562 deletions(-) create mode 100644 server/aws/authbypass_test.go create mode 100644 server/aws/cognito/public.go create mode 100644 server/aws/cognito/public_test.go diff --git a/go.mod b/go.mod index 273b190d4..f72168235 100644 --- a/go.mod +++ b/go.mod @@ -78,7 +78,6 @@ require ( github.com/aws/aws-sdk-go-v2/service/cloudwatch v1.56.2 github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0 github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0 - github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0 github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0 github.com/aws/aws-sdk-go-v2/service/configservice v1.68.4 github.com/aws/aws-sdk-go-v2/service/costexplorer v1.63.6 diff --git a/go.sum b/go.sum index 7434a7237..db65c75d3 100644 --- a/go.sum +++ b/go.sum @@ -210,8 +210,6 @@ github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0 h1:5W/KOwsnZrdi7RD97 github.com/aws/aws-sdk-go-v2/service/cloudwatchlogs v1.79.0/go.mod h1:h1Iw2nkdpmAUJaa89RvX3cg/HGLgdSkCWpMNgKvBSHA= github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0 h1:zssdCGwtqMsIrkYFCWu4yhdHHbp+OxbHoFSWp0QDsZo= github.com/aws/aws-sdk-go-v2/service/codeartifact v1.45.0/go.mod h1:b8QnRsunCLqoz5Elpypig/tbI23E2N0betuvS6wwbsE= -github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0 h1:Wp96KzoBhY4T8nIU/PK/Lf6mCD9nelkXuY6SvvKvdgY= -github.com/aws/aws-sdk-go-v2/service/cognitoidentity v1.41.0/go.mod h1:Nqfwixdr/+xgqWGHVkOcVytNVy3VCxFCzWr+8XxEn10= github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0 h1:3Vje2gVkUDNSksJ8NXLcLCSg5m/YtsTqSNfDupy3qeI= github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider v1.53.0/go.mod h1:ygltZT++6Wn2uG4+tqE0NW1MkdEtb5W2O/CFc0xJX/g= github.com/aws/aws-sdk-go-v2/service/configservice v1.68.4 h1:L37DoV4JOQFqucVUeYaJrrhIc1995EPra+8UtGRCsgo= diff --git a/internal/jwtsign/jwtsign.go b/internal/jwtsign/jwtsign.go index 8d4dd933a..34c4c40cd 100644 --- a/internal/jwtsign/jwtsign.go +++ b/internal/jwtsign/jwtsign.go @@ -32,6 +32,10 @@ const algRS256 = "RS256" const ( keyBits = 2048 jwtSegments = 3 + + // iatLeewaySeconds tolerates an issuer clock slightly ahead of the + // verifier: a token whose iat is at most this far in the future passes. + iatLeewaySeconds = 60 ) // Errors returned by Verify. ErrExpired is kept apart from ErrInvalidToken so a @@ -117,8 +121,8 @@ func Sign(k *Key, claims map[string]any) (string, error) { // Verify checks token against keys and the clock and returns its claims. It // requires alg RS256, a kid present in keys, a valid signature, and an exp -// claim. It rejects a token at or after exp, before nbf, or with iat in the -// future. Numeric claims come back as json.Number. +// claim. It rejects a token at or after exp, before nbf, or with iat more than +// 60 seconds in the future. Numeric claims come back as json.Number. func Verify(token string, keys []*Key, clock config.Clock) (map[string]any, error) { parts := strings.Split(token, ".") if len(parts) != jwtSegments { @@ -177,7 +181,7 @@ func checkTimes(claims map[string]any, now time.Time) error { return ErrNotYetValid } - if iat, ok := numericDate(claims, "iat"); ok && unix < iat { + if iat, ok := numericDate(claims, "iat"); ok && unix+iatLeewaySeconds < iat { return ErrNotYetValid } diff --git a/internal/jwtsign/jwtsign_test.go b/internal/jwtsign/jwtsign_test.go index 866ec991d..2cd7127cd 100644 --- a/internal/jwtsign/jwtsign_test.go +++ b/internal/jwtsign/jwtsign_test.go @@ -153,12 +153,21 @@ func TestVerifyNotBeforeAndIssuedInFuture(t *testing.T) { } iat := claimsAt(epoch, time.Hour) - iat["iat"] = future.Unix() delete(iat, "nbf") + + iat["iat"] = epoch.Add(2 * time.Minute).Unix() tok, _ = Sign(k, iat) if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); !errors.Is(err, ErrNotYetValid) { - t.Fatalf("iat in future: err = %v, want ErrNotYetValid", err) + t.Fatalf("iat 2m in future: err = %v, want ErrNotYetValid", err) + } + + // A small issuer clock skew is tolerated. + iat["iat"] = epoch.Add(time.Minute).Unix() + tok, _ = Sign(k, iat) + + if _, err := Verify(tok, []*Key{k}, config.NewFakeClock(epoch)); err != nil { + t.Fatalf("iat within the 60s leeway: %v", err) } } diff --git a/server/aws/apigateway/handler.go b/server/aws/apigateway/handler.go index fa80a0e8d..c8c5e3fcc 100644 --- a/server/aws/apigateway/handler.go +++ b/server/aws/apigateway/handler.go @@ -72,17 +72,31 @@ func (*Handler) Matches(r *http.Request) bool { // ServeHTTP dispatches to the data plane (execute-api host or a _user_request_ // path) or the control plane. func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { - if strings.Contains(r.Host, executeAPIMarker) { + switch { + case isHostDataPlane(r): h.serveHostDataPlane(w, r) - return - } - - if strings.Contains(r.URL.Path, "/"+userRequestMark) { + case isPathDataPlane(r): h.servePathDataPlane(w, r) - return + default: + h.serveControlPlane(w, r) } +} + +// PublicRequest reports whether r is an API invocation (either data-plane +// form), which real API Gateway accepts without SigV4 unless the method uses +// AWS_IAM authorization. Control-plane requests always need SigV4. +func (*Handler) PublicRequest(r *http.Request) bool { + return isHostDataPlane(r) || isPathDataPlane(r) +} + +// isHostDataPlane reports a request addressed to an execute-api host. +func isHostDataPlane(r *http.Request) bool { + return strings.Contains(r.Host, executeAPIMarker) +} - h.serveControlPlane(w, r) +// isPathDataPlane reports a /restapis/{apiId}/{stage}/_user_request_/ invoke. +func isPathDataPlane(r *http.Request) bool { + return strings.Contains(r.URL.Path, "/"+userRequestMark) } // serveControlPlane routes the restJson1 management API under /restapis. diff --git a/server/aws/authbypass_test.go b/server/aws/authbypass_test.go new file mode 100644 index 000000000..1173a68d7 --- /dev/null +++ b/server/aws/authbypass_test.go @@ -0,0 +1,130 @@ +package aws + +import ( + "context" + "net/http" + "strings" + "testing" + + "github.com/aws/aws-sdk-go-v2/aws" + awssts "github.com/aws/aws-sdk-go-v2/service/sts" +) + +// TestEnforcedGateBypassAttempts replays every known way of getting an unsigned +// request past --enforce-auth. Each must be rejected by the gate itself with +// 403 MissingAuthenticationToken, so it never reaches a handler. +func TestEnforcedGateBypassAttempts(t *testing.T) { + ts, _ := enforcedServer(t) + + cognitoOn := func(host, path, target string) rawReq { + rq := jsonRPC(idpTarget+target, `{"PoolName":"p"}`) + rq.host, rq.path = host, path + + return rq + } + + form := func(path, body string) rawReq { + rq := queryForm(body) + rq.path = path + + return rq + } + + cases := []struct { + name string + req rawReq + }{ + // STS noAuth ops stay gated until web-identity and SAML validation land. + {"sts AssumeRoleWithWebIdentity for a missing role", form("/", "Action=AssumeRoleWithWebIdentity&Version=2011-06-15"+ + "&RoleArn=arn%3Aaws%3Aiam%3A%3A000000000000%3Arole%2Fnonexistent&RoleSessionName=s&WebIdentityToken=junk")}, + {"sts AssumeRoleWithSAML", form("/", "Action=AssumeRoleWithSAML&Version=2011-06-15&RoleArn=x&PrincipalArn=y&SAMLAssertion=eA%3D%3D")}, + {"sts GetCallerIdentity", form("/", "Action=GetCallerIdentity&Version=2011-06-15")}, + + // Parser disagreement: a body that does not parse plus a public Action in + // the query string, and duplicated Actions. + {"bad body, public query Action, GetSessionToken", form("/?Action=AssumeRoleWithWebIdentity", + "Action=GetSessionToken&x=%zz")}, + {"bad body, public query Action, AssumeRole", form("/?Action=AssumeRoleWithWebIdentity", + "Action=AssumeRole&RoleArn=x&RoleSessionName=s&x=%zz")}, + {"duplicate Action, public first", form("/", "Action=AssumeRoleWithWebIdentity&Action=GetSessionToken")}, + {"duplicate Action, public last", form("/", "Action=GetSessionToken&Action=AssumeRoleWithWebIdentity")}, + {"public body Action, private query Action", form("/?Action=GetSessionToken", "Action=AssumeRoleWithWebIdentity")}, + {"bad query string", form("/?Action=AssumeRoleWithWebIdentity&x=%zz", "Action=AssumeRoleWithWebIdentity")}, + {"lower-case public Action", form("/", "Action=assumerolewithwebidentity")}, + + // Cognito private ops reached through the hosted-UI and well-known + // markers, or with a public target on a non-JSON-RPC route. + {"hosted-ui host, CreateUserPool", cognitoOn("x.auth.localhost", "/", "CreateUserPool")}, + {"hosted-ui amazoncognito host, CreateUserPool", cognitoOn("x.auth.us-east-1.amazoncognito.com", "/", "CreateUserPool")}, + {"/_cognito path, CreateUserPool", cognitoOn("", "/_cognito/x", "CreateUserPool")}, + {"/_cognito path, public target", cognitoOn("", "/_cognito/x", "InitiateAuth")}, + {"well-known path, ListUserPools", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json", + header: map[string]string{"X-Amz-Target": idpTarget + "ListUserPools"}}}, + {"jwks GET before Cognito serves it", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json"}}, + {"cognito CreateUserPool", cognitoOn("", "/", "CreateUserPool")}, + {"cognito AdminInitiateAuth", cognitoOn("", "/", "AdminInitiateAuth")}, + {"cognito lower-case op", cognitoOn("", "/", "initiateAuth")}, + {"cognito lower-case prefix", rawReq{method: http.MethodPost, path: "/", body: `{}`, + header: map[string]string{"X-Amz-Target": "awscognitoidentityproviderservice.InitiateAuth", "Content-Type": amzJSON11}}}, + {"cognito-identity GetId (no handler served)", jsonRPC(identTgt+"GetId", `{}`)}, + + // AppSync control plane behind a data-plane marker. + {"appsync-api host, CreateGraphqlApi", rawReq{method: http.MethodPost, path: "/v1/apis", + host: "x.appsync-api.us-east-1.amazonaws.com", body: `{"name":"a","authenticationType":"API_KEY"}`, + header: map[string]string{"Content-Type": "application/json"}}}, + {"appsync /graphql before AppSync serves it", rawReq{method: http.MethodPost, path: "/graphql", body: `{}`}}, + + // execute-api markers on requests that dispatch to another service. + {"execute-api host, DynamoDB target", func() rawReq { + rq := jsonRPC("DynamoDB_20120810.ListTables", `{}`) + rq.host = execHost + + return rq + }()}, + {"execute-api host, lambda path", rawReq{method: http.MethodGet, path: lambdaPath, host: execHost}}, + {"execute-api host, query Action", func() rawReq { + rq := form("/", "Action=ListUsers&Version=2010-05-08") + rq.host = execHost + + return rq + }()}, + + // Plain private operations. + {"ec2 DescribeInstances", form("/", "Action=DescribeInstances&Version=2016-11-15")}, + {"iam CreateUser", form("/", "Action=CreateUser&Version=2010-05-08&UserName=x")}, + {"dynamodb ListTables", jsonRPC("DynamoDB_20120810.ListTables", `{}`)}, + {"s3 ListBuckets", rawReq{method: http.MethodGet, path: "/"}}, + {"public Action on a lambda path", rawReq{method: http.MethodGet, path: lambdaPath + "?Action=AssumeRoleWithWebIdentity"}}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + status, body := doRaw(t, ts, tc.req) + if status != http.StatusForbidden || !strings.Contains(body, missingTok) { + t.Fatalf("status %d, body %s; want 403 %s", status, body, missingTok) + } + }) + } +} + +// TestSDKAnonymousSTSCallIsRejected drives the real STS client, which sends +// AssumeRoleWithWebIdentity unsigned (its model marks it noAuth). The gate +// rejects it until the token is validated against a registered provider. +func TestSDKAnonymousSTSCallIsRejected(t *testing.T) { + ts, _ := enforcedServer(t) + + client := awssts.New(awssts.Options{ + Region: "us-east-1", + BaseEndpoint: aws.String(ts.URL), + Credentials: aws.AnonymousCredentials{}, + }) + + _, err := client.AssumeRoleWithWebIdentity(context.Background(), &awssts.AssumeRoleWithWebIdentityInput{ + RoleArn: aws.String("arn:aws:iam::123456789012:role/nonexistent"), + RoleSessionName: aws.String("s"), + WebIdentityToken: aws.String("junk"), + }) + if err == nil || !strings.Contains(err.Error(), missingTok) { + t.Fatalf("unsigned AssumeRoleWithWebIdentity: err = %v, want %s", err, missingTok) + } +} diff --git a/server/aws/authgate.go b/server/aws/authgate.go index 150f3ef1f..9ed13192e 100644 --- a/server/aws/authgate.go +++ b/server/aws/authgate.go @@ -73,22 +73,21 @@ func newAuthGate( // Temporary STS credentials are verified against the secret STS recorded // when it minted them (resolved from the session store), so a forged ASIA - // credential and an expired session are both rejected. - if strings.HasPrefix(akid, tempCredentialPrefix) { - principal, aerr := verifyTempCredential(r, body, akid, accountID, sessions, clock) - - restore() - - if aerr != nil { - writeAuthError(w, r, aerr) - return r, false - } + // credential and an expired session are both rejected. The session is + // then authorized as its owner: a role session strictly against the + // role's policies, a GetSessionToken session as the user that minted it. + var ( + principal authctx.Principal + roleSession bool + aerr *sigv4.AuthError + ) - return withPrincipal(r, principal), true + if strings.HasPrefix(akid, tempCredentialPrefix) { + principal, roleSession, aerr = verifyTempCredential(r, body, akid, accountID, sessions, clock) + } else { + principal, aerr = sigv4.Verify(r, body, resolverLookup(r, resolver), clock) } - principal, aerr := sigv4.Verify(r, body, resolverLookup(r, resolver), clock) - restore() if aerr != nil { @@ -96,7 +95,7 @@ func newAuthGate( return r, false } - if !authorize(w, r, principal, iamDriver, body, accountID) { + if !authorize(w, r, principal, iamDriver, body, accountID, roleSession) { return r, false } @@ -108,10 +107,12 @@ func newAuthGate( // resolves the secret STS recorded for the presented access key id, rejects an // unknown key (InvalidClientTokenId) or an expired session (ExpiredToken), then // SigV4-verifies the signature against that secret. When no session store is -// wired the credential is unverifiable, so it fails closed. +// wired the credential is unverifiable, so it fails closed. The principal is +// the session's owner (see stssrv.SessionOwner), and roleSession reports +// whether it is a role session. func verifyTempCredential( r *http.Request, body []byte, akid, accountID string, sessions *stssrv.SessionStore, clock config.Clock, -) (authctx.Principal, *sigv4.AuthError) { +) (principal authctx.Principal, roleSession bool, aerr *sigv4.AuthError) { invalid := &sigv4.AuthError{ Code: "InvalidClientTokenId", Message: "The security token included in the request is invalid.", @@ -119,16 +120,16 @@ func verifyTempCredential( } if sessions == nil { - return authctx.Principal{}, invalid + return authctx.Principal{}, false, invalid } sess, ok := sessions.Lookup(akid) if !ok { - return authctx.Principal{}, invalid + return authctx.Principal{}, false, invalid } if clock.Now().UTC().After(sess.Expiration) { - return authctx.Principal{}, &sigv4.AuthError{ + return authctx.Principal{}, false, &sigv4.AuthError{ Code: "ExpiredToken", Message: "The security token included in the request is expired.", HTTPStatus: http.StatusForbidden, @@ -136,10 +137,18 @@ func verifyTempCredential( } lookup := func(id string) (string, authctx.Principal, bool) { - return sess.SecretAccessKey, authctx.Principal{AccessKeyID: id, AccountID: accountID}, true + return sess.SecretAccessKey, authctx.Principal{ + AccessKeyID: id, + AccountID: accountID, + UserName: sess.Owner.PolicyEntity, + ARN: sess.Owner.ARN, + UserID: sess.Owner.UserID, + }, true } - return sigv4.Verify(r, body, lookup, clock) + principal, aerr = sigv4.Verify(r, body, lookup, clock) + + return principal, sess.Owner.Role, aerr } // resolverLookup adapts the IAM access-key resolver to sigv4.LookupFunc, diff --git a/server/aws/authgate_temp_test.go b/server/aws/authgate_temp_test.go index 59861b51d..a2ead73c3 100644 --- a/server/aws/authgate_temp_test.go +++ b/server/aws/authgate_temp_test.go @@ -49,7 +49,7 @@ func TestVerifyTempCredential(t *testing.T) { clock := config.NewFakeClock(now) store := stssrv.NewSessionStore(clock) - issued, err := store.Mint(time.Hour) // Expiration = now + 1h + issued, err := store.Mint(time.Hour, stssrv.SessionOwner{}) // Expiration = now + 1h if err != nil { t.Fatalf("Mint: %v", err) } @@ -68,7 +68,7 @@ func TestVerifyTempCredential(t *testing.T) { r := newReq() signTemp(t, r, issued.AccessKeyID, issued.SecretAccessKey, issued.SessionToken, now) - p, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, store, clock) + p, _, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, store, clock) if aerr != nil { t.Fatalf("valid temp credential rejected: %v", aerr) } @@ -81,7 +81,7 @@ func TestVerifyTempCredential(t *testing.T) { r := newReq() signTemp(t, r, issued.AccessKeyID, "forged-secret-000000000000000000000000", issued.SessionToken, now) - _, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, store, clock) + _, _, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, store, clock) if aerr == nil || aerr.Code != "SignatureDoesNotMatch" { t.Fatalf("want SignatureDoesNotMatch, got %v", aerr) } @@ -91,7 +91,7 @@ func TestVerifyTempCredential(t *testing.T) { r := newReq() signTemp(t, r, "ASIAUNKNOWN0000000000", issued.SecretAccessKey, issued.SessionToken, now) - _, aerr := verifyTempCredential(r, nil, "ASIAUNKNOWN0000000000", tempTestAccount, store, clock) + _, _, aerr := verifyTempCredential(r, nil, "ASIAUNKNOWN0000000000", tempTestAccount, store, clock) if aerr == nil || aerr.Code != "InvalidClientTokenId" { t.Fatalf("want InvalidClientTokenId, got %v", aerr) } @@ -101,7 +101,7 @@ func TestVerifyTempCredential(t *testing.T) { r := newReq() signTemp(t, r, issued.AccessKeyID, issued.SecretAccessKey, issued.SessionToken, now) - _, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, nil, clock) + _, _, aerr := verifyTempCredential(r, nil, issued.AccessKeyID, tempTestAccount, nil, clock) if aerr == nil || aerr.Code != "InvalidClientTokenId" { t.Fatalf("want InvalidClientTokenId (fail closed), got %v", aerr) } @@ -110,7 +110,7 @@ func TestVerifyTempCredential(t *testing.T) { t.Run("expired-session", func(t *testing.T) { shortClock := config.NewFakeClock(now) shortStore := stssrv.NewSessionStore(shortClock) - short, err := shortStore.Mint(15 * time.Minute) // Expiration = now + 15m + short, err := shortStore.Mint(15*time.Minute, stssrv.SessionOwner{}) // Expiration = now + 15m if err != nil { t.Fatalf("Mint: %v", err) } @@ -119,7 +119,7 @@ func TestVerifyTempCredential(t *testing.T) { signTemp(t, r, short.AccessKeyID, short.SecretAccessKey, short.SessionToken, now) shortClock.Advance(time.Hour) // now past expiration - _, aerr := verifyTempCredential(r, nil, short.AccessKeyID, tempTestAccount, shortStore, shortClock) + _, _, aerr := verifyTempCredential(r, nil, short.AccessKeyID, tempTestAccount, shortStore, shortClock) if aerr == nil || aerr.Code != "ExpiredToken" { t.Fatalf("want ExpiredToken, got %v", aerr) } diff --git a/server/aws/authzgate.go b/server/aws/authzgate.go index 8abcd69de..998bd594e 100644 --- a/server/aws/authzgate.go +++ b/server/aws/authzgate.go @@ -82,6 +82,11 @@ var jsonRPCServiceByTarget = map[string]string{ // policies defined, gates the action through CheckPermission. It returns // proceed=false only when the action is denied, having already written the 403. // +// strict is set for an STS role session. Its principal is the role, which is +// evaluated on its policies alone: the root/admin and no-policies bootstrap +// shortcuts that apply to IAM users do not apply, so a role with no allowing +// policy, or a role that does not exist, is denied. +// // Authorization is enforced for the JSON-RPC protocol, where the X-Amz-Target // header both routes the request and names the service, so the service the gate // authorizes is the one the handler runs. The query and REST protocols are @@ -94,6 +99,7 @@ var jsonRPCServiceByTarget = map[string]string{ // action+resource authorization bound to the routed operation is a follow-up. func authorize( w http.ResponseWriter, r *http.Request, p authctx.Principal, iamDriver iamdriver.IAM, body []byte, accountID string, + strict bool, ) bool { service, action, decision := deriveAction(r) @@ -107,11 +113,11 @@ func authorize( return false } - if isAdminPrincipal(p) { + if !strict && isAdminPrincipal(p) { return true // account root / bootstrap admin identity: full access. } - if !principalHasPolicies(r, p, iamDriver) { + if !strict && !principalHasPolicies(r, p, iamDriver) { return true // no policies defined: unrestricted (dev-friendly bootstrap). } diff --git a/server/aws/cognito/public.go b/server/aws/cognito/public.go new file mode 100644 index 000000000..d5b5b817b --- /dev/null +++ b/server/aws/cognito/public.go @@ -0,0 +1,67 @@ +package cognito + +import ( + "net/http" + "strings" +) + +// publicOps are the cognito-idp operations whose Smithy model carries +// "auth": ["smithy.api#noAuth"] (botocore: "authtype": "none"), taken from the +// botocore model (awscli 2.31.19, botocore/data/cognito-idp/2016-04-18/service-2.json). +// They authenticate with an access token, a session, or a client secret hash +// inside the request, never with SigV4. TestPublicOpsMatchSDKModel re-derives +// the set from the generated aws-sdk-go-v2 auth resolver and fails on drift. +// +//nolint:gochecknoglobals // static protocol lookup table +var publicOps = map[string]struct{}{ + "AssociateSoftwareToken": {}, + "ChangePassword": {}, + "CompleteWebAuthnRegistration": {}, + "ConfirmDevice": {}, + "ConfirmForgotPassword": {}, + "ConfirmSignUp": {}, + "DeleteUser": {}, + "DeleteUserAttributes": {}, + "DeleteWebAuthnCredential": {}, + "ForgetDevice": {}, + "ForgotPassword": {}, + "GetDevice": {}, + "GetTokensFromRefreshToken": {}, + "GetUser": {}, + "GetUserAttributeVerificationCode": {}, + "GetUserAuthFactors": {}, + "GlobalSignOut": {}, + "InitiateAuth": {}, + "ListDevices": {}, + "ListWebAuthnCredentials": {}, + "ResendConfirmationCode": {}, + "RespondToAuthChallenge": {}, + "RevokeToken": {}, + "SetUserMFAPreference": {}, + "SetUserSettings": {}, + "SignUp": {}, + "StartWebAuthnRegistration": {}, + "UpdateAuthEventFeedback": {}, + "UpdateDeviceStatus": {}, + "UpdateUserAttributes": {}, + "VerifySoftwareToken": {}, + "VerifyUserAttribute": {}, +} + +// PublicRequest reports whether r is a noAuth cognito-idp operation on the +// normal JSON-RPC route (POST / with the operation's X-Amz-Target). Every other +// request this handler claims, whatever its Host or path, needs SigV4. +func (*Handler) PublicRequest(r *http.Request) bool { + if r.Method != http.MethodPost || r.URL.Path != "/" { + return false + } + + op, ok := strings.CutPrefix(r.Header.Get("X-Amz-Target"), targetPrefix) + if !ok { + return false + } + + _, public := publicOps[op] + + return public +} diff --git a/server/aws/cognito/public_test.go b/server/aws/cognito/public_test.go new file mode 100644 index 000000000..5cfca8102 --- /dev/null +++ b/server/aws/cognito/public_test.go @@ -0,0 +1,100 @@ +package cognito + +import ( + "context" + "net/http" + "net/http/httptest" + "reflect" + "testing" + + "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider" + smithyauth "github.com/aws/smithy-go/auth" +) + +// newerModelNoAuth lists operations the botocore model the table was taken from +// marks noAuth while the aws-sdk-go-v2 version pinned in go.mod still signs +// them. Drop an entry once the pinned SDK catches up. +// +//nolint:gochecknoglobals // test fixture +var newerModelNoAuth = map[string]struct{}{ + "GetTokensFromRefreshToken": {}, +} + +// TestPublicOpsMatchSDKModel asks the SDK client's default auth scheme +// resolver, which is generated from the Smithy model, which operations resolve +// to the anonymous scheme (smithy.api#noAuth), and requires publicOps to be +// exactly that set. Operations are the client's exported methods other than +// Options. +func TestPublicOpsMatchSDKModel(t *testing.T) { + client := cognitoidentityprovider.New(cognitoidentityprovider.Options{Region: "us-east-1"}) + resolver := client.Options().AuthSchemeResolver + typ := reflect.TypeOf(client) + anon := 0 + + for i := range typ.NumMethod() { + op := typ.Method(i).Name + if op == "Options" { + continue + } + + opts, err := resolver.ResolveAuthSchemes(context.Background(), + &cognitoidentityprovider.AuthResolverParameters{Operation: op, Region: "us-east-1"}) + if err != nil { + t.Fatalf("%s: resolve: %v", op, err) + } + + isAnon := len(opts) > 0 && opts[0].SchemeID == smithyauth.SchemeIDAnonymous + if isAnon { + anon++ + } + + _, exempt := publicOps[op] + if _, newer := newerModelNoAuth[op]; newer && !isAnon && exempt { + continue + } + + if isAnon != exempt { + t.Errorf("%s: model noAuth=%v, public=%v", op, isAnon, exempt) + } + } + + if anon == 0 { + t.Fatalf("SDK model reports no anonymous operations; resolver probe is broken") + } +} + +func TestPublicRequestOnlyOnTheJSONRPCRoute(t *testing.T) { + cases := []struct { + name, method, host, path, target string + want bool + }{ + {"InitiateAuth", http.MethodPost, "", "/", targetPrefix + "InitiateAuth", true}, + {"SignUp", http.MethodPost, "", "/", targetPrefix + "SignUp", true}, + {"CreateUserPool", http.MethodPost, "", "/", targetPrefix + "CreateUserPool", false}, + {"AdminInitiateAuth", http.MethodPost, "", "/", targetPrefix + "AdminInitiateAuth", false}, + {"lower-case op", http.MethodPost, "", "/", targetPrefix + "initiateAuth", false}, + {"lower-case prefix", http.MethodPost, "", "/", "awscognitoidentityproviderservice.InitiateAuth", false}, + {"no target", http.MethodPost, "", "/", "", false}, + {"GET", http.MethodGet, "", "/", targetPrefix + "InitiateAuth", false}, + {"hosted-ui path", http.MethodPost, "", "/_cognito/x", targetPrefix + "InitiateAuth", false}, + {"well-known path", http.MethodGet, "", "/us-east-1_abc/.well-known/jwks.json", targetPrefix + "GetUser", false}, + {"hosted-ui host, private op", http.MethodPost, "x.auth.localhost", "/", targetPrefix + "CreateUserPool", false}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + r := httptest.NewRequest(tc.method, "http://localhost"+tc.path, nil) + if tc.host != "" { + r.Host = tc.host + } + + if tc.target != "" { + r.Header.Set("X-Amz-Target", tc.target) + } + + if got := (&Handler{}).PublicRequest(r); got != tc.want { + t.Fatalf("PublicRequest = %v, want %v", got, tc.want) + } + }) + } +} diff --git a/server/aws/publicauth.go b/server/aws/publicauth.go index b86cafb70..6e0eb4fbf 100644 --- a/server/aws/publicauth.go +++ b/server/aws/publicauth.go @@ -1,275 +1,53 @@ package aws import ( - "net" + "bytes" + "io" "net/http" "net/url" - "regexp" "strings" "github.com/stackshy/cloudemu/v2/server" - apigatewaysrv "github.com/stackshy/cloudemu/v2/server/aws/apigateway" - appsyncsrv "github.com/stackshy/cloudemu/v2/server/aws/appsync" - cognitosrv "github.com/stackshy/cloudemu/v2/server/aws/cognito" - stssrv "github.com/stackshy/cloudemu/v2/server/aws/sts" ) -// Some AWS operations are called without SigV4 credentials: a user signs in to -// Cognito before it has any AWS credentials, a web-identity token is traded for -// credentials, and API Gateway / AppSync endpoints are hit by browsers. The -// operation tables below are the operations whose Smithy model carries -// "auth": ["smithy.api#noAuth"] (botocore: "authtype": "none"), taken from the -// botocore service models (awscli 2.31.19, botocore/data//*/service-2.json). -// TestPublicOpsMatchSDKModels re-derives the same sets from the generated -// aws-sdk-go-v2 auth resolvers and fails when a table drifts from the model. -// Of the other services with noAuth operations, sso and sso-oidc are not served. - -// cognitoIDPPublicOps are the cognito-idp (AWSCognitoIdentityProviderService) -// noAuth operations. They authenticate with an access token, a session, or a -// client secret hash inside the request, never with SigV4. -// -//nolint:gochecknoglobals // static protocol lookup table -var cognitoIDPPublicOps = map[string]struct{}{ - "AssociateSoftwareToken": {}, - "ChangePassword": {}, - "CompleteWebAuthnRegistration": {}, - "ConfirmDevice": {}, - "ConfirmForgotPassword": {}, - "ConfirmSignUp": {}, - "DeleteUser": {}, - "DeleteUserAttributes": {}, - "DeleteWebAuthnCredential": {}, - "ForgetDevice": {}, - "ForgotPassword": {}, - "GetDevice": {}, - "GetTokensFromRefreshToken": {}, - "GetUser": {}, - "GetUserAttributeVerificationCode": {}, - "GetUserAuthFactors": {}, - "GlobalSignOut": {}, - "InitiateAuth": {}, - "ListDevices": {}, - "ListWebAuthnCredentials": {}, - "ResendConfirmationCode": {}, - "RespondToAuthChallenge": {}, - "RevokeToken": {}, - "SetUserMFAPreference": {}, - "SetUserSettings": {}, - "SignUp": {}, - "StartWebAuthnRegistration": {}, - "UpdateAuthEventFeedback": {}, - "UpdateDeviceStatus": {}, - "UpdateUserAttributes": {}, - "VerifySoftwareToken": {}, - "VerifyUserAttribute": {}, -} - -// cognitoIdentityPublicOps are the cognito-identity (AWSCognitoIdentityService) -// noAuth operations, which authenticate with the identity-pool logins map. -// -//nolint:gochecknoglobals // static protocol lookup table -var cognitoIdentityPublicOps = map[string]struct{}{ - "GetCredentialsForIdentity": {}, - "GetId": {}, - "GetOpenIdToken": {}, - "UnlinkIdentity": {}, -} - -// stsPublicActions are the STS noAuth actions: the caller presents a web -// identity token or a SAML assertion instead of AWS credentials. -// -//nolint:gochecknoglobals // static protocol lookup table -var stsPublicActions = map[string]struct{}{ - "AssumeRoleWithSAML": {}, - "AssumeRoleWithWebIdentity": {}, -} - -const ( - cognitoIDPTargetPrefix = "AWSCognitoIdentityProviderService." - cognitoIdentityTargetPrefix = "AWSCognitoIdentityService." - wellKnownSegment = "/.well-known/" - hostedUIPathPrefix = "/_cognito/" - restAPIsPrefix = "/restapis/" - userRequestSegment = "_user_request_" - userRequestSegmentIndex = 2 // {apiId}/{stage}/_user_request_/... - userRequestSplitParts = 4 // the three leading segments plus the rest - graphQLPath = "/graphql" - urlEncodedForm = "application/x-www-form-urlencoded" -) - -// userPoolIDPattern is the Cognito user pool id shape ({region}_{suffix}). The -// underscore makes it an illegal S3 bucket name, so a .well-known path under it -// never names a bucket. -var userPoolIDPattern = regexp.MustCompile(`^[a-z]{2}(-[a-z]+)+-\d_[0-9A-Za-z]+$`) - -// publicRoute is one family of requests AWS serves without SigV4. match decides -// from the request shape alone. owner reports whether the handler that would -// actually serve the request is the one that owns this public surface; nil -// means no served handler owns it yet. -type publicRoute struct { - match func(r *http.Request, body []byte) bool - owner func(server.Handler) bool -} - -// publicRoutes lists every public surface. A route whose owner handler does not -// claim the request yet (the Cognito JWKS and hosted UI, the AppSync GraphQL -// endpoint) stays gated. It opens once that handler starts claiming it. +// exemptPublic reports whether r may skip the SigV4 gate because it is an +// operation AWS serves without credentials (a noAuth operation). // -//nolint:gochecknoglobals // static route table -var publicRoutes = []publicRoute{ - {match: targetIn(cognitoIDPTargetPrefix, cognitoIDPPublicOps), owner: ownedBy[*cognitosrv.Handler]}, - {match: targetIn(cognitoIdentityTargetPrefix, cognitoIdentityPublicOps)}, - {match: isUserPoolWellKnown, owner: ownedBy[*cognitosrv.Handler]}, - {match: isHostedUI, owner: ownedBy[*cognitosrv.Handler]}, - {match: isPublicSTSAction, owner: ownedBy[*stssrv.Handler]}, - {match: isExecuteAPI, owner: ownedBy[*apigatewaysrv.Handler]}, - {match: isUnsignedGraphQL, owner: ownedBy[*appsyncsrv.Handler]}, -} - -// exemptPublic reports whether r may skip the SigV4 gate. The request must have -// the shape of a public operation, and the handler that would serve it (per -// match, the dispatcher's own first-match lookup) must be that operation's -// owner. Binding to the dispatch target stops a request from borrowing a public -// marker, such as an execute-api Host or a public Action, while being routed to -// a different service. When no handler would serve the request it is exempt too, -// since the dispatcher then answers 501 and touches no state. +// The decision belongs to the handler that dispatch will pick for this exact +// request: the gate runs the dispatcher's own first-match lookup (match) on a +// probe copy of r and asks that handler, through server.PublicRequester, +// whether it serves r as a public operation. A handler answers true only for +// the public routes it really serves, never for its private ones, so a +// request cannot borrow a public marker (a Host, a path, a target) while being +// served as something else. When no handler would serve r, it is not exempt. // -// Matches may parse the body, so the caller restores it afterwards. Exempt -// requests skip authorization as well: IAM does not govern noAuth operations. +// The probe gets fresh form state and the same body bytes that dispatch will +// read, so the lookup and the real dispatch see identical input. A form body +// or query string that does not parse fails closed, since handlers that parse +// forms could otherwise disagree about what the request is. func exemptPublic(r *http.Request, body []byte, match func(*http.Request) server.Handler) bool { - rt := publicRouteFor(r, body) - if rt == nil { - return false - } - - h := match(r) - if h == nil { - return true - } - - return rt.owner != nil && rt.owner(h) -} - -// publicRouteFor returns the public route r's shape matches, or nil. -func publicRouteFor(r *http.Request, body []byte) *publicRoute { - for i := range publicRoutes { - if publicRoutes[i].match(r, body) { - return &publicRoutes[i] - } - } - - return nil -} - -func ownedBy[T server.Handler](h server.Handler) bool { - _, ok := h.(T) - return ok -} - -// targetIn matches a JSON-RPC request whose X-Amz-Target is prefix+op for an -// op in ops. -func targetIn(prefix string, ops map[string]struct{}) func(*http.Request, []byte) bool { - return func(r *http.Request, _ []byte) bool { - op, ok := strings.CutPrefix(r.Header.Get("X-Amz-Target"), prefix) - if !ok { - return false - } - - _, public := ops[op] - - return public - } -} - -// isPublicSTSAction matches a query-protocol request whose Action (form body -// first, then the query string, the order net/http's ParseForm gives) is a -// public STS action. -func isPublicSTSAction(r *http.Request, body []byte) bool { - if r.Header.Get("X-Amz-Target") != "" { + if _, err := url.ParseQuery(r.URL.RawQuery); err != nil { return false } - action := "" - if strings.HasPrefix(r.Header.Get("Content-Type"), urlEncodedForm) { - if form, err := url.ParseQuery(string(body)); err == nil { - action = form.Get("Action") + if _, err := url.ParseQuery(string(body)); err != nil { + return false } } - if action == "" { - action = r.URL.Query().Get("Action") - } - - _, public := stsPublicActions[action] - - return public -} - -// isUserPoolWellKnown matches GET /{userPoolId}/.well-known/... (jwks.json and -// openid-configuration). -func isUserPoolWellKnown(r *http.Request, _ []byte) bool { - if r.Method != http.MethodGet && r.Method != http.MethodHead { - return false - } + probe := r.Clone(r.Context()) + probe.Form, probe.PostForm, probe.MultipartForm = nil, nil, nil + probe.Body = io.NopCloser(bytes.NewReader(body)) - pool, rest, ok := strings.Cut(strings.TrimPrefix(r.URL.Path, "/"), "/") + pub, ok := match(probe).(server.PublicRequester) if !ok { return false } - return strings.HasPrefix("/"+rest, wellKnownSegment) && userPoolIDPattern.MatchString(pool) -} + probe.Body = io.NopCloser(bytes.NewReader(body)) -// isHostedUI matches the Cognito hosted domain ({domain}.auth.{region}.amazoncognito.com, -// or {domain}.auth.localhost locally) and its path fallback /_cognito/{domain}/... -func isHostedUI(r *http.Request, _ []byte) bool { - if strings.HasPrefix(r.URL.Path, hostedUIPathPrefix) { - return true - } - - host := hostOnly(r.Host) - - return strings.Contains(host, ".auth.") && - (strings.HasSuffix(host, ".amazoncognito.com") || strings.HasSuffix(host, ".auth.localhost")) + return pub.PublicRequest(probe) } -// isExecuteAPI matches an API Gateway invocation: an execute-api host, or the -// path form /restapis/{apiId}/{stage}/_user_request_/... -func isExecuteAPI(r *http.Request, _ []byte) bool { - if strings.Contains(r.Host, ".execute-api.") { - return true - } - - rest, ok := strings.CutPrefix(r.URL.Path, restAPIsPrefix) - if !ok { - return false - } - - segs := strings.SplitN(rest, "/", userRequestSplitParts) - - return len(segs) > userRequestSegmentIndex && segs[userRequestSegmentIndex] == userRequestSegment -} - -// isUnsignedGraphQL matches an AppSync GraphQL call that carries no SigV4 -// Authorization header (API key, Cognito or OIDC auth modes). A SigV4-signed -// call uses the IAM auth mode and goes through the gate. -func isUnsignedGraphQL(r *http.Request, _ []byte) bool { - if r.Header.Get("Authorization") != "" { - return false - } - - if strings.Contains(r.Host, ".appsync-api.") { - return true - } - - return r.URL.Path == graphQLPath || strings.HasPrefix(r.URL.Path, graphQLPath+"/") -} - -func hostOnly(hostport string) string { - if h, _, err := net.SplitHostPort(hostport); err == nil { - return h - } - - return hostport -} +const urlEncodedForm = "application/x-www-form-urlencoded" diff --git a/server/aws/publicauth_test.go b/server/aws/publicauth_test.go index d772c546a..7bfc71b7f 100644 --- a/server/aws/publicauth_test.go +++ b/server/aws/publicauth_test.go @@ -5,112 +5,25 @@ import ( "crypto/sha256" "encoding/hex" "encoding/json" + "encoding/xml" "io" "net/http" "net/http/httptest" - "reflect" + "net/url" "strings" "testing" "time" "github.com/aws/aws-sdk-go-v2/aws" v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" - "github.com/aws/aws-sdk-go-v2/service/cognitoidentity" - "github.com/aws/aws-sdk-go-v2/service/cognitoidentityprovider" awssts "github.com/aws/aws-sdk-go-v2/service/sts" - smithyauth "github.com/aws/smithy-go/auth" + ststypes "github.com/aws/aws-sdk-go-v2/service/sts/types" cloudemu "github.com/stackshy/cloudemu/v2" awsprovider "github.com/stackshy/cloudemu/v2/providers/aws" iamdriver "github.com/stackshy/cloudemu/v2/services/iam/driver" ) -// anonymousOps asks an SDK client's default auth scheme resolver, which is -// generated from the service's Smithy model, which of the client's operations -// resolve to the anonymous scheme (smithy.api#noAuth). Operations are the -// client's exported methods other than Options. -func anonymousOps(t *testing.T, client any, resolve func(op string) []*smithyauth.Option) map[string]bool { - t.Helper() - - out := map[string]bool{} - typ := reflect.TypeOf(client) - - for i := range typ.NumMethod() { - op := typ.Method(i).Name - if op == "Options" { - continue - } - - opts := resolve(op) - out[op] = len(opts) > 0 && opts[0].SchemeID == smithyauth.SchemeIDAnonymous - } - - return out -} - -// newerModelNoAuth lists operations the botocore models the tables were taken -// from mark noAuth while the aws-sdk-go-v2 version pinned in go.mod still signs -// them. Drop an entry once the pinned SDK catches up. -// -//nolint:gochecknoglobals // test fixture -var newerModelNoAuth = map[string]struct{}{ - "cognito-idp GetTokensFromRefreshToken": {}, -} - -// checkAgainstModel asserts the exemption table agrees with the SDK model: every -// operation the model marks noAuth is exempt, and every exempt operation the SDK -// knows is noAuth (bar newerModelNoAuth). Exempt operations newer than the -// pinned SDK are tolerated. -func checkAgainstModel(t *testing.T, service string, model map[string]bool, exempt map[string]struct{}) { - t.Helper() - - anon := 0 - - for op, isAnon := range model { - _, ok := exempt[op] - if isAnon { - anon++ - } - - if _, newer := newerModelNoAuth[service+" "+op]; newer && !isAnon && ok { - continue - } - - if isAnon != ok { - t.Errorf("%s %s: model noAuth=%v, exempt=%v", service, op, isAnon, ok) - } - } - - if anon == 0 { - t.Fatalf("%s: SDK model reports no anonymous operations; resolver probe is broken", service) - } -} - -func TestPublicOpsMatchSDKModels(t *testing.T) { - ctx := context.Background() - - idp := cognitoidentityprovider.New(cognitoidentityprovider.Options{Region: "us-east-1"}) - checkAgainstModel(t, "cognito-idp", anonymousOps(t, idp, func(op string) []*smithyauth.Option { - o, _ := idp.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, - &cognitoidentityprovider.AuthResolverParameters{Operation: op, Region: "us-east-1"}) - return o - }), cognitoIDPPublicOps) - - ident := cognitoidentity.New(cognitoidentity.Options{Region: "us-east-1"}) - checkAgainstModel(t, "cognito-identity", anonymousOps(t, ident, func(op string) []*smithyauth.Option { - o, _ := ident.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, - &cognitoidentity.AuthResolverParameters{Operation: op, Region: "us-east-1"}) - return o - }), cognitoIdentityPublicOps) - - sts := awssts.New(awssts.Options{Region: "us-east-1"}) - checkAgainstModel(t, "sts", anonymousOps(t, sts, func(op string) []*smithyauth.Option { - o, _ := sts.Options().AuthSchemeResolver.ResolveAuthSchemes(ctx, - &awssts.AuthResolverParameters{Operation: op, Region: "us-east-1"}) - return o - }), stsPublicActions) -} - // enforcedServer starts the full AWS wire server with --enforce-auth on. func enforcedServer(t *testing.T) (*httptest.Server, *awsprovider.Provider) { t.Helper() @@ -164,6 +77,10 @@ const ( idpTarget = "AWSCognitoIdentityProviderService." identTgt = "AWSCognitoIdentityService." lambdaPath = "/2015-03-31/functions" + execHost = "abc123.execute-api.us-east-1.amazonaws.com" + + // defaultTestAccount is the account cloudemu.NewAWS() uses by default. + defaultTestAccount = "123456789012" ) func jsonRPC(target, body string) rawReq { @@ -188,24 +105,15 @@ func TestEnforcedGateAdmitsUnsignedPublicOps(t *testing.T) { req rawReq want int }{ - {"sts AssumeRoleWithWebIdentity", queryForm("Action=AssumeRoleWithWebIdentity&Version=2011-06-15" + - "&RoleArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fweb&RoleSessionName=s&WebIdentityToken=tok"), http.StatusOK}, - {"sts AssumeRoleWithSAML", queryForm("Action=AssumeRoleWithSAML&Version=2011-06-15" + - "&RoleArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Arole%2Fsaml" + - "&PrincipalArn=arn%3Aaws%3Aiam%3A%3A123456789012%3Asaml-provider%2Fidp&SAMLAssertion=eA%3D%3D"), http.StatusOK}, // Cognito user-pool public ops are not routed yet, so the handler answers // UnknownOperationException (400): the point is the gate let them through. {"cognito-idp InitiateAuth", jsonRPC(idpTarget+"InitiateAuth", `{}`), http.StatusBadRequest}, {"cognito-idp SignUp", jsonRPC(idpTarget+"SignUp", `{}`), http.StatusBadRequest}, {"cognito-idp RespondToAuthChallenge", jsonRPC(idpTarget+"RespondToAuthChallenge", `{}`), http.StatusBadRequest}, - // No identity-pool handler is served yet: the request falls to the - // dispatcher's 501, not the gate's 403. - {"cognito-identity GetId", jsonRPC(identTgt+"GetId", `{}`), http.StatusNotImplemented}, // An unknown API reaches the API Gateway data plane, which answers like // real API Gateway: 403 {"message":"Missing Authentication Token"}. That // body differs from the gate's MissingAuthenticationToken error. - {"execute-api host", rawReq{method: http.MethodGet, path: "/prod/pets", host: "abc123.execute-api.us-east-1.amazonaws.com"}, - http.StatusForbidden}, + {"execute-api host", rawReq{method: http.MethodGet, path: "/prod/pets", host: execHost}, http.StatusForbidden}, {"execute-api path", rawReq{method: http.MethodGet, path: "/restapis/abc123/prod/_user_request_/pets"}, http.StatusForbidden}, } @@ -219,220 +127,242 @@ func TestEnforcedGateAdmitsUnsignedPublicOps(t *testing.T) { } } -// TestEnforcedGateRejectsUnsignedPrivateOps asserts everything that is not a -// public operation of the handler that actually serves it is still rejected -// when unsigned, including requests that borrow a public marker (an -// execute-api Host, a public Action) but would dispatch elsewhere. -func TestEnforcedGateRejectsUnsignedPrivateOps(t *testing.T) { - ts, _ := enforcedServer(t) +// signedJSONRPC sends a SigV4-signed JSON-RPC call with creds (a session token +// in creds is sent as X-Amz-Security-Token) and returns the status and __type. +func signedJSONRPC(t *testing.T, ts *httptest.Server, creds aws.Credentials, service, target string) (int, string) { + t.Helper() - cases := []struct { - name string - req rawReq - }{ - {"ec2 DescribeInstances", queryForm("Action=DescribeInstances&Version=2016-11-15")}, - {"sts GetCallerIdentity", queryForm("Action=GetCallerIdentity&Version=2011-06-15")}, - {"sts AssumeRole", queryForm("Action=AssumeRole&Version=2011-06-15&RoleArn=x&RoleSessionName=s")}, - {"cognito-idp CreateUserPool", jsonRPC(idpTarget+"CreateUserPool", `{"PoolName":"p"}`)}, - {"cognito-idp AdminInitiateAuth", jsonRPC(idpTarget+"AdminInitiateAuth", `{}`)}, - {"cognito-identity CreateIdentityPool", jsonRPC(identTgt+"CreateIdentityPool", `{}`)}, - {"dynamodb ListTables", jsonRPC("DynamoDB_20120810.ListTables", `{}`)}, - {"s3 ListBuckets", rawReq{method: http.MethodGet, path: "/"}}, - {"execute-api host on a lambda path", rawReq{method: http.MethodGet, path: lambdaPath, - host: "abc123.execute-api.us-east-1.amazonaws.com"}}, - {"public Action on a lambda path", rawReq{method: http.MethodGet, path: lambdaPath + "?Action=AssumeRoleWithWebIdentity"}}, - // Surfaces that are exempt in real AWS but not served yet fall to S3 here, - // so they must stay gated until their own handler claims them. - {"jwks before cognito serves it", rawReq{method: http.MethodGet, path: "/us-east-1_abcDEF123/.well-known/jwks.json"}}, - {"appsync graphql before appsync serves it", rawReq{method: http.MethodPost, path: "/graphql", body: `{}`}}, - {"hosted ui before cognito serves it", rawReq{method: http.MethodGet, path: "/oauth2/userInfo", - host: "mydomain.auth.us-east-1.amazoncognito.com"}}, - } + ctx := context.Background() + body := `{}` - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - status, body := doRaw(t, ts, tc.req) - if status != http.StatusForbidden || !strings.Contains(body, missingTok) { - t.Fatalf("status %d, body %s; want 403 %s", status, body, missingTok) - } - }) + req, err := http.NewRequestWithContext(ctx, http.MethodPost, ts.URL+"/", strings.NewReader(body)) + if err != nil { + t.Fatalf("new request: %v", err) } -} - -func TestPublicRequestShapes(t *testing.T) { - mk := func(method, target, host, path, auth string) *http.Request { - r := httptest.NewRequest(method, "http://localhost"+path, nil) - if host != "" { - r.Host = host - } - if target != "" { - r.Header.Set("X-Amz-Target", target) - } + req.Header.Set("X-Amz-Target", target) + req.Header.Set("Content-Type", amzJSON11) - if auth != "" { - r.Header.Set("Authorization", auth) - } + sum := sha256.Sum256([]byte(body)) + if err := v4.NewSigner().SignHTTP(ctx, creds, req, hex.EncodeToString(sum[:]), service, "us-east-1", time.Now()); err != nil { + t.Fatalf("sign: %v", err) + } - return r + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatalf("do: %v", err) } + defer resp.Body.Close() - idp := targetIn(cognitoIDPTargetPrefix, cognitoIDPPublicOps) - ident := targetIn(cognitoIdentityTargetPrefix, cognitoIdentityPublicOps) + b, _ := io.ReadAll(resp.Body) - cases := []struct { - name string - match func(*http.Request, []byte) bool - r *http.Request - want bool - }{ - {"idp public target", idp, mk(http.MethodPost, idpTarget+"GetUser", "", "/", ""), true}, - {"idp private target", idp, mk(http.MethodPost, idpTarget+"AdminGetUser", "", "/", ""), false}, - {"idp op under the identity prefix", idp, mk(http.MethodPost, identTgt+"InitiateAuth", "", "/", ""), false}, - {"identity public target", ident, mk(http.MethodPost, identTgt+"GetCredentialsForIdentity", "", "/", ""), true}, - {"identity private target", ident, mk(http.MethodPost, identTgt+"DescribeIdentity", "", "/", ""), false}, - {"jwks", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/eu-west-2_Ab12/.well-known/jwks.json", ""), true}, - {"openid-configuration", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/us-east-1_x/.well-known/openid-configuration", ""), true}, - {"well-known on a bucket-shaped id", isUserPoolWellKnown, mk(http.MethodGet, "", "", "/mybucket/.well-known/jwks.json", ""), false}, - {"well-known POST", isUserPoolWellKnown, mk(http.MethodPost, "", "", "/us-east-1_x/.well-known/jwks.json", ""), false}, - {"hosted domain host", isHostedUI, mk(http.MethodGet, "", "d.auth.us-east-1.amazoncognito.com", "/oauth2/token", ""), true}, - {"hosted domain localhost", isHostedUI, mk(http.MethodPost, "", "d.auth.localhost:4566", "/oauth2/token", ""), true}, - {"hosted path fallback", isHostedUI, mk(http.MethodPost, "", "", "/_cognito/d/oauth2/token", ""), true}, - {"other amazoncognito host", isHostedUI, mk(http.MethodGet, "", "cognito-idp.us-east-1.amazoncognito.com", "/", ""), false}, - {"execute-api host", isExecuteAPI, mk(http.MethodGet, "", "a.execute-api.us-east-1.amazonaws.com", "/p", ""), true}, - {"execute-api path", isExecuteAPI, mk(http.MethodGet, "", "", "/restapis/a/s/_user_request_/p", ""), true}, - {"restapis control plane", isExecuteAPI, mk(http.MethodGet, "", "", "/restapis/a/resources/r", ""), false}, - {"appsync host unsigned", isUnsignedGraphQL, mk(http.MethodPost, "", "a.appsync-api.us-east-1.amazonaws.com", "/graphql", ""), true}, - {"appsync path unsigned", isUnsignedGraphQL, mk(http.MethodPost, "", "", "/graphql", ""), true}, - {"appsync signed (IAM auth mode)", isUnsignedGraphQL, mk(http.MethodPost, "", "", "/graphql", "AWS4-HMAC-SHA256 Credential=x"), false}, - {"graphql-prefixed bucket", isUnsignedGraphQL, mk(http.MethodGet, "", "", "/graphqlbucket/k", ""), false}, + var e struct { + Type string `json:"__type"` } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - if got := tc.match(tc.r, nil); got != tc.want { - t.Fatalf("match = %v, want %v", got, tc.want) - } - }) - } + _ = json.Unmarshal(b, &e) - if rt := publicRouteFor(mk(http.MethodGet, "", "", "/", ""), nil); rt != nil { - t.Fatalf("plain GET / matched a public route") - } + return resp.StatusCode, e.Type +} - form := func() *http.Request { - r := httptest.NewRequest(http.MethodPost, "http://localhost/", nil) - r.Header.Set("Content-Type", formCT) +// userWithPolicy creates an IAM user, attaches doc as a managed policy when it +// is non-empty, and returns a long-term key for it. +func userWithPolicy(t *testing.T, cloud *awsprovider.Provider, name, doc string) aws.Credentials { + t.Helper() - return r - } + ctx := context.Background() - if !isPublicSTSAction(form(), []byte("Action=AssumeRoleWithWebIdentity")) { - t.Fatalf("form Action AssumeRoleWithWebIdentity not recognized as public STS") + if _, err := cloud.IAM.CreateUser(ctx, iamdriver.UserConfig{Name: name}); err != nil { + t.Fatalf("CreateUser: %v", err) } - if isPublicSTSAction(form(), []byte("Action=GetCallerIdentity")) { - t.Fatalf("GetCallerIdentity recognized as public") + if doc != "" { + pol, err := cloud.IAM.CreatePolicy(ctx, iamdriver.PolicyConfig{Name: name + "-policy", PolicyDocument: doc}) + if err != nil { + t.Fatalf("CreatePolicy: %v", err) + } + + if err := cloud.IAM.AttachUserPolicy(ctx, name, pol.ARN); err != nil { + t.Fatalf("AttachUserPolicy: %v", err) + } } - if !isPublicSTSAction(httptest.NewRequest(http.MethodGet, "http://localhost/?Action=AssumeRoleWithSAML", nil), nil) { - t.Fatalf("query-string Action AssumeRoleWithSAML not recognized as public STS") + ak, err := cloud.IAM.CreateAccessKey(ctx, iamdriver.AccessKeyConfig{UserName: name}) + if err != nil { + t.Fatalf("CreateAccessKey: %v", err) } + + return aws.Credentials{AccessKeyID: ak.AccessKeyID, SecretAccessKey: ak.SecretAccessKey} } +const ( + allowDynamo = `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"dynamodb:*","Resource":"*"}]}` + allowSQS = `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"sqs:*","Resource":"*"}]}` + listTables = "DynamoDB_20120810.ListTables" + listQueues = "AmazonSQS.ListQueues" + accessDeny = "AccessDeniedException" +) + // TestAuthzSkipsPublicOps: a signed caller whose IAM policy allows only // DynamoDB is still served a public Cognito operation (IAM never governs it), // while a private Cognito operation is denied by the authorization gate. func TestAuthzSkipsPublicOps(t *testing.T) { ts, cloud := enforcedServer(t) - ctx := context.Background() + creds := userWithPolicy(t, cloud, "dynonly", allowDynamo) - if _, err := cloud.IAM.CreateUser(ctx, iamdriver.UserConfig{Name: "dynonly"}); err != nil { - t.Fatalf("CreateUser: %v", err) + if status, typ := signedJSONRPC(t, ts, creds, "cognito-idp", idpTarget+"InitiateAuth"); status == http.StatusForbidden { + t.Fatalf("public InitiateAuth denied: %d %s", status, typ) + } + + if status, typ := signedJSONRPC(t, ts, creds, "cognito-idp", idpTarget+"ListUserPools"); status != http.StatusForbidden || + typ != accessDeny { + t.Fatalf("private ListUserPools: %d %s, want 403 %s", status, typ, accessDeny) + } +} + +func stsClient(ts *httptest.Server, creds aws.Credentials) *awssts.Client { + return awssts.New(awssts.Options{ + Region: "us-east-1", + BaseEndpoint: aws.String(ts.URL), + Credentials: aws.CredentialsProviderFunc(func(context.Context) (aws.Credentials, error) { return creds, nil }), + }) +} + +func sessionCreds(c *ststypes.Credentials) aws.Credentials { + return aws.Credentials{ + AccessKeyID: aws.ToString(c.AccessKeyId), + SecretAccessKey: aws.ToString(c.SecretAccessKey), + SessionToken: aws.ToString(c.SessionToken), } +} - doc := `{"Version":"2012-10-17","Statement":[{"Effect":"Allow","Action":"dynamodb:*","Resource":"*"}]}` +// signedAssumeWebIdentity sends a SigV4-signed AssumeRoleWithWebIdentity for +// roleArn and returns the session credentials from the XML response. +func signedAssumeWebIdentity(t *testing.T, ts *httptest.Server, creds aws.Credentials, roleArn string) aws.Credentials { + t.Helper() - pol, err := cloud.IAM.CreatePolicy(ctx, iamdriver.PolicyConfig{Name: "dynonly", PolicyDocument: doc}) + ctx := context.Background() + body := url.Values{ + "Action": {"AssumeRoleWithWebIdentity"}, "Version": {"2011-06-15"}, "RoleArn": {roleArn}, + "RoleSessionName": {"s"}, "WebIdentityToken": {"junk"}, + }.Encode() + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, ts.URL+"/", strings.NewReader(body)) if err != nil { - t.Fatalf("CreatePolicy: %v", err) + t.Fatalf("new request: %v", err) } - if err := cloud.IAM.AttachUserPolicy(ctx, "dynonly", pol.ARN); err != nil { - t.Fatalf("AttachUserPolicy: %v", err) + req.Header.Set("Content-Type", formCT) + + sum := sha256.Sum256([]byte(body)) + if err := v4.NewSigner().SignHTTP(ctx, creds, req, hex.EncodeToString(sum[:]), "sts", "us-east-1", time.Now()); err != nil { + t.Fatalf("sign: %v", err) } - ak, err := cloud.IAM.CreateAccessKey(ctx, iamdriver.AccessKeyConfig{UserName: "dynonly"}) + resp, err := http.DefaultClient.Do(req) if err != nil { - t.Fatalf("CreateAccessKey: %v", err) + t.Fatalf("do: %v", err) + } + defer resp.Body.Close() + + var out struct { + Result struct { + Credentials struct { + AccessKeyID string `xml:"AccessKeyId"` + SecretAccessKey string `xml:"SecretAccessKey"` + SessionToken string `xml:"SessionToken"` + } `xml:"Credentials"` + } `xml:"AssumeRoleWithWebIdentityResult"` } - creds := aws.Credentials{AccessKeyID: ak.AccessKeyID, SecretAccessKey: ak.SecretAccessKey} + raw, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK || xml.Unmarshal(raw, &out) != nil { + t.Fatalf("signed AssumeRoleWithWebIdentity: %d %s", resp.StatusCode, raw) + } - send := func(op string) (int, string) { - body := `{}` + c := out.Result.Credentials - req, _ := http.NewRequestWithContext(ctx, http.MethodPost, ts.URL+"/", strings.NewReader(body)) - req.Header.Set("X-Amz-Target", idpTarget+op) - req.Header.Set("Content-Type", amzJSON11) + return aws.Credentials{AccessKeyID: c.AccessKeyID, SecretAccessKey: c.SecretAccessKey, SessionToken: c.SessionToken} +} - sum := sha256.Sum256([]byte(body)) - if err := v4.NewSigner().SignHTTP(ctx, creds, req, hex.EncodeToString(sum[:]), "cognito-idp", "us-east-1", time.Now()); err != nil { - t.Fatalf("sign: %v", err) - } +// TestSessionCredentialsAreAuthorized proves an STS session is authorized as +// its owner. A role session gets exactly its role's policies: a role that does +// not exist, or has no allowing policy, is denied. A GetSessionToken session +// gets the calling user's policies. +func TestSessionCredentialsAreAuthorized(t *testing.T) { + ts, cloud := enforcedServer(t) + ctx := context.Background() - resp, err := http.DefaultClient.Do(req) - if err != nil { - t.Fatalf("do: %v", err) - } - defer resp.Body.Close() + // "boot" has no policies, so its own key is unrestricted (bootstrap). + bootCreds := userWithPolicy(t, cloud, "boot", "") + boot := stsClient(ts, bootCreds) - b, _ := io.ReadAll(resp.Body) + trust := `{"Statement":[{"Effect":"Allow","Principal":{"AWS":"arn:aws:iam::` + defaultTestAccount + + `:root"},"Action":"sts:AssumeRole"}]}` - var e struct { - Type string `json:"__type"` + for _, role := range []string{"noperm", "dynrole"} { + if _, err := cloud.IAM.CreateRole(ctx, iamdriver.RoleConfig{Name: role, AssumeRolePolicyDoc: trust}); err != nil { + t.Fatalf("CreateRole %s: %v", role, err) } + } - _ = json.Unmarshal(b, &e) - - return resp.StatusCode, e.Type + pol, err := cloud.IAM.CreatePolicy(ctx, iamdriver.PolicyConfig{Name: "dynrole-policy", PolicyDocument: allowDynamo}) + if err != nil { + t.Fatalf("CreatePolicy: %v", err) } - if status, typ := send("InitiateAuth"); status == http.StatusForbidden || typ == "AccessDeniedException" { - t.Fatalf("public InitiateAuth denied: %d %s", status, typ) + if err := cloud.IAM.AttachRolePolicy(ctx, "dynrole", pol.ARN); err != nil { + t.Fatalf("AttachRolePolicy: %v", err) } - if status, typ := send("ListUserPools"); status != http.StatusForbidden || typ != "AccessDeniedException" { - t.Fatalf("private ListUserPools: %d %s, want 403 AccessDeniedException", status, typ) + assume := func(role string) aws.Credentials { + out, err := boot.AssumeRole(ctx, &awssts.AssumeRoleInput{ + RoleArn: aws.String("arn:aws:iam::" + defaultTestAccount + ":role/" + role), RoleSessionName: aws.String("s"), + }) + if err != nil { + t.Fatalf("AssumeRole %s: %v", role, err) + } + + return sessionCreds(out.Credentials) } -} -// TestSDKAnonymousSTSCallPassesEnforcedGate drives the real STS client, which -// sends AssumeRoleWithWebIdentity unsigned because its model marks it noAuth. -func TestSDKAnonymousSTSCallPassesEnforcedGate(t *testing.T) { - ts, _ := enforcedServer(t) + // The SDK always sends AssumeRoleWithWebIdentity unsigned (noAuth), so sign + // it by hand: an authenticated caller asking for a role that does not exist. + web := signedAssumeWebIdentity(t, ts, bootCreds, "arn:aws:iam::"+defaultTestAccount+":role/nonexistent") - client := awssts.New(awssts.Options{ - Region: "us-east-1", - BaseEndpoint: aws.String(ts.URL), - Credentials: aws.AnonymousCredentials{}, - }) + sessionFor := func(user string, doc string) aws.Credentials { + out, err := stsClient(ts, userWithPolicy(t, cloud, user, doc)).GetSessionToken(ctx, &awssts.GetSessionTokenInput{}) + if err != nil { + t.Fatalf("GetSessionToken %s: %v", user, err) + } - out, err := client.AssumeRoleWithWebIdentity(context.Background(), &awssts.AssumeRoleWithWebIdentityInput{ - RoleArn: aws.String("arn:aws:iam::123456789012:role/web"), - RoleSessionName: aws.String("s"), - WebIdentityToken: aws.String("header.payload.sig"), - }) - if err != nil { - t.Fatalf("AssumeRoleWithWebIdentity unsigned under --enforce-auth: %v", err) + return sessionCreds(out.Credentials) } - if out.Credentials == nil || aws.ToString(out.Credentials.AccessKeyId) == "" { - t.Fatalf("no credentials returned") + cases := []struct { + name string + creds aws.Credentials + service string + target string + denied bool + }{ + {"web identity session for a missing role", web, "dynamodb", listTables, true}, + {"role with no policies", assume("noperm"), "dynamodb", listTables, true}, + {"role allowed its action", assume("dynrole"), "dynamodb", listTables, false}, + {"role outside its policy", assume("dynrole"), "sqs", listQueues, true}, + {"session of a policy-limited user, outside policy", sessionFor("sqsonly", allowSQS), "dynamodb", listTables, true}, + {"session of a policy-limited user, inside policy", sessionFor("sqsonly2", allowSQS), "sqs", listQueues, false}, + {"session of an unrestricted user", sessionFor("free", ""), "dynamodb", listTables, false}, } - if _, err := client.GetCallerIdentity(context.Background(), &awssts.GetCallerIdentityInput{}); err == nil || - !strings.Contains(err.Error(), missingTok) { - t.Fatalf("unsigned GetCallerIdentity: err = %v, want %s", err, missingTok) + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + status, typ := signedJSONRPC(t, ts, tc.creds, tc.service, tc.target) + denied := status == http.StatusForbidden && typ == accessDeny + + if denied != tc.denied { + t.Fatalf("status %d %s: denied = %v, want %v", status, typ, denied, tc.denied) + } + }) } } diff --git a/server/aws/sts/operations.go b/server/aws/sts/operations.go index 1739af215..98fd781fc 100644 --- a/server/aws/sts/operations.go +++ b/server/aws/sts/operations.go @@ -6,6 +6,7 @@ import ( "strings" "time" + "github.com/stackshy/cloudemu/v2/server/authctx" "github.com/stackshy/cloudemu/v2/server/wire/awsidentity" "github.com/stackshy/cloudemu/v2/server/wire/awsquery" ) @@ -74,7 +75,8 @@ func (h *Handler) assumeRole(w http.ResponseWriter, r *http.Request) { assumedArn := "arn:aws:sts::" + h.accountID + ":assumed-role/" + roleName + "/" + sessionName assumedRoleID := assumedRoleIDPrefix + ":" + sessionName - creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}) + creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}, + roleOwner(roleName)) if !ok { return } @@ -133,7 +135,8 @@ func (h *Handler) assumeRoleWithWebIdentity(w http.ResponseWriter, r *http.Reque assumedRoleID := assumedRoleIDPrefix + ":" + sessionName - creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}) + creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}, + roleOwner(roleName)) if !ok { return } @@ -162,7 +165,8 @@ func (h *Handler) assumeRoleWithSAML(w http.ResponseWriter, r *http.Request) { assumedArn := "arn:aws:sts::" + h.accountID + ":assumed-role/" + roleName + "/" + sessionName assumedRoleID := assumedRoleIDPrefix + ":" + sessionName - creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}) + creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: assumedArn, UserID: assumedRoleID}, + roleOwner(roleName)) if !ok { return } @@ -196,7 +200,8 @@ func (h *Handler) getFederationToken(w http.ResponseWriter, r *http.Request) { fedArn := "arn:aws:sts::" + h.accountID + ":federated-user/" + name fedUserID := h.accountID + ":" + name - creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: fedArn, UserID: fedUserID}) + creds, ok := h.mintCredentials(w, durationFromForm(r), awsidentity.Identity{ARN: fedArn, UserID: fedUserID}, + h.callerOwner(r)) if !ok { return } @@ -254,7 +259,7 @@ func durationFromForm(r *http.Request) time.Duration { // or a federated user, so the minted credentials are recorded under the // identity resolveCallerIdentity resolves for the request that asked for them. func (h *Handler) getSessionToken(w http.ResponseWriter, r *http.Request) { - creds, ok := h.mintCredentials(w, durationFromForm(r), h.resolveCallerIdentity(r)) + creds, ok := h.mintCredentials(w, durationFromForm(r), h.resolveCallerIdentity(r), h.callerOwner(r)) if !ok { return } @@ -273,13 +278,15 @@ func (h *Handler) getSessionToken(w http.ResponseWriter, r *http.Request) { // always has (default, auth-off behavior is byte-for-byte unchanged). Either // way, the returned access key id is recorded under identity so a later // GetCallerIdentity call made with these credentials reflects it. -func (h *Handler) synthCredentials(dur time.Duration, identity awsidentity.Identity) (credentials, error) { +func (h *Handler) synthCredentials(dur time.Duration, identity awsidentity.Identity, owner SessionOwner) (credentials, error) { if dur <= 0 { dur = sessionDuration } if h.sessions != nil { - sess, err := h.sessions.Mint(dur) + owner.ARN, owner.UserID = identity.ARN, identity.UserID + + sess, err := h.sessions.Mint(dur, owner) if err != nil { return credentials{}, err } @@ -309,8 +316,10 @@ func (h *Handler) synthCredentials(dur time.Duration, identity awsidentity.Ident // mintCredentials builds temporary credentials representing identity for a // handler, writing an InternalFailure error response and reporting ok=false // when credential generation fails closed (a crypto/rand read error). -func (h *Handler) mintCredentials(w http.ResponseWriter, dur time.Duration, identity awsidentity.Identity) (credentials, bool) { - creds, err := h.synthCredentials(dur, identity) +func (h *Handler) mintCredentials( + w http.ResponseWriter, dur time.Duration, identity awsidentity.Identity, owner SessionOwner, +) (credentials, bool) { + creds, err := h.synthCredentials(dur, identity, owner) if err != nil { awsquery.WriteXMLError(w, http.StatusInternalServerError, "InternalFailure", "could not generate temporary credentials") @@ -321,6 +330,27 @@ func (h *Handler) mintCredentials(w http.ResponseWriter, dur time.Duration, iden return creds, true } +// roleOwner is the policy owner of a session for the assumed role roleName. +func roleOwner(roleName string) SessionOwner { + return SessionOwner{PolicyEntity: roleName, Role: true} +} + +// callerOwner is the policy owner of a session the caller mints for itself +// (GetSessionToken, GetFederationToken): the calling IAM user. A caller that is +// itself signing with a session passes that session's owner on, so a session +// can never widen its own permissions. +func (h *Handler) callerOwner(r *http.Request) SessionOwner { + p, _ := authctx.PrincipalFrom(r.Context()) + + if h.sessions != nil { + if sess, ok := h.sessions.Lookup(p.AccessKeyID); ok { + return sess.Owner + } + } + + return SessionOwner{PolicyEntity: p.UserName} +} + // roleNameFromArn extracts the role name (last path segment) from a role ARN // such as "arn:aws:iam::123456789012:role/path/MyRole". Falls back to a stable // placeholder when the ARN is missing or malformed. diff --git a/server/aws/sts/sessions.go b/server/aws/sts/sessions.go index cf661204a..53814d1f6 100644 --- a/server/aws/sts/sessions.go +++ b/server/aws/sts/sessions.go @@ -17,6 +17,24 @@ type Session struct { SecretAccessKey string SessionToken string Expiration time.Time + Owner SessionOwner +} + +// SessionOwner is who a session acts as, and whose IAM policies the +// authorization gate evaluates for requests signed with it. +type SessionOwner struct { + // ARN and UserID are the session's own identity (assumed-role/... or the + // calling user's), as GetCallerIdentity reports it. + ARN string + UserID string + // PolicyEntity is the IAM user or role friendly name whose policies govern + // the session: the assumed role, or the user that called GetSessionToken or + // GetFederationToken. + PolicyEntity string + // Role marks a role session. Its role's policies are evaluated strictly: a + // role with no allowing policy (or no such role) is denied, with none of the + // no-policy bootstrap leniency a long-term user key gets. + Role bool } // SessionStore records the temporary credentials STS issues so their signatures @@ -51,12 +69,12 @@ const secretLen = 40 // sessionTokenRandomLen is the random suffix length of a generated session token. const sessionTokenRandomLen = 32 -// Mint generates a unique temporary credential set valid for dur, records it, -// and returns it. Each call yields a distinct access key id and a fresh +// Mint generates a unique temporary credential set valid for dur acting as +// owner, records it, and returns it. Each call yields a distinct access key id and a fresh // high-entropy secret, so a caller that does not hold the issued secret cannot // forge a valid signature. It fails closed on a crypto/rand read error rather // than issuing a predictable, forgeable credential. -func (s *SessionStore) Mint(dur time.Duration) (Session, error) { +func (s *SessionStore) Mint(dur time.Duration, owner SessionOwner) (Session, error) { if dur <= 0 { dur = sessionDuration } @@ -81,6 +99,7 @@ func (s *SessionStore) Mint(dur time.Duration) (Session, error) { SecretAccessKey: secret, SessionToken: "cloudemu-session-" + token, Expiration: s.clock.Now().UTC().Add(dur), + Owner: owner, } s.mu.Lock() diff --git a/server/aws/sts/sessions_test.go b/server/aws/sts/sessions_test.go index 413581add..24ddedc14 100644 --- a/server/aws/sts/sessions_test.go +++ b/server/aws/sts/sessions_test.go @@ -16,12 +16,12 @@ import ( func TestMintYieldsDistinctHighEntropyCredentials(t *testing.T) { store := sts.NewSessionStore(config.NewFakeClock(time.Unix(0, 0))) - a, err := store.Mint(time.Hour) + a, err := store.Mint(time.Hour, sts.SessionOwner{}) if err != nil { t.Fatalf("Mint: %v", err) } - b, err := store.Mint(time.Hour) + b, err := store.Mint(time.Hour, sts.SessionOwner{}) if err != nil { t.Fatalf("Mint: %v", err) } diff --git a/server/server.go b/server/server.go index 8465d1717..ff950d7ff 100644 --- a/server/server.go +++ b/server/server.go @@ -18,6 +18,16 @@ type Handler interface { ServeHTTP(w http.ResponseWriter, r *http.Request) } +// PublicRequester is an optional Handler capability for services that serve some +// operations without credentials (AWS noAuth operations such as Cognito +// InitiateAuth, or an API Gateway invoke). PublicRequest reports whether the +// handler serves r as one of those public operations. It must answer true only +// for the exact public routes the handler itself serves, since an +// authentication hook lets such requests through unsigned. +type PublicRequester interface { + PublicRequest(r *http.Request) bool +} + // Server routes incoming HTTP requests to registered Handlers. Server itself // implements http.Handler, so httptest.NewServer(srv) works. type Server struct {