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/internal/jwtsign/jwtsign.go b/internal/jwtsign/jwtsign.go new file mode 100644 index 000000000..34c4c40cd --- /dev/null +++ b/internal/jwtsign/jwtsign.go @@ -0,0 +1,280 @@ +// 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 + + // 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 +// 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 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 { + 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+iatLeewaySeconds < 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..2cd7127cd --- /dev/null +++ b/internal/jwtsign/jwtsign_test.go @@ -0,0 +1,317 @@ +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) + 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 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) + } +} + +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/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 3e4292031..9ed13192e 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() @@ -55,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) + // 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 + ) - restore() - - if aerr != nil { - writeAuthError(w, r, aerr) - return r, false - } - - 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 { @@ -78,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 } @@ -90,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.", @@ -101,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, @@ -118,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 cd6e68cf6..c5e115bd4 100644 --- a/server/aws/authzgate.go +++ b/server/aws/authzgate.go @@ -83,6 +83,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 @@ -95,6 +100,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) @@ -108,11 +114,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/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/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 new file mode 100644 index 000000000..6e0eb4fbf --- /dev/null +++ b/server/aws/publicauth.go @@ -0,0 +1,53 @@ +package aws + +import ( + "bytes" + "io" + "net/http" + "net/url" + "strings" + + "github.com/stackshy/cloudemu/v2/server" +) + +// exemptPublic reports whether r may skip the SigV4 gate because it is an +// operation AWS serves without credentials (a noAuth operation). +// +// 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. +// +// 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 { + if _, err := url.ParseQuery(r.URL.RawQuery); err != nil { + return false + } + + if strings.HasPrefix(r.Header.Get("Content-Type"), urlEncodedForm) { + if _, err := url.ParseQuery(string(body)); err != nil { + return false + } + } + + probe := r.Clone(r.Context()) + probe.Form, probe.PostForm, probe.MultipartForm = nil, nil, nil + probe.Body = io.NopCloser(bytes.NewReader(body)) + + pub, ok := match(probe).(server.PublicRequester) + if !ok { + return false + } + + probe.Body = io.NopCloser(bytes.NewReader(body)) + + return pub.PublicRequest(probe) +} + +const urlEncodedForm = "application/x-www-form-urlencoded" diff --git a/server/aws/publicauth_test.go b/server/aws/publicauth_test.go new file mode 100644 index 000000000..7bfc71b7f --- /dev/null +++ b/server/aws/publicauth_test.go @@ -0,0 +1,368 @@ +package aws + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "encoding/xml" + "io" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4" + awssts "github.com/aws/aws-sdk-go-v2/service/sts" + 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" +) + +// 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" + 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 { + 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 + }{ + // 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}, + // 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: execHost}, 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) + } + }) + } +} + +// 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() + + ctx := context.Background() + body := `{}` + + req, err := http.NewRequestWithContext(ctx, http.MethodPost, ts.URL+"/", strings.NewReader(body)) + if err != nil { + t.Fatalf("new request: %v", err) + } + + req.Header.Set("X-Amz-Target", target) + req.Header.Set("Content-Type", amzJSON11) + + 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) + } + + 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 +} + +// 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() + + ctx := context.Background() + + if _, err := cloud.IAM.CreateUser(ctx, iamdriver.UserConfig{Name: name}); err != nil { + t.Fatalf("CreateUser: %v", err) + } + + 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) + } + } + + 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) + creds := userWithPolicy(t, cloud, "dynonly", allowDynamo) + + 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), + } +} + +// 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() + + 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("new request: %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) + } + + resp, err := http.DefaultClient.Do(req) + if err != nil { + 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"` + } + + raw, _ := io.ReadAll(resp.Body) + if resp.StatusCode != http.StatusOK || xml.Unmarshal(raw, &out) != nil { + t.Fatalf("signed AssumeRoleWithWebIdentity: %d %s", resp.StatusCode, raw) + } + + c := out.Result.Credentials + + return aws.Credentials{AccessKeyID: c.AccessKeyID, SecretAccessKey: c.SecretAccessKey, SessionToken: c.SessionToken} +} + +// 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() + + // "boot" has no policies, so its own key is unrestricted (bootstrap). + bootCreds := userWithPolicy(t, cloud, "boot", "") + boot := stsClient(ts, bootCreds) + + trust := `{"Statement":[{"Effect":"Allow","Principal":{"AWS":"arn:aws:iam::` + defaultTestAccount + + `:root"},"Action":"sts:AssumeRole"}]}` + + 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) + } + } + + pol, err := cloud.IAM.CreatePolicy(ctx, iamdriver.PolicyConfig{Name: "dynrole-policy", PolicyDocument: allowDynamo}) + if err != nil { + t.Fatalf("CreatePolicy: %v", err) + } + + if err := cloud.IAM.AttachRolePolicy(ctx, "dynrole", pol.ARN); err != nil { + t.Fatalf("AttachRolePolicy: %v", err) + } + + 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) + } + + // 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") + + 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) + } + + return sessionCreds(out.Credentials) + } + + 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}, + } + + 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 388079b3b..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 { @@ -71,17 +81,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 h := s.Match(r); h != nil { + h.ServeHTTP(w, r) - if s.observer != nil { - s.observer(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 +}