diff --git a/CLAUDE.md b/CLAUDE.md index 8f978c1c4..50ea00612 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -236,7 +236,9 @@ mcp-data-platform/ │ ├── agentinstructions/ # The deployment's customized agent-instruction layer as a policy rather than a config value (#1607): the byte bound and size advisory both its writers enforce, the config-store adapter it is read and written through, and the `mcp:knowledge_page:` index-entry form BOTH instruction layers point at a page with │ ├── admin/ # Admin-API seams built only by pkg/admin: auditapi/ (events + metrics), callapi/ (the call catalog + its review actions), catalogapi/ (OpenAPI spec bundles + embedding jobs), connoauthapi/ (connection OAuth, unified + legacy per-kind), notifyapi/ (notification delivery history + status counts), settingsapi/ (SMTP + review-queue-alert settings REST) — extracted by #1078 │ ├── apigwmetrics/ # The api gateway's outbound HTTP instrumentation as an http.RoundTripper: the connection/status labels, and the persona the call was authorized under, read off the request context the tool call carries (#1615). Holds no gateway types, and was extracted when pkg/toolkits/apigateway reached its package-size budget -│ ├── apigwtls/ # The api gateway's TLS material: what a connection's mTLS keypair and CA bundle must satisfy, and the *tls.Config its outbound transport is built with. Knows nothing about API connections — it takes the four values one carries — and was extracted when pkg/toolkits/apigateway reached its package-size budget (#1626) +│ ├── apigwtls/ # The TLS material an HTTP upstream connection carries: what its mTLS keypair and CA bundle must satisfy, and the *tls.Config its outbound transport is built with. Knows nothing about API connections — it takes the values one carries, including the kind name its refusals speak in — and is reached through internal/upstreamauth (#1626, #1647) +│ ├── upstreamauth/ # What every HTTP-based connection kind does to reach its upstream, in one copy: the Authenticator and its modes (none, bearer, api_key, basic, oauth, mtls), the operator-owned static headers and the header names a model may not claim, the connect/call timeouts and response read cap, and the TLS material. Owns the config keys those rules belong to and validates them; error text is the caller's, through Config.ErrPrefix. Extracted from pkg/toolkits/apigateway so a second kind reuses it instead of forking it (#1647) +│ ├── cfgmap/ # The typed readers over the map[string]any a connection is stored as: the tolerant String/Duration/Int64/Bool/StringMap that absorb what a JSON round-trip did to an operator's value, in one place so two kinds cannot disagree about what `"call_timeout": 30` means (#1647) │ ├── httpjson/ # RFC 9457 Problem Details responder + admin list-query param parsing, shared by the admin/portal decomposition seams (#1078) │ ├── httpserver/ # HTTP composition root: mux/route assembly (MCP streamable+SSE, OAuth, admin/portal/resources/gateway/observability REST, portal UI), CORS, drain/shutdown sequencing — extracted from main.go (#895). Subpackages are the adapters it mounts: accessgate/, attachhttp/, datahubapi/, gatewayhttp/, health/, httpauth/, mentionhttp/, notifyhttp/ (self-scoped notification prefs), scripthttp/ (managed-script admin + portal routes, including the administrator's owner transfer), sources/, unsubhttp/ (no-login unsubscribe + its tokens), versionhttp/ (#1076, #1080) │ ├── sqlgate/ # Collects the module's SQL and hands each statement to a real PostgreSQL to parse and plan (#1512); integration-tagged, so it is absent from the default build diff --git a/docs/library/stability.md b/docs/library/stability.md index 0b14b91c7..7971a4916 100644 --- a/docs/library/stability.md +++ b/docs/library/stability.md @@ -68,6 +68,22 @@ aliased back in `pkg/portal`, so `portal.Asset`, `portal.Collection`, were never a supported integration surface; the location now enforces that so their evolution cannot break an external build. +What a connection kind does to reach an HTTP upstream is one of these seams: +`internal/upstreamauth` holds the outbound authenticator and every auth mode +(`none`, `bearer`, `api_key`, `basic`, `oauth`, `mtls`), the operator-owned +static headers and the header names a model may not claim, the connect and call +timeouts, the response read cap, and the TLS material the handshake presents. +It is shared rather than copied because a second HTTP-based connection kind +answers the same questions, and one copy is what keeps a fix to, say, the +token-fetch error scrubber from landing in one kind and not the other. The +generic readers that pull a typed value out of a stored connection's +`map[string]any` sit beside it in `internal/cfgmap`. The API gateway's own +names are aliased back, so `apigateway.Authenticator`, +`apigateway.NewAuthenticator`, `apigateway.ErrNeedsReauth` and the +`AuthMode*`, `CredentialPlacement*` and `OAuth2AuthStyle*` constants are +spelled exactly as before, and every configuration key and error message an +operator sees is unchanged. + If you were importing one of these while it still lived under `pkg/`, the package moved but its API did not: the type and function names are unchanged, and the functionality is reachable through the supported surface, which diff --git a/docs/llms-full.txt b/docs/llms-full.txt index 88e09e581..12fd0ba69 100644 --- a/docs/llms-full.txt +++ b/docs/llms-full.txt @@ -2196,7 +2196,7 @@ Supported import surface (breaking changes only in a major release): | `pkg/middleware` | Request/response middleware contracts | | `pkg/toolkits/*` | Toolkit adapters' exported config types | -Other exported packages under `pkg/` are importable but are implementation packages, not a committed integration surface; their API may change in a minor release with a release-note callout. The set is bounded by a build gate (`TestPublicSurfacePolicy`, `pkg_stability_policy_test.go`) that fails when a package is added under `pkg/` outside the supported table with a single first-party importer, since that shape is an implementation seam and belongs under `internal/`; the remaining exemptions are the reference store and provider implementations a consumer passes to `platform.WithSessionStore`/`WithQueryProvider`/`WithStorageProvider` and their siblings, plus `pkg/admin` (a mountable router) and `pkg/database/migrate` (the embedded schema migrations). Facade-internal seams live under `internal/platform/`, the HTTP adapters the server mounts under `internal/httpserver/`, and the portal's own seams under `internal/portal/` (its domain types and store contracts, PostgreSQL and no-database stores, authorization core, feedback surface, public-viewer templates, rate limiter and share cache — every moved name aliased back so `portal.Asset`, `portal.Collection`, `portal.User` and the store constructors are spelled as before); all are unimportable from outside the module by Go's internal rule and were never a supported surface. +Other exported packages under `pkg/` are importable but are implementation packages, not a committed integration surface; their API may change in a minor release with a release-note callout. The set is bounded by a build gate (`TestPublicSurfacePolicy`, `pkg_stability_policy_test.go`) that fails when a package is added under `pkg/` outside the supported table with a single first-party importer, since that shape is an implementation seam and belongs under `internal/`; the remaining exemptions are the reference store and provider implementations a consumer passes to `platform.WithSessionStore`/`WithQueryProvider`/`WithStorageProvider` and their siblings, plus `pkg/admin` (a mountable router) and `pkg/database/migrate` (the embedded schema migrations). Facade-internal seams live under `internal/platform/`, the HTTP adapters the server mounts under `internal/httpserver/`, and the portal's own seams under `internal/portal/` (its domain types and store contracts, PostgreSQL and no-database stores, authorization core, feedback surface, public-viewer templates, rate limiter and share cache — every moved name aliased back so `portal.Asset`, `portal.Collection`, `portal.User` and the store constructors are spelled as before); all are unimportable from outside the module by Go's internal rule and were never a supported surface. What a connection kind does to reach an HTTP upstream is one of these seams too: `internal/upstreamauth` holds the outbound authenticator and every auth mode (`none`, `bearer`, `api_key`, `basic`, `oauth`, `mtls`), the operator-owned static headers and the header names a model may not claim, the connect/call timeouts, the response read cap and the TLS material, shared by the HTTP-based kinds rather than copied into each; `internal/cfgmap` holds the typed readers over a stored connection's `map[string]any`. `apigateway.Authenticator`, `apigateway.NewAuthenticator`, `apigateway.ErrNeedsReauth` and the `AuthMode*`, `CredentialPlacement*` and `OAuth2AuthStyle*` constants are aliased back, so the toolkit's API, its configuration keys and its error messages are unchanged. Configuration-file compatibility is handled more conservatively: additive keys ship in minor releases, and breaking renames or removals are called out in that version's release notes with an admonition (old key, new key, runtime effect). Precedent: `workflow.require_search` replaced the former `workflow.require_discovery_before_query` as a hard, non-aliased rename documented in the release notes. Set `config.strict: true` to turn an unaccounted-for rename into a hard startup error instead of a silent no-op. diff --git a/internal/apigwtls/apigwtls.go b/internal/apigwtls/apigwtls.go index a94eeb4ca..4e8c03dd5 100644 --- a/internal/apigwtls/apigwtls.go +++ b/internal/apigwtls/apigwtls.go @@ -3,11 +3,12 @@ // transport is built with. // // It is a seam of pkg/toolkits/apigateway, extracted when that package reached -// its size budget. Nothing here knows what an API connection is -- it takes -// the four values one carries -- which is why X.509 parsing, key-strength +// its size budget, and is now reached through internal/upstreamauth by every +// HTTP-based connection kind. Nothing here knows what an API connection is -- +// it takes the values one carries -- which is why X.509 parsing, key-strength // policy and PEM handling sit together rather than beside operation discovery. -// The messages keep their "apigateway:" prefix: an operator sees them when a -// connection is refused, and the prefix names that subsystem. +// Material.ErrPrefix names the kind in every message, because an operator +// reading a refused connection save should see the surface they configured. package apigwtls import ( @@ -35,11 +36,32 @@ const minRSABits = 2048 // presents, the extra CA bundle it trusts, and whether that keypair is the // connection's credential (auth_mode=mtls) rather than an addition to it. The // caller resolves ClientPairRequired from the auth mode. +// +// ErrPrefix names the connection kind in the messages Validate and Build +// produce. Empty falls back to this package's own name, which no caller should +// let an operator see: pass the kind through. type Material struct { ClientCertPEM string ClientKeyPEM string CABundlePEM string ClientPairRequired bool + ErrPrefix string +} + +// prefix returns the caller-supplied error prefix, or this package's name when +// a caller left it unset. +func (m Material) prefix() string { + if m.ErrPrefix == "" { + return "apigwtls" + } + return m.ErrPrefix +} + +// errf builds an error in the calling kind's voice. The prefix is joined to +// the format string rather than passed as an argument so the literal a reader +// (and the string-format linter) sees is the message itself. +func errf(prefix, format string, a ...any) error { + return fmt.Errorf(prefix+": "+format, a...) //nolint:err113,perfsprint // one formatting helper for the whole package } // Validate enforces the mTLS and CA-trust rules. Three independent checks: @@ -59,21 +81,22 @@ type Material struct { // certificate when set. Empty string means "no extra CAs", which is // the existing default. func Validate(m Material) error { + prefix := m.prefix() if m.ClientPairRequired { if m.ClientCertPEM == "" || m.ClientKeyPEM == "" { - return errors.New("apigateway: mtls_client_cert_pem and mtls_client_key_pem are required when auth_mode is \"mtls\"") + return errf(prefix, "mtls_client_cert_pem and mtls_client_key_pem are required when auth_mode is %q", "mtls") } } if (m.ClientCertPEM == "") != (m.ClientKeyPEM == "") { - return errors.New("apigateway: mtls_client_cert_pem and mtls_client_key_pem must both be set or both be empty") + return errf(prefix, "mtls_client_cert_pem and mtls_client_key_pem must both be set or both be empty") } if m.ClientCertPEM != "" { - if err := validateClientKeyPair(m.ClientCertPEM, m.ClientKeyPEM); err != nil { + if err := validateClientKeyPair(prefix, m.ClientCertPEM, m.ClientKeyPEM); err != nil { return err } } if m.CABundlePEM != "" { - if err := validateCABundle(m.CABundlePEM); err != nil { + if err := validateCABundle(prefix, m.CABundlePEM); err != nil { return err } } @@ -85,19 +108,19 @@ func Validate(m Material) error { // x509.ParseCertificate, and the key-matches-cert signature check; the // extra leaf-cert parse here gives a clean place to enforce minimum // key strength without re-deriving the key from raw bytes. -func validateClientKeyPair(certPEM, keyPEM string) error { +func validateClientKeyPair(prefix, certPEM, keyPEM string) error { pair, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) if err != nil { - return fmt.Errorf("apigateway: mtls cert/key invalid: %s", sanitizeKeyPairError(err)) + return errf(prefix, "mtls cert/key invalid: %s", sanitizeKeyPairError(err)) } if len(pair.Certificate) == 0 { - return errors.New("apigateway: mtls_client_cert_pem contained no certificates") + return errf(prefix, "mtls_client_cert_pem contained no certificates") } leaf, err := x509.ParseCertificate(pair.Certificate[0]) if err != nil { - return fmt.Errorf("apigateway: mtls leaf certificate unreadable: %s", err.Error()) + return errf(prefix, "mtls leaf certificate unreadable: %s", err.Error()) } - return checkKeyStrength(leaf.PublicKey) + return checkKeyStrength(prefix, leaf.PublicKey) } // sanitizeKeyPairError strips any PEM content from tls.X509KeyPair's @@ -122,11 +145,11 @@ func sanitizeKeyPairError(err error) string { // is not interoperable with most peers and stronger curves are not // supported by Go's TLS stack as of this writing); Ed25519 is // always accepted. Unknown key algorithms are rejected loudly. -func checkKeyStrength(pub any) error { +func checkKeyStrength(prefix string, pub any) error { switch k := pub.(type) { case *rsa.PublicKey: if k.N == nil || k.N.BitLen() < minRSABits { - return fmt.Errorf("apigateway: mtls private key RSA-%d is below the minimum %d bits", k.N.BitLen(), minRSABits) + return errf(prefix, "mtls private key RSA-%d is below the minimum %d bits", k.N.BitLen(), minRSABits) } return nil case *ecdsa.PublicKey: @@ -134,11 +157,11 @@ func checkKeyStrength(pub any) error { case elliptic.P256(), elliptic.P384(), elliptic.P521(): return nil } - return errors.New("apigateway: mtls private key uses an unsupported ECDSA curve (want P-256, P-384, or P-521)") + return errf(prefix, "mtls private key uses an unsupported ECDSA curve (want P-256, P-384, or P-521)") case ed25519.PublicKey: return nil default: - return fmt.Errorf("apigateway: mtls private key uses an unsupported algorithm %T", pub) + return errf(prefix, "mtls private key uses an unsupported algorithm %T", pub) } } @@ -147,7 +170,7 @@ func checkKeyStrength(pub any) error { // already filtered out the no-bundle case); a bundle with zero // CERTIFICATE blocks (e.g., one that contains only PRIVATE KEY blocks) // is rejected as misconfigured. -func validateCABundle(bundle string) error { +func validateCABundle(prefix, bundle string) error { rest := []byte(bundle) count := 0 for len(rest) > 0 { @@ -160,12 +183,12 @@ func validateCABundle(bundle string) error { continue } if _, err := x509.ParseCertificate(block.Bytes); err != nil { - return fmt.Errorf("apigateway: tls_ca_bundle_pem contains an unparseable certificate: %s", err.Error()) + return errf(prefix, "tls_ca_bundle_pem contains an unparseable certificate: %s", err.Error()) } count++ } if count == 0 { - return errors.New("apigateway: tls_ca_bundle_pem must contain at least one CERTIFICATE block") + return errf(prefix, "tls_ca_bundle_pem must contain at least one CERTIFICATE block") } return nil } @@ -189,18 +212,19 @@ func Build(m Material) (*tls.Config, error) { if !hasClient && !hasCABundle { return nil, nil //nolint:nilnil // nil config = use http.Transport defaults } + prefix := m.prefix() out := &tls.Config{MinVersion: tls.VersionTLS12} if hasClient { pair, err := tls.X509KeyPair([]byte(m.ClientCertPEM), []byte(m.ClientKeyPEM)) if err != nil { - return nil, fmt.Errorf("apigateway: building mtls keypair: %s", sanitizeKeyPairError(err)) + return nil, errf(prefix, "building mtls keypair: %s", sanitizeKeyPairError(err)) } out.Certificates = []tls.Certificate{pair} } if hasCABundle { pool, err := RootPool(m.CABundlePEM) if err != nil { - return nil, err + return nil, errf(prefix, "%w", err) } out.RootCAs = pool } @@ -221,7 +245,7 @@ func RootPool(bundle string) (*x509.CertPool, error) { pool = x509.NewCertPool() } if ok := pool.AppendCertsFromPEM([]byte(bundle)); !ok { - return nil, errors.New("apigateway: tls_ca_bundle_pem contained no valid certificates") + return nil, errors.New("tls_ca_bundle_pem contained no valid certificates") } return pool, nil } diff --git a/internal/apigwtls/apigwtls_test.go b/internal/apigwtls/apigwtls_test.go index e0f2b4f98..80a462d8d 100644 --- a/internal/apigwtls/apigwtls_test.go +++ b/internal/apigwtls/apigwtls_test.go @@ -271,3 +271,64 @@ func TestRootPoolWithBundle_RejectsInvalidPEM(t *testing.T) { require.Error(t, err) assert.Contains(t, err.Error(), "no valid certificates") } + +// TestMaterialPrefix_NamesTheCallingKind covers the field that decides +// whose voice a refusal speaks in. Every caller sets it; the fallback +// exists so a Material built without one still produces a readable +// message rather than one beginning with ": ". +func TestMaterialPrefix_NamesTheCallingKind(t *testing.T) { + assert.Equal(t, "graphql", Material{ErrPrefix: "graphql"}.prefix()) + assert.Equal(t, "apigwtls", Material{}.prefix()) +} + +// TestValidate_PrefixReachesEveryMessage pins the threading: a kind's +// name must appear on the refusals from all three arms of Validate, not +// just the first one. +func TestValidate_PrefixReachesEveryMessage(t *testing.T) { + cert, key, _ := generateCertPair(t, keyECDSAP256) + cases := []struct { + name string + m Material + }{ + {"required pair missing", Material{ErrPrefix: "graphql", ClientPairRequired: true}}, + {"ambiguous pair", Material{ErrPrefix: "graphql", ClientCertPEM: cert}}, + {"unparseable keypair", Material{ErrPrefix: "graphql", ClientCertPEM: "no", ClientKeyPEM: "pem"}}, + {"empty CA bundle", Material{ErrPrefix: "graphql", ClientCertPEM: cert, ClientKeyPEM: key, CABundlePEM: "-----BEGIN PRIVATE KEY-----\nAA==\n-----END PRIVATE KEY-----\n"}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := Validate(tc.m) + require.Error(t, err) + assert.Contains(t, err.Error(), "graphql: ") + }) + } +} + +// TestValidate_CABundleWithUnparseableCertificate covers the bundle arm +// that reaches x509: a PEM block correctly labeled CERTIFICATE whose +// bytes are not one must be refused at write time, since the pool would +// otherwise silently drop it and the connection would fail its +// handshake against a CA the operator believes is trusted. +func TestValidate_CABundleWithUnparseableCertificate(t *testing.T) { + bundle := string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: []byte("not a certificate")})) + err := Validate(Material{ErrPrefix: "apigateway", CABundlePEM: bundle}) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: tls_ca_bundle_pem contains an unparseable certificate") +} + +// TestBuild_SurfacesMaterialFaults covers Build's two error arms. They +// are reachable only when a caller skipped Validate, which is exactly +// when a clear message matters most: the alternative is a nil config +// and a handshake failure with no explanation. +func TestBuild_SurfacesMaterialFaults(t *testing.T) { + t.Run("unusable keypair", func(t *testing.T) { + _, err := Build(Material{ErrPrefix: "apigateway", ClientCertPEM: "not", ClientKeyPEM: "pem"}) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: building mtls keypair") + }) + t.Run("CA bundle with no usable certificate", func(t *testing.T) { + _, err := Build(Material{ErrPrefix: "apigateway", CABundlePEM: "not pem at all"}) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: tls_ca_bundle_pem contained no valid certificates") + }) +} diff --git a/internal/cfgmap/cfgmap.go b/internal/cfgmap/cfgmap.go new file mode 100644 index 000000000..f10cd2646 --- /dev/null +++ b/internal/cfgmap/cfgmap.go @@ -0,0 +1,143 @@ +// Package cfgmap reads typed values out of the map[string]any form a +// connection configuration takes in the platform's connection_instances +// store. +// +// A connection is stored as JSON and reaches a toolkit as a generic map, +// so every kind that parses one needs the same handful of tolerant +// readers: a value may arrive as the Go type the author intended +// (programmatic construction, YAML with a duration string) or as +// whatever JSON round-tripping produced (a float64 for an integer, a +// string for a bool). These readers absorb that spread in one place so +// two kinds cannot disagree about what `"call_timeout": 30` means. +// +// Every reader is total: an absent key, a nil map or a value of an +// unrecognized type yields the supplied default (or the zero value) +// rather than an error. Required-field enforcement belongs to the +// caller's validation, not to reading. +package cfgmap + +import ( + "maps" + "strconv" + "time" +) + +// String reads a string value. Absent, nil or non-string yields "". +func String(cfg map[string]any, key string) string { + if v, ok := cfg[key].(string); ok { + return v + } + return "" +} + +// StringDefault reads a string value, falling back to defaultVal when +// the key is absent or holds an empty string. The empty-string fallback +// matters for admin-saved connections: a form that submits every field +// writes "" for the ones the operator left blank, and those must take +// the default rather than override it with emptiness. +func StringDefault(cfg map[string]any, key, defaultVal string) string { + if v, ok := cfg[key].(string); ok && v != "" { + return v + } + return defaultVal +} + +// Duration reads a duration. A string is parsed with +// time.ParseDuration ("30s", "5m"); a bare number is read as seconds, +// which is the form a JSON round-trip of an operator's "30" produces. +// An unparseable string yields defaultVal rather than an error, so one +// malformed value cannot refuse a connection that is otherwise sound. +func Duration(cfg map[string]any, key string, defaultVal time.Duration) time.Duration { + raw, ok := cfg[key] + if !ok { + return defaultVal + } + switch v := raw.(type) { + case string: + if d, err := time.ParseDuration(v); err == nil { + return d + } + case time.Duration: + return v + case int: + return time.Duration(v) * time.Second + case int64: + return time.Duration(v) * time.Second + case float64: + return time.Duration(v) * time.Second + } + return defaultVal +} + +// Int64 reads an integer. float64 is accepted because encoding/json +// decodes every JSON number into one. +func Int64(cfg map[string]any, key string, defaultVal int64) int64 { + raw, ok := cfg[key] + if !ok { + return defaultVal + } + switch v := raw.(type) { + case int: + return int64(v) + case int64: + return v + case float64: + return int64(v) + } + return defaultVal +} + +// Bool reads a boolean flag. Absent or unrecognized values yield false. +// A string is parsed leniently (strconv.ParseBool) so YAML or JSON that +// round-trips a flag as "true"/"false" is honored alongside a native +// bool. +func Bool(cfg map[string]any, key string) bool { + raw, ok := cfg[key] + if !ok { + return false + } + switch v := raw.(type) { + case bool: + return v + case string: + b, err := strconv.ParseBool(v) + return err == nil && b + } + return false +} + +// StringMap reads a map of strings. Accepts map[string]string +// (programmatic construction) or map[string]any (YAML/JSON +// unmarshaling); non-string values inside a map[string]any are skipped. +// An absent, empty or wholly non-string map yields nil, so callers can +// test the result with len() and need no separate "was it set" flag. +func StringMap(cfg map[string]any, key string) map[string]string { + raw, ok := cfg[key] + if !ok { + return nil + } + switch v := raw.(type) { + case map[string]string: + if len(v) == 0 { + return nil + } + out := make(map[string]string, len(v)) + maps.Copy(out, v) + return out + case map[string]any: + if len(v) == 0 { + return nil + } + out := make(map[string]string, len(v)) + for k, val := range v { + if s, isStr := val.(string); isStr { + out[k] = s + } + } + if len(out) == 0 { + return nil + } + return out + } + return nil +} diff --git a/internal/cfgmap/cfgmap_test.go b/internal/cfgmap/cfgmap_test.go new file mode 100644 index 000000000..3a1714686 --- /dev/null +++ b/internal/cfgmap/cfgmap_test.go @@ -0,0 +1,189 @@ +package cfgmap + +import ( + "testing" + "time" +) + +func TestBool(t *testing.T) { + tests := []struct { + name string + val any + want bool + }{ + {"native true", true, true}, + {"native false", false, false}, + {"string true", "true", true}, + {"string 1", "1", true}, + {"string false", "false", false}, + {"string garbage", "nope", false}, + {"absent", nil, false}, + {"wrong type", 42, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := map[string]any{} + if tt.val != nil { + cfg["k"] = tt.val + } + if got := Bool(cfg, "k"); got != tt.want { + t.Errorf("Bool(%v) = %v; want %v", tt.val, got, tt.want) + } + }) + } +} + +// TestString covers the two shapes a connection config presents a +// string in: the value the operator saved, and the absence of one. +func TestString(t *testing.T) { + cfg := map[string]any{"present": "value", "wrong-type": 42} + if got := String(cfg, "present"); got != "value" { + t.Errorf("String(present) = %q; want %q", got, "value") + } + if got := String(cfg, "absent"); got != "" { + t.Errorf("String(absent) = %q; want empty", got) + } + if got := String(cfg, "wrong-type"); got != "" { + t.Errorf("String(wrong-type) = %q; want empty", got) + } + if got := String(nil, "any"); got != "" { + t.Errorf("String(nil map) = %q; want empty", got) + } +} + +// TestStringDefault pins the empty-string fallback, which is the whole +// reason this reader is distinct from String: an admin form that +// submits every field writes "" for the ones left blank, and those +// must take the default rather than override it with emptiness. +func TestStringDefault(t *testing.T) { + cases := []struct { + name string + cfg map[string]any + want string + }{ + {"set", map[string]any{"k": "given"}, "given"}, + {"empty string falls back", map[string]any{"k": ""}, "fallback"}, + {"absent falls back", map[string]any{}, "fallback"}, + {"wrong type falls back", map[string]any{"k": 7}, "fallback"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := StringDefault(tc.cfg, "k", "fallback"); got != tc.want { + t.Errorf("StringDefault = %q; want %q", got, tc.want) + } + }) + } +} + +// TestDuration covers every wire shape a duration arrives in. The bare +// number cases are what a JSON round-trip of an operator's "30" +// produces, and they must read as seconds rather than nanoseconds. +func TestDuration(t *testing.T) { + cases := []struct { + name string + val any + want time.Duration + }{ + {"duration string", "45s", 45 * time.Second}, + {"minutes string", "2m", 2 * time.Minute}, + {"unparseable string falls back", "later", 10 * time.Second}, + {"native duration", 3 * time.Second, 3 * time.Second}, + {"int seconds", 7, 7 * time.Second}, + {"int64 seconds", int64(8), 8 * time.Second}, + {"float64 seconds", float64(9), 9 * time.Second}, + {"wrong type falls back", true, 10 * time.Second}, + {"absent falls back", nil, 10 * time.Second}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := map[string]any{} + if tc.val != nil { + cfg["k"] = tc.val + } + if got := Duration(cfg, "k", 10*time.Second); got != tc.want { + t.Errorf("Duration(%v) = %v; want %v", tc.val, got, tc.want) + } + }) + } +} + +// TestInt64 covers float64, which is the only shape encoding/json ever +// produces for a JSON number and therefore the shape every saved +// connection presents. +func TestInt64(t *testing.T) { + cases := []struct { + name string + val any + want int64 + }{ + {"int", 5, 5}, + {"int64", int64(6), 6}, + {"float64 from JSON", float64(1048576), 1048576}, + {"wrong type falls back", "8", 99}, + {"absent falls back", nil, 99}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cfg := map[string]any{} + if tc.val != nil { + cfg["k"] = tc.val + } + if got := Int64(cfg, "k", 99); got != tc.want { + t.Errorf("Int64(%v) = %d; want %d", tc.val, got, tc.want) + } + }) + } +} + +// TestStringMap covers both map shapes and the nil-on-empty contract +// callers rely on to test the result with len() alone. +func TestStringMap(t *testing.T) { + t.Run("map[string]string", func(t *testing.T) { + got := StringMap(map[string]any{"k": map[string]string{"A": "1"}}, "k") + if len(got) != 1 || got["A"] != "1" { + t.Errorf("StringMap = %#v; want {A:1}", got) + } + }) + t.Run("map[string]any skips non-strings", func(t *testing.T) { + got := StringMap(map[string]any{"k": map[string]any{"A": "1", "B": 2}}, "k") + if len(got) != 1 || got["A"] != "1" { + t.Errorf("StringMap = %#v; want only the string entry", got) + } + }) + t.Run("all non-string yields nil", func(t *testing.T) { + if got := StringMap(map[string]any{"k": map[string]any{"B": 2}}, "k"); got != nil { + t.Errorf("StringMap = %#v; want nil", got) + } + }) + for _, tc := range []struct { + name string + val any + }{ + {"empty map[string]string", map[string]string{}}, + {"empty map[string]any", map[string]any{}}, + {"wrong type", "not a map"}, + } { + t.Run(tc.name, func(t *testing.T) { + if got := StringMap(map[string]any{"k": tc.val}, "k"); got != nil { + t.Errorf("StringMap = %#v; want nil", got) + } + }) + } + t.Run("absent", func(t *testing.T) { + if got := StringMap(map[string]any{}, "k"); got != nil { + t.Errorf("StringMap = %#v; want nil", got) + } + }) +} + +// TestStringMap_CopiesTheSource proves the reader hands back a copy: a +// connection's stored map must not be mutable through the value a +// toolkit reads out of it. +func TestStringMap_CopiesTheSource(t *testing.T) { + src := map[string]string{"A": "1"} + got := StringMap(map[string]any{"k": src}, "k") + got["A"] = "mutated" + if src["A"] != "1" { + t.Errorf("source map mutated through the returned copy: %#v", src) + } +} diff --git a/pkg/toolkits/apigateway/auth.go b/internal/upstreamauth/auth.go similarity index 76% rename from pkg/toolkits/apigateway/auth.go rename to internal/upstreamauth/auth.go index ceca61b5f..1ff8ad7a7 100644 --- a/pkg/toolkits/apigateway/auth.go +++ b/internal/upstreamauth/auth.go @@ -1,11 +1,10 @@ -package apigateway +package upstreamauth import ( "context" "crypto/tls" "encoding/base64" "errors" - "fmt" "net/http" "net/url" "strings" @@ -25,7 +24,7 @@ import ( // safe for concurrent use — a single Authenticator is shared across // all in-flight invocations of a connection. // -// Implementations MUST NOT log credential material. The toolkit's +// Implementations MUST NOT log credential material. The platform's // audit pipeline expects no Authorization or X-API-Key value to ever // appear in slog output, error messages, or audit rows; carelessly // formatted error strings are the most common leak path. @@ -34,7 +33,7 @@ type Authenticator interface { } // NewAuthenticator returns the Authenticator implementation for a -// validated Config. ParseConfig has already rejected unknown auth +// validated Config. ValidateAuth has already rejected unknown auth // modes, so the default branch only fires if a future mode is added // without a matching case here. func NewAuthenticator(c Config) (Authenticator, error) { @@ -42,43 +41,43 @@ func NewAuthenticator(c Config) (Authenticator, error) { case AuthModeNone: return noneAuth{}, nil case AuthModeBearer: - return bearerAuth{credential: c.Credential}, nil + return bearerAuth{cfg: c, credential: c.Credential}, nil case AuthModeAPIKey: return newAPIKeyAuth(c) case AuthModeBasic: return newBasicAuth(c) case AuthModeOAuth: // The canonical OAuth mode dispatches on the grant. The - // authorization_code variant requires a TokenStore; + // authorization_code variant requires a token store; // NewAuthenticator alone cannot supply one (it has no DB - // handle), so the toolkit's addParsedConnection wires the - // TokenStore via SetTokenStore immediately after this returns. + // handle), so the kind's connection wiring calls + // SetConnOAuthStore immediately after this returns. if c.OAuth2.Grant == connoauth.GrantAuthorizationCode { return newOAuth2AuthorizationCodeAuth(c), nil } return newOAuth2ClientCredentialsAuth(c), nil case AuthModeOAuth2ClientCredentials: - // Legacy auth_mode (hand-built Configs that bypass ParseConfig). + // Legacy auth_mode (hand-built Configs that bypass Parse). return newOAuth2ClientCredentialsAuth(c), nil case AuthModeOAuth2AuthorizationCode: - // Legacy auth_mode (hand-built Configs that bypass ParseConfig). + // Legacy auth_mode (hand-built Configs that bypass Parse). return newOAuth2AuthorizationCodeAuth(c), nil case AuthModeMTLS: // The client certificate IS the credential. No header is // added; the TLS handshake at transport setup - // (newHTTPTransport) attaches the cert. The no-op Apply + // (NewHTTPTransport) attaches the cert. The no-op Apply // keeps the invocation path's "always call Apply" contract // without a nil check. return mtlsAuth{}, nil default: - return nil, fmt.Errorf("apigateway: no authenticator for auth_mode %q", c.AuthMode) + return nil, c.errf("no authenticator for auth_mode %q", c.AuthMode) } } // mtlsAuth is a no-op Authenticator used when the connection presents // a client certificate during the TLS handshake instead of an // Authorization header. The cert + key are attached to the -// http.Transport's TLSClientConfig by newHTTPTransport; this struct +// http.Transport's TLSClientConfig by NewHTTPTransport; this struct // exists only so the invocation path can call Apply unconditionally // across every auth mode. type mtlsAuth struct{} @@ -87,7 +86,7 @@ type mtlsAuth struct{} func (mtlsAuth) Apply(_ *http.Request) error { return nil } // noneAuth applies no credential. Distinct from a nil Authenticator so -// the toolkit's invocation path can call Apply unconditionally. +// the invocation path can call Apply unconditionally. type noneAuth struct{} // Apply is a no-op; the connection requested no outbound auth. @@ -95,21 +94,23 @@ func (noneAuth) Apply(_ *http.Request) error { return nil } // bearerAuth sets the Authorization header to "Bearer ". type bearerAuth struct { + cfg Config credential string } // Apply attaches the bearer token as the Authorization header. func (b bearerAuth) Apply(req *http.Request) error { if b.credential == "" { - return errors.New("apigateway: bearer credential is empty") + return b.cfg.err("bearer credential is empty") } - req.Header.Set(authorizationHeader, "Bearer "+b.credential) + req.Header.Set(AuthorizationHeader, "Bearer "+b.credential) return nil } // apiKeyAuth attaches the credential as either a header or a query // parameter, per the connection's CredentialPlacement setting. type apiKeyAuth struct { + cfg Config credential string placement string header string @@ -118,25 +119,26 @@ type apiKeyAuth struct { func newAPIKeyAuth(c Config) (apiKeyAuth, error) { a := apiKeyAuth{ + cfg: c, credential: c.Credential, placement: c.CredentialPlacement, header: c.APIKeyHeader, param: c.APIKeyParam, } if a.credential == "" { - return apiKeyAuth{}, errors.New("apigateway: api_key credential is empty") + return apiKeyAuth{}, c.err("api_key credential is empty") } switch a.placement { case CredentialPlacementHeader: if a.header == "" { - return apiKeyAuth{}, errors.New("apigateway: api_key_header is empty") + return apiKeyAuth{}, c.err("api_key_header is empty") } case CredentialPlacementQuery: if a.param == "" { - return apiKeyAuth{}, errors.New("apigateway: api_key_param is empty") + return apiKeyAuth{}, c.err("api_key_param is empty") } default: - return apiKeyAuth{}, fmt.Errorf("apigateway: invalid api_key_placement %q", a.placement) + return apiKeyAuth{}, c.errf("invalid api_key_placement %q", a.placement) } return a, nil } @@ -152,7 +154,7 @@ func (a apiKeyAuth) Apply(req *http.Request) error { q.Set(a.param, a.credential) req.URL.RawQuery = q.Encode() default: - return fmt.Errorf("apigateway: invalid api_key_placement %q", a.placement) + return a.cfg.errf("invalid api_key_placement %q", a.placement) } return nil } @@ -168,20 +170,20 @@ type basicAuth struct { } // newBasicAuth constructs the authenticator with all validation -// re-checked against the parsed Config. Config.Validate() has already +// re-checked against the parsed Config. ValidateAuth has already // rejected ":" in the userid and CR/LF/NUL in either field, but // authenticators construct from the (validated) Config without seeing // the validator path, so the guards live here too as defense in depth -// against a future caller that bypasses Validate. +// against a future caller that bypasses validation. func newBasicAuth(c Config) (basicAuth, error) { if c.Username == "" { - return basicAuth{}, errors.New("apigateway: basic auth requires a username") + return basicAuth{}, c.err("basic auth requires a username") } if strings.Contains(c.Username, ":") { - return basicAuth{}, errors.New("apigateway: basic auth username must not contain \":\"") + return basicAuth{}, c.err("basic auth username must not contain \":\"") } if strings.ContainsAny(c.Username, "\r\n\x00") || strings.ContainsAny(c.Password, "\r\n\x00") { - return basicAuth{}, errors.New("apigateway: basic auth credentials contain CR/LF/NUL") + return basicAuth{}, c.err("basic auth credentials contain CR/LF/NUL") } encoded := base64.StdEncoding.EncodeToString([]byte(c.Username + ":" + c.Password)) return basicAuth{header: "Basic " + encoded}, nil @@ -191,7 +193,7 @@ func newBasicAuth(c Config) (basicAuth, error) { // header. newBasicAuth has already validated the inputs and computed // the encoded value, so this hot path is just a Header.Set. func (b basicAuth) Apply(req *http.Request) error { - req.Header.Set(authorizationHeader, b.header) + req.Header.Set(AuthorizationHeader, b.header) return nil } @@ -212,6 +214,7 @@ func (b basicAuth) Apply(req *http.Request) error { // any credential the operator embedded in the URL (an unfortunate // pattern but one that exists in the wild). type oauth2ClientCredentialsAuth struct { + cfg Config src oauth2.TokenSource } @@ -290,7 +293,7 @@ func newOAuth2ClientCredentialsAuth(c Config) oauth2ClientCredentialsAuth { // already does this internally but the wrap is explicit // defense against future library changes. src := oauth2.ReuseTokenSource(nil, cfg.TokenSource(ctx)) - return oauth2ClientCredentialsAuth{src: src} + return oauth2ClientCredentialsAuth{cfg: c, src: src} } // Apply fetches (or returns the cached) access token and attaches @@ -302,12 +305,12 @@ func newOAuth2ClientCredentialsAuth(c Config) oauth2ClientCredentialsAuth { func (a oauth2ClientCredentialsAuth) Apply(req *http.Request) error { tok, err := a.src.Token() if err != nil { - return tokenFetchError(err) + return tokenFetchError(a.cfg, err) } if tok == nil || tok.AccessToken == "" { - return errors.New("apigateway: oauth2 token source returned no access token") + return a.cfg.err("oauth2 token source returned no access token") } - req.Header.Set(authorizationHeader, "Bearer "+tok.AccessToken) + req.Header.Set(AuthorizationHeader, "Bearer "+tok.AccessToken) return nil } @@ -319,42 +322,44 @@ func (a oauth2ClientCredentialsAuth) Apply(req *http.Request) error { // (e.g., https://user:secret@idp.example/token), those would // leak. We rebuild the message keeping only the non-sensitive // pieces. -func tokenFetchError(err error) error { +func tokenFetchError(cfg Config, err error) error { var re *oauth2.RetrieveError if errors.As(err, &re) { - return fmt.Errorf("apigateway: oauth2 token fetch failed: status=%d", re.Response.StatusCode) + return cfg.errf("oauth2 token fetch failed: status=%d", re.Response.StatusCode) } var ue *url.Error if errors.As(err, &ue) { parsed, perr := url.Parse(ue.URL) if perr != nil { - return fmt.Errorf("apigateway: oauth2 token fetch %s: %w", ue.Op, ue.Err) + return cfg.errf("oauth2 token fetch %s: %w", ue.Op, ue.Err) } parsed.RawQuery = "" parsed.User = nil - return fmt.Errorf("apigateway: oauth2 token fetch %s %q: %w", ue.Op, parsed.String(), ue.Err) + return cfg.errf("oauth2 token fetch %s %q: %w", ue.Op, parsed.String(), ue.Err) } // Fallback: redact anything that looks URL-shaped just in case // a future library version wraps in a different error type. msg := err.Error() if strings.Contains(msg, "://") { - return errors.New("apigateway: oauth2 token fetch failed (details redacted)") + return cfg.err("oauth2 token fetch failed (details redacted)") } - return fmt.Errorf("apigateway: oauth2 token fetch failed: %s", msg) + return cfg.errf("oauth2 token fetch failed: %s", msg) } -// ErrNeedsReauth is the structured error api_invoke_endpoint surfaces +// ErrNeedsReauth is the structured error a kind's invoke tool surfaces // when an authorization_code connection's stored refresh token is // missing, expired beyond refresh_expires_at, or definitively rejected // by the IdP (RFC 6749 §5.2 invalid_grant on the refresh_token grant). // Transient failures (network, 5xx, request cancellation) DO NOT // produce this error. // -// The error message intentionally points the operator at the -// platform's reauth path rather than echoing the underlying IdP -// response (which can include sensitive material from a partial -// grant exchange). -var ErrNeedsReauth = errors.New("apigateway: oauth2 connection needs admin reconnect") +// The message carries no package prefix of its own: Apply wraps it in +// the calling kind's voice, so an operator sees "apigateway: oauth2 +// connection needs admin reconnect" while errors.Is still matches. The +// wording intentionally points at the platform's reauth path rather +// than echoing the underlying IdP response (which can include sensitive +// material from a partial grant exchange). +var ErrNeedsReauth = errors.New("oauth2 connection needs admin reconnect") // oauth2AuthorizationCodeAuth applies an OAuth 2.1 access token // acquired via the user-driven authorization_code grant. All token @@ -364,14 +369,14 @@ var ErrNeedsReauth = errors.New("apigateway: oauth2 connection needs admin recon // transparently when the cached access token is near expiry and // persists the rotated refresh token (RFC 6749 §6) back to the row. // -// The Source is built per-call from the toolkit's connOAuthStore + +// The Source is built per-call from the kind's connOAuthStore + // authevents.Writer, so the authenticator carries no token state of // its own — exactly the property that makes refresh coherent with the // background refresher across replicas. type oauth2AuthorizationCodeAuth struct { cfg Config // mu guards the store + events pair against concurrent - // SetConnOAuthStore / SetAuthEvents (called from the toolkit's + // SetConnOAuthStore / SetAuthEvents (called from the kind's // platform-side wiring) while Apply reads them on every outbound // request. RWMutex so concurrent reads don't serialize. mu sync.RWMutex @@ -384,7 +389,7 @@ func newOAuth2AuthorizationCodeAuth(c Config) *oauth2AuthorizationCodeAuth { } // SetConnOAuthStore wires the unified token store. Required before -// Apply can be called; the toolkit's SetConnOAuthStore method threads +// Apply can be called; the kind's SetConnOAuthStore method threads // this through. func (a *oauth2AuthorizationCodeAuth) SetConnOAuthStore(s connoauth.Store) { a.mu.Lock() @@ -417,44 +422,47 @@ func (a *oauth2AuthorizationCodeAuth) snapshot() (connoauth.Store, *authevents.W func (a *oauth2AuthorizationCodeAuth) Apply(req *http.Request) error { store, events := a.snapshot() if store == nil { - return errors.New("apigateway: oauth2 authorization_code: token store not wired") + return a.cfg.err("oauth2 authorization_code: token store not wired") } src := connoauth.NewSource(store, connoauth.Key{ - Kind: connoauth.KindAPI, + Kind: a.cfg.Kind, Name: a.cfg.ConnectionName, - }, connoauthConfigFromOAuth2(a.cfg)). + }, a.cfg.ConnOAuthConfig()). WithEvents(events). WithActor(authevents.SystemToolCall) token, err := src.Token(req.Context()) if err != nil { if errors.Is(err, connoauth.ErrNeedsReauth) { - return ErrNeedsReauth + return errf(a.cfg.ErrPrefix, "%w", ErrNeedsReauth) } - return fmt.Errorf("apigateway: oauth token: %w", err) + return a.cfg.errf("oauth token: %w", err) } - req.Header.Set(authorizationHeader, "Bearer "+token) + req.Header.Set(AuthorizationHeader, "Bearer "+token) return nil } -// connoauthConfigFromOAuth2 maps the toolkit's OAuth2 config slice to -// the unified connoauth.Config the Source consumes. The full Config -// (including the CA bundle for IdPs behind a private CA) is read so -// the token-exchange and refresh paths can verify the IdP's TLS cert -// against the operator's bundle without falling back to system trust. -func connoauthConfigFromOAuth2(c Config) connoauth.Config { - authStyle := oauth2.AuthStyleInHeader - if c.OAuth2.EndpointAuthStyle == OAuth2AuthStyleParams { - authStyle = oauth2.AuthStyleInParams +// SetConnOAuthStore wires the persisted OAuth token store onto an +// authenticator that needs one, and reports whether it did. Only the +// authorization_code authenticator does; every other mode returns +// false, so a kind can call this unconditionally after +// NewAuthenticator instead of type-switching at each call site. +func SetConnOAuthStore(a Authenticator, s connoauth.Store) bool { + setter, ok := a.(interface{ SetConnOAuthStore(connoauth.Store) }) + if !ok { + return false } - return connoauth.Config{ - Grant: c.OAuth2.Grant, - AuthorizationURL: c.OAuth2.AuthorizationURL, - TokenURL: c.OAuth2.TokenURL, - ClientID: c.OAuth2.ClientID, - ClientSecret: c.OAuth2.ClientSecret, - Scopes: c.OAuth2.Scopes, - EndpointAuthStyle: authStyle, - Prompt: c.OAuth2.Prompt, - CABundlePEM: c.TLSCABundlePEM, + setter.SetConnOAuthStore(s) + return true +} + +// SetAuthEvents wires the audit-event writer onto an authenticator +// that emits OAuth lifecycle events, and reports whether it did. The +// companion to SetConnOAuthStore; see its comment. +func SetAuthEvents(a Authenticator, w *authevents.Writer) bool { + setter, ok := a.(interface{ SetAuthEvents(*authevents.Writer) }) + if !ok { + return false } + setter.SetAuthEvents(w) + return true } diff --git a/pkg/toolkits/apigateway/auth_test.go b/internal/upstreamauth/auth_test.go similarity index 51% rename from pkg/toolkits/apigateway/auth_test.go rename to internal/upstreamauth/auth_test.go index 402223bc1..08f61ddfa 100644 --- a/pkg/toolkits/apigateway/auth_test.go +++ b/internal/upstreamauth/auth_test.go @@ -1,15 +1,18 @@ -package apigateway +package upstreamauth import ( "context" + "errors" "net/http" "net/http/httptest" + "net/url" "sync/atomic" "testing" "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "golang.org/x/oauth2" "github.com/txn2/mcp-data-platform/pkg/connoauth" ) @@ -103,6 +106,7 @@ func TestOAuth2AuthCode_ApplyRefreshesStaleAccessToken(t *testing.T) { })) auth := newOAuth2AuthorizationCodeAuth(Config{ + Kind: connoauth.KindAPI, ConnectionName: "fixture", OAuth2: OAuth2Config{ TokenURL: idp.URL, @@ -138,6 +142,7 @@ func TestOAuth2AuthCode_ApplyRefreshesStaleAccessToken(t *testing.T) { func TestOAuth2AuthCode_ApplyWithoutStoreErrors(t *testing.T) { auth := newOAuth2AuthorizationCodeAuth(Config{ + Kind: connoauth.KindAPI, ConnectionName: "fixture", OAuth2: OAuth2Config{TokenURL: "https://idp", ClientID: "id"}, }) @@ -150,6 +155,7 @@ func TestOAuth2AuthCode_ApplyWithoutStoreErrors(t *testing.T) { func TestOAuth2AuthCode_ApplyWithoutPersistedTokenReturnsNeedsReauth(t *testing.T) { auth := newOAuth2AuthorizationCodeAuth(Config{ + Kind: connoauth.KindAPI, ConnectionName: "fixture", OAuth2: OAuth2Config{TokenURL: "https://idp", ClientID: "id"}, }) @@ -226,8 +232,8 @@ func TestNewBasicAuth_DefenseInDepth(t *testing.T) { } } -func TestConnoauthConfigFromOAuth2_MapsAuthStyleAndScopes(t *testing.T) { - got := connoauthConfigFromOAuth2(Config{ +func TestConnOAuthConfig_MapsAuthStyleAndScopes(t *testing.T) { + got := Config{ OAuth2: OAuth2Config{ Grant: "authorization_code", AuthorizationURL: "https://idp/authorize", @@ -239,7 +245,7 @@ func TestConnoauthConfigFromOAuth2_MapsAuthStyleAndScopes(t *testing.T) { Prompt: "consent", }, TLSCABundlePEM: "-----BEGIN CERTIFICATE-----\nfake\n-----END CERTIFICATE-----\n", - }) + }.ConnOAuthConfig() assert.Equal(t, "authorization_code", got.Grant) assert.Equal(t, []string{"api", "refresh_token"}, got.Scopes) assert.Equal(t, "consent", got.Prompt) @@ -272,3 +278,207 @@ func intToString(n int32) string { } return string(buf[i:]) } + +// TestOAuth2ClientCredentials_ApplyAttachesTheFetchedToken covers the +// server-to-server grant end to end against a fake IdP: the token is +// fetched once, attached as a bearer, and reused from the library's +// cache on the next call rather than re-fetched. +func TestOAuth2ClientCredentials_ApplyAttachesTheFetchedToken(t *testing.T) { + var fetches atomic.Int32 + idp := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fetches.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`{"access_token":"cc-token","token_type":"Bearer","expires_in":3600}`)) + })) + defer idp.Close() + + auth, err := NewAuthenticator(Config{ + AuthMode: AuthModeOAuth, + OAuth2: OAuth2Config{ + Grant: connoauth.GrantClientCredentials, + TokenURL: idp.URL, + ClientID: "id", + ClientSecret: "secret", + EndpointAuthStyle: OAuth2AuthStyleHeader, + }, + }) + require.NoError(t, err) + + for range 2 { + req, rerr := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://upstream.example/x", http.NoBody) + require.NoError(t, rerr) + require.NoError(t, auth.Apply(req)) + assert.Equal(t, "Bearer cc-token", req.Header.Get("Authorization")) + } + assert.Equal(t, int32(1), fetches.Load(), "the cached token must be reused inside its lifetime") +} + +// TestOAuth2ClientCredentials_ApplySurfacesAScrubbedFetchFailure is the +// security contract on the token-fetch error path: an IdP rejection +// reaches the model as a status code, never as the IdP's response body, +// which can carry material from a partial grant exchange. +func TestOAuth2ClientCredentials_ApplySurfacesAScrubbedFetchFailure(t *testing.T) { + idp := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":"invalid_client","leaked_refresh_token":"do-not-echo"}`)) + })) + defer idp.Close() + + auth, err := NewAuthenticator(Config{ + ErrPrefix: "apigateway", + AuthMode: AuthModeOAuth, + OAuth2: OAuth2Config{ + Grant: connoauth.GrantClientCredentials, + TokenURL: idp.URL, + ClientID: "id", + ClientSecret: "secret", + EndpointAuthStyle: OAuth2AuthStyleHeader, + }, + }) + require.NoError(t, err) + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://upstream.example/x", http.NoBody) + require.NoError(t, err) + err = auth.Apply(req) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: oauth2 token fetch failed: status=401") + assert.NotContains(t, err.Error(), "do-not-echo") + assert.Empty(t, req.Header.Get("Authorization")) +} + +// TestTokenFetchError_ScrubsEveryShape walks the error types the oauth2 +// library can hand back. Each branch exists to keep a credential out of +// the message: the IdP's body, userinfo or a query string embedded in +// the token URL, and anything URL-shaped from a future library version. +func TestTokenFetchError_ScrubsEveryShape(t *testing.T) { + cfg := Config{ErrPrefix: "apigateway"} + + t.Run("RetrieveError keeps only the status", func(t *testing.T) { + err := tokenFetchError(cfg, &oauth2.RetrieveError{ + Response: &http.Response{StatusCode: http.StatusForbidden}, + Body: []byte(`{"secret":"leaked"}`), + }) + assert.Equal(t, "apigateway: oauth2 token fetch failed: status=403", err.Error()) + }) + + t.Run("url.Error drops userinfo and query", func(t *testing.T) { + err := tokenFetchError(cfg, &url.Error{ + Op: "Post", + URL: "https://user:s3cret@idp.example/token?client_secret=alsosecret", + Err: errors.New("dial tcp: connection refused"), + }) + msg := err.Error() + assert.Contains(t, msg, "https://idp.example/token") + assert.NotContains(t, msg, "s3cret") + assert.NotContains(t, msg, "alsosecret") + }) + + t.Run("unparseable url.Error keeps only the op", func(t *testing.T) { + err := tokenFetchError(cfg, &url.Error{ + Op: "Post", + URL: "://not a url", + Err: errors.New("boom"), + }) + msg := err.Error() + assert.Contains(t, msg, "oauth2 token fetch Post") + assert.NotContains(t, msg, "not a url") + }) + + t.Run("unknown URL-shaped error is redacted wholesale", func(t *testing.T) { + err := tokenFetchError(cfg, errors.New("boom talking to https://user:pw@idp.example/token")) + assert.Equal(t, "apigateway: oauth2 token fetch failed (details redacted)", err.Error()) + }) + + t.Run("unknown plain error passes through", func(t *testing.T) { + err := tokenFetchError(cfg, errors.New("context deadline exceeded")) + assert.Equal(t, "apigateway: oauth2 token fetch failed: context deadline exceeded", err.Error()) + }) +} + +// TestOAuth2ClientCredentials_ApplyRejectsAnEmptyAccessToken covers the +// last-resort guard on a token source that succeeds but yields nothing. +// The oauth2 library refuses an empty access_token itself, so this is +// reachable only by substituting the source — which is exactly the +// point: sending "Bearer " would produce an opaque 401 from the +// upstream instead of naming the real fault. +func TestOAuth2ClientCredentials_ApplyRejectsAnEmptyAccessToken(t *testing.T) { + cases := []struct { + name string + src oauth2.TokenSource + }{ + {name: "empty access token", src: stubTokenSource{token: &oauth2.Token{}}}, + {name: "nil token", src: stubTokenSource{}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + auth := oauth2ClientCredentialsAuth{cfg: Config{ErrPrefix: "apigateway"}, src: tc.src} + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://upstream.example/x", http.NoBody) + require.NoError(t, err) + err = auth.Apply(req) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: oauth2 token source returned no access token") + assert.Empty(t, req.Header.Get("Authorization")) + }) + } +} + +// stubTokenSource hands back whatever the test configured, including +// nothing at all, so the authenticator's own guards are reachable +// without an IdP that violates the OAuth spec. +type stubTokenSource struct { + token *oauth2.Token +} + +func (s stubTokenSource) Token() (*oauth2.Token, error) { return s.token, nil } + +// TestOAuth2ConfigFromConnoauth_MapsTheParamsAuthStyle covers the +// non-default endpoint auth style, which some IdPs require and which +// travels as an enum on one side and a string on the other. +func TestOAuth2ConfigFromConnoauth_MapsTheParamsAuthStyle(t *testing.T) { + got := oauth2ConfigFromConnoauth(connoauth.Config{EndpointAuthStyle: oauth2.AuthStyleInParams}) + assert.Equal(t, OAuth2AuthStyleParams, got.EndpointAuthStyle) + + got = oauth2ConfigFromConnoauth(connoauth.Config{EndpointAuthStyle: oauth2.AuthStyleInHeader}) + assert.Equal(t, OAuth2AuthStyleHeader, got.EndpointAuthStyle) +} + +// TestOAuth2AuthCode_ApplyKeepsATransientFailureOutOfReauth is the +// classification contract: only a definitively dead credential asks the +// operator to reconnect. A store that cannot be read is transient, and +// reporting it as ErrNeedsReauth would send an admin to redo a browser +// flow over what is probably a database blip. +func TestOAuth2AuthCode_ApplyKeepsATransientFailureOutOfReauth(t *testing.T) { + auth := newOAuth2AuthorizationCodeAuth(Config{ + ErrPrefix: "apigateway", + Kind: connoauth.KindAPI, + ConnectionName: "fixture", + OAuth2: OAuth2Config{TokenURL: "https://idp", ClientID: "id"}, + }) + auth.SetConnOAuthStore(failingStore{err: errors.New("database unreachable")}) + + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://upstream.example/x", http.NoBody) + require.NoError(t, err) + err = auth.Apply(req) + require.Error(t, err) + assert.NotErrorIs(t, err, ErrNeedsReauth) + assert.Contains(t, err.Error(), "apigateway: oauth token:") + assert.Empty(t, req.Header.Get("Authorization")) +} + +// failingStore fails every read so the authenticator's transient-error +// arm is reachable without a database. +type failingStore struct { + err error +} + +func (s failingStore) Get(context.Context, connoauth.Key) (*connoauth.PersistedToken, error) { + return nil, s.err +} +func (s failingStore) Set(context.Context, connoauth.PersistedToken) error { return s.err } +func (s failingStore) Delete(context.Context, connoauth.Key) error { return s.err } +func (s failingStore) List(context.Context) ([]connoauth.PersistedToken, error) { + return nil, s.err +} + +func (failingStore) Lock(context.Context, connoauth.Key) (func(), error) { + return func() {}, nil +} diff --git a/internal/upstreamauth/config.go b/internal/upstreamauth/config.go new file mode 100644 index 000000000..8e934955d --- /dev/null +++ b/internal/upstreamauth/config.go @@ -0,0 +1,650 @@ +// Package upstreamauth holds the outbound authentication and transport +// policy shared by the platform's HTTP-based connection kinds. +// +// A connection kind that reaches an upstream over HTTP has to answer the +// same questions whatever it carries on the wire: which credential goes +// on the request, which headers the operator pins and the model may not +// touch, how long a dial and a call may take, how much of a response may +// be read, and which TLS material the handshake presents. Those answers +// are the same for an OpenAPI-described REST API and for a GraphQL +// endpoint, and keeping one copy of them is what stops a fix to, say, +// the token-fetch error scrubber from landing in one kind and not the +// other. +// +// The package deliberately knows nothing about tools, catalogs, +// operations or documents. It takes the slice of a connection's +// configuration that concerns credentials and transport, and it produces +// an Authenticator and an *http.Client. Each kind keeps its own exported +// Config with its own keys and maps the overlapping fields onto this +// one, so this package never appears in a kind's public API. +// +// Error text is owned by the caller. Config.ErrPrefix names the kind in +// every message this package produces ("apigateway: credential is +// required ..."), because an operator reading a refused connection save +// should see the surface they configured, not the seam behind it. +package upstreamauth + +import ( + "errors" + "fmt" + "strings" + "time" + + "golang.org/x/oauth2" + + "github.com/txn2/mcp-data-platform/internal/cfgmap" + "github.com/txn2/mcp-data-platform/pkg/connoauth" +) + +const ( + // AuthModeNone disables outbound authentication. + AuthModeNone = "none" + // AuthModeBearer sends "Authorization: Bearer ". + AuthModeBearer = "bearer" + // AuthModeAPIKey sends the credential as a header (default + // "X-API-Key") or as a query parameter; placement and key name are + // per-connection so APIs that use non-standard schemes (e.g. an + // "api_key" query parameter, or a custom "X-Api-Token" header) can + // be onboarded without code changes. + AuthModeAPIKey = "api_key" + // AuthModeBasic sends "Authorization: Basic base64(username:password)" + // per RFC 7617. Required for the long tail of older REST APIs (Jenkins, + // on-prem Jira / Confluence Server / DC, internal apps) that never moved + // to bearer or OAuth. RFC 7617 §2 forbids ":" in the userid; password + // may be empty (some APIs accept "token:" as a bearer-token-in-username + // pattern). Password is encrypted at rest via the platform's + // FieldEncryptor (the "password" config key is already in the + // sensitive-keys list). + AuthModeBasic = "basic" + + // AuthModeOAuth is the canonical OAuth auth_mode shared across + // every toolkit kind (see connoauth.AuthModeOAuth). The specific + // flow is carried separately in OAuth2Config.Grant. Parse + // normalizes the legacy api-only auth_mode values below to this + // form, so a parsed Config always reports AuthModeOAuth for an + // OAuth connection. + AuthModeOAuth = connoauth.AuthModeOAuth + + // AuthModeOAuth2ClientCredentials is the legacy api-only auth_mode + // that encoded the client_credentials grant in the mode string. + // Retained so raw config authored before the schema unified (and + // hand-built test Configs) still parse; Parse normalizes it to + // AuthModeOAuth + Grant=client_credentials. + // + // client_credentials acquires a bearer token — server-to-server, + // no human in the loop. The platform exchanges the configured + // client_id + client_secret for a token at OAuth.TokenURL and + // applies it as "Authorization: Bearer " on outbound + // calls. Tokens are cached + refreshed automatically by the + // underlying golang.org/x/oauth2 library; no DB state is + // required because every restart can re-acquire from credentials. + AuthModeOAuth2ClientCredentials = "oauth2_client_credentials" // #nosec G101 -- mode name, not a credential + + // AuthModeOAuth2AuthorizationCode runs the user-driven OAuth 2.1 + // authorization-code grant: an admin completes a one-time browser + // flow at connection setup; the resulting refresh token is + // persisted (encrypted) so subsequent platform restarts and + // background workloads keep working without further interaction. + // Tokens are refreshed automatically before expiry. Requires the + // platform's database (refresh-token state survives restarts). + AuthModeOAuth2AuthorizationCode = "oauth2_authorization_code" // #nosec G101 -- mode name, not a credential + + // AuthModeMTLS authenticates by presenting an X.509 client + // certificate during the TLS handshake per RFC 5246 / 8446. No + // Authorization header is sent: the cert IS the credential. + // Used by upstreams that map the cert's subject DN (or a SAN) to + // a user identity in their authorizer, including service-mesh + // peers, PKI-fronted internal APIs, healthcare integration + // engines, financial messaging endpoints, and FedRAMP / DoD- + // boundary services. Requires both mtls_client_cert_pem and + // mtls_client_key_pem on the connection config; the mTLS material + // can also be present alongside other auth modes (bearer + mTLS, + // etc.), but auth_mode=mtls is the explicit "no header + // credential" signal. + AuthModeMTLS = "mtls" + + // CredentialPlacementHeader (default) sends the credential as an HTTP + // header named by APIKeyHeader. + CredentialPlacementHeader = "header" + // CredentialPlacementQuery sends the credential as a URL query parameter + // named by APIKeyParam. + CredentialPlacementQuery = "query" + + // DefaultAPIKeyHeader is the conventional API-key header name when + // the connection does not specify one. + DefaultAPIKeyHeader = "X-API-Key" // #nosec G101 -- header name, not a credential + + // DefaultConnectTimeout caps the time spent establishing the + // outbound connection (TCP + TLS handshake) on each invocation. + DefaultConnectTimeout = 10 * time.Second + // DefaultCallTimeout caps the total per-call time including + // upstream processing and response read. + DefaultCallTimeout = 60 * time.Second + + // DefaultMaxResponseBytes is the upstream read cap: the most a kind + // reads of any one response. It bounds transfer and buffering, not + // what reaches the model; a kind's own inline budget does that. + DefaultMaxResponseBytes = int64(10 * 1024 * 1024) +) + +// EndpointAuthStyle values. +const ( + OAuth2AuthStyleHeader = "header" + OAuth2AuthStyleParams = "params" +) + +// cfgKey* constants name the keys this package reads from the +// map[string]any form a connection takes in the platform's +// connection_instances store. A kind that embeds this policy does not +// re-declare them: it hands the whole map to Parse and keeps only its +// own keys. +const ( + cfgKeyAuthMode = "auth_mode" + cfgKeyCredential = "credential" // #nosec G101 -- map key, not a secret + cfgKeyAPIKeyHeader = "api_key_header" // #nosec G101 -- map key, not a credential + cfgKeyAPIKeyParam = "api_key_param" // #nosec G101 -- map key, not a credential + cfgKeyAPIKeyPlacement = "api_key_placement" + cfgKeyUsername = "username" + cfgKeyPassword = "password" // #nosec G101 -- map key, not a credential; encryption handled by platform FieldEncryptor sensitive-keys list + cfgKeyConnectTimeout = "connect_timeout" + cfgKeyCallTimeout = "call_timeout" + + cfgKeyMaxResponseBytes = "max_response_bytes" + + // cfgKeyStaticHeaders holds operator-configured headers appended to + // every outbound request. Required for upstreams that demand BOTH + // an Authorization bearer AND a separate subscription/key header + // (Google Cloud's x-goog-user-project quota-billing header, vendor + // subscription keys, a GraphQL server's folder-routing header). + // Stored as a map[string]any whose values are encrypted at rest by + // platform.FieldEncryptor (see CfgKeyStaticHeaders in + // pkg/platform/fieldcrypt.go). + cfgKeyStaticHeaders = "static_headers" + + // OAuth config keys are owned by pkg/connoauth (the canonical + // oauth_* vocabulary plus the legacy oauth2_* fallback). This + // package delegates OAuth parsing to connoauth.ParseConfig rather + // than declaring its own key constants. The keys remain top-level + // (not nested) so the platform's FieldEncryptor, which walks only + // the top level of the config map, encrypts oauth_client_secret at + // rest without changes to the encryptor. + + // mTLS material config keys. Cert and CA bundle are public + // material (plain text at rest); the private key is in the + // platform's sensitive-keys list (see pkg/platform/fieldcrypt.go) + // and encrypted via FieldEncryptor like every other secret on a + // connection. + cfgKeyMTLSClientCertPEM = "mtls_client_cert_pem" // #nosec G101 -- map key, not a credential + cfgKeyMTLSClientKeyPEM = "mtls_client_key_pem" // #nosec G101 -- map key, not a credential + cfgKeyTLSCABundlePEM = "tls_ca_bundle_pem" + + cfgKeyIdentityPassthrough = "identity_passthrough" +) + +// Config is the authentication and transport slice of a connection's +// configuration. A kind builds one from its own Config and hands it to +// NewAuthenticator and NewHTTPClient. +type Config struct { + // Kind is the connoauth connection kind ("api", "graphql"). It + // keys the persisted OAuth token row and identifies the connection + // in connoauth's deduplicated configuration warnings, so two kinds + // with a same-named connection do not share a token. + Kind string + // ErrPrefix names the calling kind in every error this package + // produces. Empty falls back to the package name. Set it: an + // operator whose connection save is refused should read + // "apigateway: credential is required ...", not the name of an + // internal seam they cannot see in any configuration file. + ErrPrefix string + // ConnectionName is the audit-visible connection identifier. Used + // as the OAuth token-row key for the authorization_code grant. + // Kinds populate it from the toolkit instance name after Parse. + ConnectionName string + + // AuthMode selects the credential scheme: one of the AuthMode* + // constants. + AuthMode string + // Credential is the bearer token or API key. Ignored when AuthMode + // is "none". Encrypted at rest via the platform's FieldEncryptor. + Credential string + // CredentialPlacement is "header" (default) or "query" — only consulted + // when AuthMode is "api_key". + CredentialPlacement string + // APIKeyHeader is the header name to set when CredentialPlacement is + // "header". Defaults to DefaultAPIKeyHeader. + APIKeyHeader string + // APIKeyParam is the query parameter name when CredentialPlacement is + // "query". No default — required when placement is "query". + APIKeyParam string + // Username is the userid for HTTP Basic auth (RFC 7617). Required + // when AuthMode is "basic". Ignored otherwise. Not a secret on its + // own (per RFC 7617 §2 the userid is sent in clear after base64 + // decoding regardless), so it is not encrypted at rest. + Username string + // Password is the password for HTTP Basic auth. May be empty: some + // legacy APIs accept a bearer token in the userid slot with an empty + // password (the "token:" pattern). Encrypted at rest via the + // platform's FieldEncryptor. + Password string + // OAuth2 carries the OAuth 2.1 parameters used when AuthMode is + // AuthModeOAuth. Empty for non-OAuth modes. + OAuth2 OAuth2Config + + // ConnectTimeout caps the dial step (TCP + TLS handshake) on each + // invocation. + ConnectTimeout time.Duration + // CallTimeout caps the total per-invocation time. + CallTimeout time.Duration + // MaxResponseBytes is the upstream read cap: the most a kind reads + // of any one response. Defaults to DefaultMaxResponseBytes. + MaxResponseBytes int64 + // StaticHeaders are operator-configured headers attached to every + // outbound request, in addition to whatever AuthMode contributes. + // Operator-supplied; the model never sets or overrides these. + // Values are encrypted at rest. + StaticHeaders map[string]string + + // MTLSClientCertPEM is the PEM-encoded X.509 client certificate + // chain (leaf first) presented during the TLS handshake. Public + // material, stored in plain text. Required alongside + // MTLSClientKeyPEM and optional otherwise; an ambiguous config + // (one set, the other empty) is refused. + MTLSClientCertPEM string + // MTLSClientKeyPEM is the PEM-encoded private key matching + // MTLSClientCertPEM. Encrypted at rest via the platform's + // FieldEncryptor. Validation runs the cert + key through + // tls.X509KeyPair so a key that does not match the cert is + // rejected at write time, not on first outbound call. + MTLSClientKeyPEM string + // TLSCABundlePEM is an optional PEM bundle of root CA + // certificates added to the TLS trust store for outbound + // requests on this connection. Appended to the system root + // pool, not substituted: public CAs remain trusted. Required + // when the upstream's TLS certificate is signed by a private + // CA (cluster-internal CA, mesh CA, corporate root) that the + // host's default cert store does not carry. + TLSCABundlePEM string + + // IdentityPassthrough forwards the acting caller's inbound bearer + // token as the outbound Authorization header, instead of applying + // this connection's shared credential. When set, AuthMode must be + // "none": the shared-credential Authenticator is skipped, and two + // sources for one header would be ambiguous. Reading the caller's + // token off the request context is the kind's job, not this + // package's — only the invariant lives here. + IdentityPassthrough bool +} + +// OAuth2Config describes the OAuth 2.1 grant parameters. For +// client_credentials the platform exchanges ClientID + ClientSecret at +// TokenURL for an access token (cached + refreshed by the +// golang.org/x/oauth2 library). For authorization_code an admin +// completes a one-time browser flow and the persisted refresh token is +// read through connoauth.Source on every call. +type OAuth2Config struct { + // Grant is the OAuth flow, populated by Parse from the canonical + // oauth_grant (or derived from a legacy auth_mode). One of + // connoauth.GrantClientCredentials or + // connoauth.GrantAuthorizationCode. The authenticator and + // validation dispatch on this rather than on the auth_mode string. + Grant string + // TokenURL is the upstream's token endpoint. Required. + TokenURL string + // ClientID is the platform's registered client id. Required. + ClientID string + // ClientSecret is the platform's registered client secret. + // Required. Encrypted at rest via the platform's FieldEncryptor. + ClientSecret string + // Scopes is an optional list of OAuth scopes to request. + Scopes []string + // EndpointAuthStyle controls how the client credentials are + // transmitted at token-fetch time. "header" (default) sends + // them as HTTP Basic auth on the token request; "params" + // sends them as POST body parameters. Some IdPs require one + // or the other; "header" is the OAuth 2.1 default. + EndpointAuthStyle string + // AuthorizationURL is the upstream's authorization endpoint. + // Required only for the authorization_code grant — that's + // where the platform redirects the admin's browser to start + // the flow. + AuthorizationURL string + // Prompt is an optional OIDC prompt parameter (RFC OIDC + // §3.1.2.1). Common values: "login" (force credential prompt), + // "consent" (force consent screen), "select_account", + // "none" (silent auth). Empty by default — the IdP decides. + // Operators of strict OIDC realms (Keycloak, Auth0, Okta) + // typically set this to "login" so admin Reconnect actions + // always re-prompt the user. Pure-OAuth (non-OIDC) providers + // often reject unknown parameters with invalid_request, so + // leave empty for those. + Prompt string +} + +// Parse reads the authentication and transport keys out of a +// connection's config map and applies the defaults, leaving the kind to +// read its own keys from the same map. The returned Config is NOT +// validated: a kind interleaves these checks with its own so the first +// error an operator sees is the same one it was before this policy was +// shared (see Validate and the individual validators). +// +// kind is the connoauth connection kind, errPrefix names the kind in +// error text, and endpointURL identifies the connection in connoauth's +// deduplicated warnings about legacy keys and cleartext endpoints. +func Parse(kind, errPrefix, endpointURL string, cfg map[string]any) (Config, error) { + c := Config{ + Kind: kind, + ErrPrefix: errPrefix, + AuthMode: AuthModeNone, + CredentialPlacement: CredentialPlacementHeader, + APIKeyHeader: DefaultAPIKeyHeader, + ConnectTimeout: DefaultConnectTimeout, + CallTimeout: DefaultCallTimeout, + MaxResponseBytes: DefaultMaxResponseBytes, + } + c.AuthMode = cfgmap.StringDefault(cfg, cfgKeyAuthMode, c.AuthMode) + c.Credential = cfgmap.String(cfg, cfgKeyCredential) + c.CredentialPlacement = cfgmap.StringDefault(cfg, cfgKeyAPIKeyPlacement, c.CredentialPlacement) + c.APIKeyHeader = cfgmap.StringDefault(cfg, cfgKeyAPIKeyHeader, c.APIKeyHeader) + c.APIKeyParam = cfgmap.String(cfg, cfgKeyAPIKeyParam) + c.Username = cfgmap.String(cfg, cfgKeyUsername) + c.Password = cfgmap.String(cfg, cfgKeyPassword) + c.ConnectTimeout = cfgmap.Duration(cfg, cfgKeyConnectTimeout, c.ConnectTimeout) + c.CallTimeout = cfgmap.Duration(cfg, cfgKeyCallTimeout, c.CallTimeout) + c.MaxResponseBytes = cfgmap.Int64(cfg, cfgKeyMaxResponseBytes, c.MaxResponseBytes) + if isOAuthAuthMode(c.AuthMode) { + // Delegate OAuth parsing to the shared connoauth.ParseConfig + // (canonical oauth_* keys, legacy oauth2_* fallback, grant + // derivation) and normalize the auth_mode to the canonical + // AuthModeOAuth so the authenticator and validation dispatch on + // the grant rather than on three divergent mode strings. + parsed, err := connoauth.ParseConfig(kind, endpointURL, cfg) + if err != nil { + return Config{}, errf(errPrefix, "%w", err) + } + c.AuthMode = AuthModeOAuth + c.OAuth2 = oauth2ConfigFromConnoauth(parsed) + } + c.StaticHeaders = cfgmap.StringMap(cfg, cfgKeyStaticHeaders) + c.MTLSClientCertPEM = cfgmap.String(cfg, cfgKeyMTLSClientCertPEM) + c.MTLSClientKeyPEM = cfgmap.String(cfg, cfgKeyMTLSClientKeyPEM) + c.TLSCABundlePEM = cfgmap.String(cfg, cfgKeyTLSCABundlePEM) + c.IdentityPassthrough = cfgmap.Bool(cfg, cfgKeyIdentityPassthrough) + return c, nil +} + +// Validate runs every check this package owns, in the order a kind with +// no additional keys would want them. A kind that interleaves its own +// checks calls the individual validators instead. +func (c Config) Validate() error { + if err := c.ValidateAuth(); err != nil { + return err + } + if err := c.ValidateTransport(); err != nil { + return err + } + if err := c.ValidateStaticHeaders(); err != nil { + return err + } + if err := c.ValidateIdentityPassthrough(); err != nil { + return err + } + return c.ValidateTLSMaterial() +} + +// ValidateTransport enforces that the timeouts and the read cap are +// positive. Zero would mean "no timeout" to net/http and "read nothing" +// to the response reader, neither of which any operator intends. +func (c Config) ValidateTransport() error { + if c.ConnectTimeout <= 0 { + return c.err("connect_timeout must be positive") + } + if c.CallTimeout <= 0 { + return c.err("call_timeout must be positive") + } + if c.MaxResponseBytes <= 0 { + return c.err("max_response_bytes must be positive") + } + return nil +} + +// ValidateAuth enforces the per-mode credential requirements. +func (c Config) ValidateAuth() error { + switch c.AuthMode { + case AuthModeNone: + return nil + case AuthModeBearer: + if c.Credential == "" { + return c.err("credential is required when auth_mode is \"bearer\"") + } + return nil + case AuthModeAPIKey: + return c.validateAPIKeyAuth() + case AuthModeBasic: + return c.validateBasicAuth() + case AuthModeOAuth, AuthModeOAuth2ClientCredentials, AuthModeOAuth2AuthorizationCode: + return c.validateOAuthAuth() + case AuthModeMTLS: + // The mTLS material is validated centrally by + // ValidateTLSMaterial so the same rules apply whether mTLS is + // the credential or layered on top of + // bearer/api_key/basic/oauth. The mode-specific requirement + // (cert + key MUST be present) is enforced there via + // Config.AuthMode inspection. + return nil + default: + return c.errf("invalid auth_mode %q (want none, bearer, api_key, basic, oauth2_client_credentials, oauth2_authorization_code, or mtls)", c.AuthMode) + } +} + +// validateOAuthAuth dispatches to the grant-specific rules. A parsed +// Config always reports the canonical mode and carries the grant in +// OAuth2.Grant, so the grant decides. A hand-built Config that bypassed +// Parse still encodes the grant in the mode string, and there the mode +// decides — including when it contradicts the grant field, which is how +// the mode-per-grant validation behaved before the two collapsed. +func (c Config) validateOAuthAuth() error { + switch c.AuthMode { + case AuthModeOAuth2ClientCredentials: + return c.validateOAuth2() + case AuthModeOAuth2AuthorizationCode: + return c.validateOAuth2AuthCode() + } + if c.OAuth2.Grant == connoauth.GrantAuthorizationCode { + return c.validateOAuth2AuthCode() + } + return c.validateOAuth2() +} + +// validateBasicAuth enforces RFC 7617 + the platform's smuggling +// defenses for the "basic" auth mode. The userid (username) must be +// non-empty and contain no ":" (RFC 7617 §2 forbids it because the +// decoder splits on the first colon). Both fields must be free of +// CR/LF/NUL because neither RFC 7617 nor base64 stops an operator from +// pasting a "username\r\nX-Smuggled: 1" string that would inject +// extra headers after the Authorization line. The password may be +// empty: some legacy APIs accept a bearer token in the userid slot with +// an empty password (the "token:" pattern), so refusing empty here +// would block a real use case. +func (c Config) validateBasicAuth() error { + if c.Username == "" { + return c.err("username is required when auth_mode is \"basic\"") + } + // Smuggling defenses run before the colon check: a payload like + // "alice\r\nX-Smuggled: 1" contains both CRLF and ":" and we want + // the security-relevant error to surface, not the RFC compliance + // one. + if strings.ContainsAny(c.Username, "\r\n\x00") { + return c.err("username contains CR/LF/NUL header smuggling vector") + } + if strings.ContainsAny(c.Password, "\r\n\x00") { + return c.err("password contains CR/LF/NUL header smuggling vector") + } + if strings.Contains(c.Username, ":") { + return c.err("username must not contain \":\" (RFC 7617 §2 forbids it in the userid)") + } + return nil +} + +// validateOAuth2AuthCode adds the authorization_code-specific +// requirement (AuthorizationURL) on top of the client_credentials +// validation. ClientSecret is still required because OAuth 2.1 +// authorization-code with confidential clients exchanges +// (client_id, client_secret, code) for tokens. +func (c Config) validateOAuth2AuthCode() error { + if err := c.validateOAuth2(); err != nil { + return err + } + if c.OAuth2.AuthorizationURL == "" { + return c.err("oauth2.authorization_url is required when auth_mode is \"oauth2_authorization_code\"") + } + return nil +} + +func (c Config) validateOAuth2() error { + if c.OAuth2.TokenURL == "" { + return c.err("oauth2.token_url is required when auth_mode is \"oauth2_client_credentials\"") + } + if c.OAuth2.ClientID == "" { + return c.err("oauth2.client_id is required when auth_mode is \"oauth2_client_credentials\"") + } + if c.OAuth2.ClientSecret == "" { + return c.err("oauth2.client_secret is required when auth_mode is \"oauth2_client_credentials\"") + } + switch c.OAuth2.EndpointAuthStyle { + case OAuth2AuthStyleHeader, OAuth2AuthStyleParams: + return nil + default: + return c.errf("invalid oauth2.endpoint_auth_style %q (want %q or %q)", + c.OAuth2.EndpointAuthStyle, OAuth2AuthStyleHeader, OAuth2AuthStyleParams) + } +} + +func (c Config) validateAPIKeyAuth() error { + if c.Credential == "" { + return c.err("credential is required when auth_mode is \"api_key\"") + } + switch c.CredentialPlacement { + case CredentialPlacementHeader: + if c.APIKeyHeader == "" { + return c.err("api_key_header must not be empty") + } + case CredentialPlacementQuery: + if c.APIKeyParam == "" { + return c.err("api_key_param is required when api_key_placement is \"query\"") + } + default: + return c.errf("invalid api_key_placement %q (want header or query)", c.CredentialPlacement) + } + return nil +} + +// ValidateIdentityPassthrough enforces that a passthrough connection +// carries no shared credential. Passthrough forwards the caller's inbound +// token as the Authorization header, so a configured auth_mode would +// either be ignored (confusing) or fight for the same header. Requiring +// auth_mode=none keeps the single-credential-source invariant explicit. +func (c Config) ValidateIdentityPassthrough() error { + if c.IdentityPassthrough && c.AuthMode != AuthModeNone { + return c.errf("identity_passthrough requires auth_mode=none, got %q", c.AuthMode) + } + return nil +} + +// IsOAuthAuthorizationCode reports whether the connection uses the +// OAuth authorization_code grant (canonical AuthModeOAuth plus that +// grant). The admin redirect handler and the kind handlers gate the +// one-time browser flow on this, so they do not depend on the raw +// auth_mode string shape. +func (c Config) IsOAuthAuthorizationCode() bool { + return c.AuthMode == AuthModeOAuth && c.OAuth2.Grant == connoauth.GrantAuthorizationCode +} + +// isOAuthAuthMode reports whether mode names an OAuth connection in any +// of the recognized input shapes: the canonical AuthModeOAuth or either +// legacy api-only mode that encoded the grant. Parse uses this to decide +// when to delegate to connoauth.ParseConfig and normalize. +func isOAuthAuthMode(mode string) bool { + switch mode { + case AuthModeOAuth, AuthModeOAuth2ClientCredentials, AuthModeOAuth2AuthorizationCode: + return true + default: + return false + } +} + +// oauth2ConfigFromConnoauth projects the shared connoauth.Config onto +// this package's OAuth2Config. The endpoint auth style is mapped back to +// the operator-facing string form the authenticators and validation +// expect. +func oauth2ConfigFromConnoauth(c connoauth.Config) OAuth2Config { + style := OAuth2AuthStyleHeader + if c.EndpointAuthStyle == oauth2.AuthStyleInParams { + style = OAuth2AuthStyleParams + } + return OAuth2Config{ + Grant: c.Grant, + TokenURL: c.TokenURL, + ClientID: c.ClientID, + ClientSecret: c.ClientSecret, + Scopes: c.Scopes, + EndpointAuthStyle: style, + AuthorizationURL: c.AuthorizationURL, + Prompt: c.Prompt, + } +} + +// ConnOAuthConfig maps this connection's OAuth settings to the unified +// connoauth.Config the Source consumes. The CA bundle travels with it +// so the token-exchange and refresh paths can verify an IdP behind a +// private CA without falling back to system trust. +// +// Exported because the initial authorization-code exchange (the admin +// OAuth kind handler) and the per-call silent refresh (the +// authenticator) must read every field through the same translator: a +// regression here would otherwise drop CABundlePEM, Prompt, or a future +// field from one path but not the other. +func (c Config) ConnOAuthConfig() connoauth.Config { + authStyle := oauth2.AuthStyleInHeader + if c.OAuth2.EndpointAuthStyle == OAuth2AuthStyleParams { + authStyle = oauth2.AuthStyleInParams + } + return connoauth.Config{ + Grant: c.OAuth2.Grant, + AuthorizationURL: c.OAuth2.AuthorizationURL, + TokenURL: c.OAuth2.TokenURL, + ClientID: c.OAuth2.ClientID, + ClientSecret: c.OAuth2.ClientSecret, + Scopes: c.OAuth2.Scopes, + EndpointAuthStyle: authStyle, + Prompt: c.OAuth2.Prompt, + CABundlePEM: c.TLSCABundlePEM, + } +} + +// prefixOr returns the caller-supplied error prefix, or this package's +// name when a caller left it unset. +func prefixOr(p string) string { + if p == "" { + return "upstreamauth" + } + return p +} + +// err builds an error in the calling kind's voice. +func (c Config) err(msg string) error { + return errors.New(prefixOr(c.ErrPrefix) + ": " + msg) +} + +// errf builds a formatted error in the calling kind's voice. +func (c Config) errf(format string, a ...any) error { + return errf(c.ErrPrefix, format, a...) +} + +// errf builds a formatted error in a kind's voice from the prefix +// alone, for the paths that have one before they have a Config. The +// prefix is joined to the format string rather than passed as an +// argument so the literal a reader (and the string-format linter) sees +// is the message itself. +func errf(prefix, format string, a ...any) error { + return fmt.Errorf(prefixOr(prefix)+": "+format, a...) +} diff --git a/internal/upstreamauth/config_test.go b/internal/upstreamauth/config_test.go new file mode 100644 index 000000000..c67fbd6cf --- /dev/null +++ b/internal/upstreamauth/config_test.go @@ -0,0 +1,518 @@ +package upstreamauth + +import ( + "maps" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/txn2/mcp-data-platform/pkg/connoauth" +) + +// TestParse_AppliesDefaults pins what a connection that configures +// nothing gets: no credential, a header-placed API key name, and the +// three transport bounds. An operator who saves an empty form must land +// on these, not on zero values that would mean "no timeout" to +// net/http and "read nothing" to the body reader. +func TestParse_AppliesDefaults(t *testing.T) { + c, err := Parse(connoauth.KindAPI, "apigateway", "https://upstream.example", map[string]any{}) + require.NoError(t, err) + assert.Equal(t, AuthModeNone, c.AuthMode) + assert.Equal(t, CredentialPlacementHeader, c.CredentialPlacement) + assert.Equal(t, DefaultAPIKeyHeader, c.APIKeyHeader) + assert.Equal(t, DefaultConnectTimeout, c.ConnectTimeout) + assert.Equal(t, DefaultCallTimeout, c.CallTimeout) + assert.Equal(t, DefaultMaxResponseBytes, c.MaxResponseBytes) + assert.Equal(t, connoauth.KindAPI, c.Kind) + assert.Equal(t, "apigateway", c.ErrPrefix) +} + +// TestParse_ReadsEveryKey walks one connection carrying every key this +// package owns, in the shapes a saved connection presents them in. +func TestParse_ReadsEveryKey(t *testing.T) { + c, err := Parse(connoauth.KindAPI, "apigateway", "https://upstream.example", map[string]any{ + "auth_mode": AuthModeAPIKey, + "credential": "k3y", + "api_key_placement": CredentialPlacementQuery, + "api_key_header": "X-Custom", + "api_key_param": "apikey", + "username": "alice", + "password": "s3cret", + "connect_timeout": "3s", + "call_timeout": float64(90), + "max_response_bytes": float64(2048), + "static_headers": map[string]any{"X-Goog-User-Project": "proj"}, + "mtls_client_cert_pem": "cert", + "mtls_client_key_pem": "key", + "tls_ca_bundle_pem": "bundle", + "identity_passthrough": true, + }) + require.NoError(t, err) + assert.Equal(t, AuthModeAPIKey, c.AuthMode) + assert.Equal(t, "k3y", c.Credential) + assert.Equal(t, CredentialPlacementQuery, c.CredentialPlacement) + assert.Equal(t, "X-Custom", c.APIKeyHeader) + assert.Equal(t, "apikey", c.APIKeyParam) + assert.Equal(t, "alice", c.Username) + assert.Equal(t, "s3cret", c.Password) + assert.Equal(t, 3*time.Second, c.ConnectTimeout) + assert.Equal(t, 90*time.Second, c.CallTimeout) + assert.Equal(t, int64(2048), c.MaxResponseBytes) + assert.Equal(t, map[string]string{"X-Goog-User-Project": "proj"}, c.StaticHeaders) + assert.Equal(t, "cert", c.MTLSClientCertPEM) + assert.Equal(t, "key", c.MTLSClientKeyPEM) + assert.Equal(t, "bundle", c.TLSCABundlePEM) + assert.True(t, c.IdentityPassthrough) +} + +// TestParse_NormalizesEveryOAuthModeSpelling proves the two legacy +// auth_mode strings that encoded the grant collapse to the canonical +// mode with the grant carried separately, so the authenticator and the +// validators dispatch on one value rather than three. +func TestParse_NormalizesEveryOAuthModeSpelling(t *testing.T) { + base := map[string]any{ + "oauth_token_url": "https://idp.example/token", + "oauth_authorization_url": "https://idp.example/authorize", + "oauth_client_id": "id", + "oauth_client_secret": "secret", + } + cases := []struct { + mode string + wantGrant string + }{ + {AuthModeOAuth2ClientCredentials, connoauth.GrantClientCredentials}, + {AuthModeOAuth2AuthorizationCode, connoauth.GrantAuthorizationCode}, + } + for _, tc := range cases { + t.Run(tc.mode, func(t *testing.T) { + cfg := map[string]any{"auth_mode": tc.mode} + maps.Copy(cfg, base) + c, err := Parse(connoauth.KindAPI, "apigateway", "https://upstream.example", cfg) + require.NoError(t, err) + assert.Equal(t, AuthModeOAuth, c.AuthMode, "legacy mode must normalize to the canonical one") + assert.Equal(t, tc.wantGrant, c.OAuth2.Grant) + assert.Equal(t, OAuth2AuthStyleHeader, c.OAuth2.EndpointAuthStyle) + }) + } +} + +// TestParse_OAuthConfigErrorCarriesTheKindPrefix covers the one path +// Parse can fail on: connoauth refuses a malformed endpoint URL, and +// the refusal must reach the operator in the calling kind's voice. +func TestParse_OAuthConfigErrorCarriesTheKindPrefix(t *testing.T) { + _, err := Parse(connoauth.KindAPI, "apigateway", "https://upstream.example", map[string]any{ + "auth_mode": AuthModeOAuth, + "oauth_token_url": "://not a url", + }) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: ") +} + +// TestValidate_RunsEveryCheck exercises the composite entry point a +// kind with no keys of its own uses: a sound config passes, and a +// failure in any arm surfaces. +func TestValidate_RunsEveryCheck(t *testing.T) { + sound := Config{ + ErrPrefix: "graphql", + AuthMode: AuthModeBearer, + Credential: "tok", + ConnectTimeout: DefaultConnectTimeout, + CallTimeout: DefaultCallTimeout, + MaxResponseBytes: DefaultMaxResponseBytes, + } + require.NoError(t, sound.Validate()) + + cases := []struct { + name string + mutate func(c Config) Config + wantMsg string + }{ + {"auth", func(c Config) Config { c.Credential = ""; return c }, "credential is required"}, + {"transport", func(c Config) Config { c.CallTimeout = 0; return c }, "call_timeout must be positive"}, + {"static headers", func(c Config) Config { + c.StaticHeaders = map[string]string{"Authorization": "x"} + return c + }, "static_headers must not set Authorization"}, + {"identity passthrough", func(c Config) Config { c.IdentityPassthrough = true; return c }, "identity_passthrough requires auth_mode=none"}, + {"tls material", func(c Config) Config { c.MTLSClientCertPEM = "cert-only"; return c }, "must both be set or both be empty"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.mutate(sound).Validate() + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + assert.Contains(t, err.Error(), "graphql: ", "the message must be in the calling kind's voice") + }) + } +} + +// TestValidateAuth_PerMode walks every credential rule the package +// enforces. The messages are the ones an operator reads when a +// connection save is refused, so each case pins the wording that names +// the offending key. +func TestValidateAuth_PerMode(t *testing.T) { + cases := []struct { + name string + cfg Config + wantMsg string // empty = must pass + }{ + {name: "none", cfg: Config{AuthMode: AuthModeNone}}, + {name: "bearer ok", cfg: Config{AuthMode: AuthModeBearer, Credential: "t"}}, + {name: "bearer without credential", cfg: Config{AuthMode: AuthModeBearer}, wantMsg: "credential is required"}, + { + name: "api_key header ok", + cfg: Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: DefaultAPIKeyHeader}, + }, + { + name: "api_key without credential", + cfg: Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: DefaultAPIKeyHeader}, + wantMsg: "credential is required", + }, + { + name: "api_key header empty", + cfg: Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementHeader}, + wantMsg: "api_key_header must not be empty", + }, + { + name: "api_key query ok", + cfg: Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementQuery, APIKeyParam: "key"}, + }, + { + name: "api_key query without param", + cfg: Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementQuery}, + wantMsg: "api_key_param is required", + }, + { + name: "api_key bad placement", + cfg: Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: "body"}, + wantMsg: "invalid api_key_placement", + }, + {name: "basic ok", cfg: Config{AuthMode: AuthModeBasic, Username: "alice", Password: "p"}}, + {name: "basic empty password ok", cfg: Config{AuthMode: AuthModeBasic, Username: "token"}}, + {name: "basic without username", cfg: Config{AuthMode: AuthModeBasic}, wantMsg: "username is required"}, + { + name: "basic username with colon", + cfg: Config{AuthMode: AuthModeBasic, Username: "a:b"}, + wantMsg: "must not contain", + }, + { + name: "basic username with CRLF", + cfg: Config{AuthMode: AuthModeBasic, Username: "a\r\nX-Smuggled: 1"}, + wantMsg: "username contains CR/LF/NUL", + }, + { + name: "basic password with NUL", + cfg: Config{AuthMode: AuthModeBasic, Username: "alice", Password: "p\x00"}, + wantMsg: "password contains CR/LF/NUL", + }, + {name: "mtls defers to the TLS validator", cfg: Config{AuthMode: AuthModeMTLS}}, + {name: "unknown mode", cfg: Config{AuthMode: "future"}, wantMsg: "invalid auth_mode"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.cfg.ValidateAuth() + if tc.wantMsg == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + }) + } +} + +// TestValidateAuth_OAuthGrants covers the OAuth arm across the +// canonical mode and both legacy spellings, since each reaches the +// grant-specific validator by a different route. +func TestValidateAuth_OAuthGrants(t *testing.T) { + full := OAuth2Config{ + Grant: connoauth.GrantAuthorizationCode, + TokenURL: "https://idp/token", + AuthorizationURL: "https://idp/authorize", + ClientID: "id", + ClientSecret: "secret", + EndpointAuthStyle: OAuth2AuthStyleHeader, + } + cases := []struct { + name string + cfg Config + wantMsg string + }{ + {name: "canonical authorization_code", cfg: Config{AuthMode: AuthModeOAuth, OAuth2: full}}, + {name: "legacy authorization_code", cfg: Config{AuthMode: AuthModeOAuth2AuthorizationCode, OAuth2: full}}, + { + name: "canonical client_credentials", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: OAuth2Config{ + Grant: connoauth.GrantClientCredentials, TokenURL: "https://idp/token", + ClientID: "id", ClientSecret: "s", EndpointAuthStyle: OAuth2AuthStyleParams, + }}, + }, + { + name: "legacy client_credentials", + cfg: Config{AuthMode: AuthModeOAuth2ClientCredentials, OAuth2: OAuth2Config{ + TokenURL: "https://idp/token", ClientID: "id", ClientSecret: "s", + EndpointAuthStyle: OAuth2AuthStyleHeader, + }}, + }, + { + name: "authorization_code without authorization_url", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: withField(full, func(o *OAuth2Config) { o.AuthorizationURL = "" })}, + wantMsg: "oauth2.authorization_url is required", + }, + { + name: "missing token_url", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: withField(full, func(o *OAuth2Config) { o.TokenURL = "" })}, + wantMsg: "oauth2.token_url is required", + }, + { + name: "missing client_id", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: withField(full, func(o *OAuth2Config) { o.ClientID = "" })}, + wantMsg: "oauth2.client_id is required", + }, + { + name: "missing client_secret", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: withField(full, func(o *OAuth2Config) { o.ClientSecret = "" })}, + wantMsg: "oauth2.client_secret is required", + }, + { + name: "invalid endpoint auth style", + cfg: Config{AuthMode: AuthModeOAuth, OAuth2: withField(full, func(o *OAuth2Config) { o.EndpointAuthStyle = "cookie" })}, + wantMsg: "invalid oauth2.endpoint_auth_style", + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.cfg.ValidateAuth() + if tc.wantMsg == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + }) + } +} + +// withField returns a copy of o with one field changed, so the table +// above can express "the full config, minus this" without restating +// every field per case. +func withField(o OAuth2Config, mutate func(*OAuth2Config)) OAuth2Config { + mutate(&o) + return o +} + +func TestValidateTransport(t *testing.T) { + sound := Config{ConnectTimeout: time.Second, CallTimeout: time.Second, MaxResponseBytes: 1} + require.NoError(t, sound.ValidateTransport()) + + cases := []struct { + name string + mutate func(c Config) Config + wantMsg string + }{ + {"connect timeout", func(c Config) Config { c.ConnectTimeout = 0; return c }, "connect_timeout must be positive"}, + {"call timeout", func(c Config) Config { c.CallTimeout = -time.Second; return c }, "call_timeout must be positive"}, + {"read cap", func(c Config) Config { c.MaxResponseBytes = 0; return c }, "max_response_bytes must be positive"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.mutate(sound).ValidateTransport() + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + }) + } +} + +// TestValidateIdentityPassthrough enforces the single-credential-source +// invariant: passthrough forwards the caller's own token, so a +// configured auth mode would be a second, contradictory source for the +// same header. +func TestValidateIdentityPassthrough(t *testing.T) { + require.NoError(t, Config{IdentityPassthrough: true, AuthMode: AuthModeNone}.ValidateIdentityPassthrough()) + require.NoError(t, Config{IdentityPassthrough: false, AuthMode: AuthModeBearer}.ValidateIdentityPassthrough()) + err := Config{IdentityPassthrough: true, AuthMode: AuthModeBearer}.ValidateIdentityPassthrough() + require.Error(t, err) + assert.Contains(t, err.Error(), "identity_passthrough requires auth_mode=none") +} + +func TestIsOAuthAuthorizationCode(t *testing.T) { + assert.True(t, Config{ + AuthMode: AuthModeOAuth, + OAuth2: OAuth2Config{Grant: connoauth.GrantAuthorizationCode}, + }.IsOAuthAuthorizationCode()) + assert.False(t, Config{ + AuthMode: AuthModeOAuth, + OAuth2: OAuth2Config{Grant: connoauth.GrantClientCredentials}, + }.IsOAuthAuthorizationCode()) + assert.False(t, Config{AuthMode: AuthModeBearer}.IsOAuthAuthorizationCode()) +} + +// TestErrPrefix_FallsBackToThePackageName covers the unset case. No +// kind should reach it, but a Config built by hand must still produce a +// readable error rather than one that starts with ": ". +func TestErrPrefix_FallsBackToThePackageName(t *testing.T) { + err := Config{AuthMode: AuthModeBearer}.ValidateAuth() + require.Error(t, err) + assert.Contains(t, err.Error(), "upstreamauth: ") +} + +// --- authenticator behavior ----------------------------------------- + +// TestAuthenticatorApply_PerMode asserts what each mode actually puts +// on the wire. The header values are pinned literally so a refactor +// that changes an encoding fails here rather than against an upstream. +func TestAuthenticatorApply_PerMode(t *testing.T) { + cases := []struct { + name string + cfg Config + wantHeader string + wantValue string + wantQuery string + }{ + { + name: "bearer", + cfg: Config{AuthMode: AuthModeBearer, Credential: "T0K3N"}, + wantHeader: "Authorization", wantValue: "Bearer T0K3N", + }, + { + name: "api_key in a header", + cfg: Config{ + AuthMode: AuthModeAPIKey, Credential: "k3y", + CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: "X-Custom-Key", + }, + wantHeader: "X-Custom-Key", wantValue: "k3y", + }, + { + name: "api_key in the query", + cfg: Config{ + AuthMode: AuthModeAPIKey, Credential: "k3y", + CredentialPlacement: CredentialPlacementQuery, APIKeyParam: "apikey", + }, + wantQuery: "apikey=k3y", + }, + { + name: "basic", + cfg: Config{AuthMode: AuthModeBasic, Username: "alice", Password: "s3cret"}, + wantHeader: "Authorization", wantValue: "Basic YWxpY2U6czNjcmV0", + }, + {name: "none touches nothing", cfg: Config{AuthMode: AuthModeNone}}, + {name: "mtls touches nothing", cfg: Config{AuthMode: AuthModeMTLS}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + auth, err := NewAuthenticator(tc.cfg) + require.NoError(t, err) + req, err := http.NewRequest(http.MethodGet, "https://upstream.example/x", http.NoBody) //nolint:noctx // no request is sent + require.NoError(t, err) + require.NoError(t, auth.Apply(req)) + if tc.wantHeader != "" { + assert.Equal(t, tc.wantValue, req.Header.Get(tc.wantHeader)) + } else { + assert.Empty(t, req.Header.Get("Authorization")) + } + if tc.wantQuery != "" { + assert.Equal(t, tc.wantQuery, req.URL.RawQuery) + } + }) + } +} + +// TestNewAuthenticator_RejectsIncompleteCredentials covers the +// authenticator's own defense-in-depth guards, which run even when a +// caller constructed the Config without going through validation. +func TestNewAuthenticator_RejectsIncompleteCredentials(t *testing.T) { + cases := []struct { + name string + cfg Config + wantMsg string + }{ + {"api_key empty credential", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: "X"}, "api_key credential is empty"}, + {"api_key empty header", Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementHeader}, "api_key_header is empty"}, + {"api_key empty param", Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: CredentialPlacementQuery}, "api_key_param is empty"}, + {"api_key bad placement", Config{AuthMode: AuthModeAPIKey, Credential: "k", CredentialPlacement: "body"}, "invalid api_key_placement"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + _, err := NewAuthenticator(tc.cfg) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + }) + } +} + +// TestBearerAuth_ApplyRefusesAnEmptyCredential covers the runtime guard +// on a Config that reached the authenticator without validation: an +// empty token must be refused rather than sent as a bare "Bearer ". +func TestBearerAuth_ApplyRefusesAnEmptyCredential(t *testing.T) { + req, err := http.NewRequest(http.MethodGet, "https://upstream.example/x", http.NoBody) //nolint:noctx // no request is sent + require.NoError(t, err) + err = bearerAuth{cfg: Config{ErrPrefix: "apigateway"}}.Apply(req) + require.Error(t, err) + assert.Contains(t, err.Error(), "bearer credential is empty") +} + +// TestAPIKeyAuth_ApplyRefusesAnInvalidPlacement covers the switch's +// default arm, reachable only by mutating a constructed authenticator — +// the guard that keeps a future placement value from silently sending +// no credential at all. +func TestAPIKeyAuth_ApplyRefusesAnInvalidPlacement(t *testing.T) { + req, err := http.NewRequest(http.MethodGet, "https://upstream.example/x", http.NoBody) //nolint:noctx // no request is sent + require.NoError(t, err) + a := apiKeyAuth{cfg: Config{ErrPrefix: "apigateway"}, credential: "k", placement: "body"} + err = a.Apply(req) + require.Error(t, err) + assert.Contains(t, err.Error(), "invalid api_key_placement") +} + +// TestSetConnOAuthStore_OnlyWiresTheGrantThatNeedsIt lets a kind call +// the wiring helpers over every connection without type-switching: +// they report whether the authenticator took the value. +func TestSetConnOAuthStore_OnlyWiresTheGrantThatNeedsIt(t *testing.T) { + authCode, err := NewAuthenticator(Config{ + AuthMode: AuthModeOAuth, + OAuth2: OAuth2Config{Grant: connoauth.GrantAuthorizationCode, TokenURL: "https://idp/token", ClientID: "id"}, + }) + require.NoError(t, err) + assert.True(t, SetConnOAuthStore(authCode, connoauth.NewMemoryStore())) + assert.True(t, SetAuthEvents(authCode, nil)) + + bearer, err := NewAuthenticator(Config{AuthMode: AuthModeBearer, Credential: "t"}) + require.NoError(t, err) + assert.False(t, SetConnOAuthStore(bearer, connoauth.NewMemoryStore())) + assert.False(t, SetAuthEvents(bearer, nil)) +} + +// TestValidateOAuthAuth_LegacyModeOutranksTheGrantField pins the one +// case where the mode string and the grant field can disagree: a +// hand-built Config that bypassed Parse. Parse never produces one — it +// normalizes the legacy spellings to the canonical mode and carries the +// grant separately — so the rule exists to keep a Config assembled in +// code validating the way its mode says, as it did before the two +// dispatch paths collapsed into one. +func TestValidateOAuthAuth_LegacyModeOutranksTheGrantField(t *testing.T) { + // client_credentials mode, authorization_code grant: the mode wins, + // so the missing authorization_url is not required. + err := Config{ + AuthMode: AuthModeOAuth2ClientCredentials, + OAuth2: OAuth2Config{ + Grant: connoauth.GrantAuthorizationCode, TokenURL: "https://idp/token", + ClientID: "id", ClientSecret: "s", EndpointAuthStyle: OAuth2AuthStyleHeader, + }, + }.ValidateAuth() + require.NoError(t, err) + + // authorization_code mode, client_credentials grant: the mode wins + // the other way, and the missing authorization_url is refused. + err = Config{ + AuthMode: AuthModeOAuth2AuthorizationCode, + OAuth2: OAuth2Config{ + Grant: connoauth.GrantClientCredentials, TokenURL: "https://idp/token", + ClientID: "id", ClientSecret: "s", EndpointAuthStyle: OAuth2AuthStyleHeader, + }, + }.ValidateAuth() + require.Error(t, err) + assert.Contains(t, err.Error(), "oauth2.authorization_url is required") +} diff --git a/pkg/toolkits/apigateway/tls.go b/internal/upstreamauth/tls.go similarity index 67% rename from pkg/toolkits/apigateway/tls.go rename to internal/upstreamauth/tls.go index 29f40259b..fae691255 100644 --- a/pkg/toolkits/apigateway/tls.go +++ b/internal/upstreamauth/tls.go @@ -1,4 +1,4 @@ -package apigateway +package upstreamauth import ( "crypto/tls" @@ -15,20 +15,21 @@ func (c Config) tlsMaterial() apigwtls.Material { ClientKeyPEM: c.MTLSClientKeyPEM, CABundlePEM: c.TLSCABundlePEM, ClientPairRequired: c.AuthMode == AuthModeMTLS, + ErrPrefix: prefixOr(c.ErrPrefix), } } -// validateTLSMaterial enforces the per-connection mTLS and CA-trust rules, so +// ValidateTLSMaterial enforces the per-connection mTLS and CA-trust rules, so // a misconfiguration is refused at admin write time rather than on the first // outbound call. -func (c Config) validateTLSMaterial() error { - //nolint:wrapcheck // the message names the subsystem an operator sees refuse the write; wrapping would prefix it twice +func (c Config) ValidateTLSMaterial() error { + //nolint:wrapcheck // the message already names the kind an operator sees refuse the write; wrapping would prefix it twice return apigwtls.Validate(c.tlsMaterial()) } -// buildTLSConfig returns the *tls.Config this connection's transport presents, +// BuildTLSConfig returns the *tls.Config this connection's transport presents, // or nil when it carries neither a client keypair nor a CA bundle. -func buildTLSConfig(c Config) (*tls.Config, error) { - //nolint:wrapcheck // as above: the error is already an operator-facing "apigateway: ..." message +func (c Config) BuildTLSConfig() (*tls.Config, error) { + //nolint:wrapcheck // as above: the error is already an operator-facing, kind-prefixed message return apigwtls.Build(c.tlsMaterial()) } diff --git a/internal/upstreamauth/transport.go b/internal/upstreamauth/transport.go new file mode 100644 index 000000000..05e48d1fa --- /dev/null +++ b/internal/upstreamauth/transport.go @@ -0,0 +1,255 @@ +package upstreamauth + +import ( + "io" + "net" + "net/http" + "strings" + "time" + + "github.com/txn2/mcp-data-platform/internal/membudget" +) + +// AuthorizationHeader is the HTTP header bearer-mode auth populates. +// Named so the same literal is not repeated across the auth dispatch +// and the header-spoof rejection. +const AuthorizationHeader = "Authorization" + +// idleConnectionTimeout caps how long an idle keep-alive connection +// can sit in the pool before being closed. Independent of the +// per-call timeouts; a generous default reduces reconnect churn for +// chatty connections. +const idleConnectionTimeout = 90 * time.Second + +// maxIdleConnections caps the per-host pool of reusable keep-alive +// sockets. Modest because each connection's typical workload is +// occasional fan-out from MCP tool calls, not high-throughput. +const maxIdleConnections = 10 + +// NewHTTPClient builds the per-connection *http.Client: the call +// timeout as the client deadline, and the connection's transport. +// +// Redirects are explicitly disallowed so a kind does not blindly +// re-issue a request (and re-attach the connection's credential) to a +// host the operator did not authorize. The model can follow a redirect +// manually by reading the upstream Location header from the response +// and issuing a new call with the redirected URL. +// +// Metrics wrapping is applied by the caller rather than here so test +// helpers can construct a bare client without threading a metrics +// handle through every call site. +// +// TLS-config build errors are intentionally not surfaced from this +// constructor. Validation has already checked cert + key + CA bundle, +// so BuildTLSConfig only fails here if a caller has constructed a +// Config by hand and bypassed it. The fallback returns a transport +// with the system default tls.Config and the first outbound call will +// fail loudly with the underlying tls error, which is the same surface +// a misconfigured transport would produce on any other auth mode. +func NewHTTPClient(cfg Config) *http.Client { + return &http.Client{ + Timeout: cfg.CallTimeout, + Transport: NewHTTPTransport(cfg), + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + } +} + +// NewHTTPTransport builds the per-connection http.Transport. The +// dial step (TCP + TLS handshake) is bound by cfg.ConnectTimeout so +// an unreachable upstream fails fast instead of consuming the full +// CallTimeout budget. Exposed separately from NewHTTPClient so unit +// tests can verify the wiring without standing up a network listener. +// +// When the connection carries mTLS material (cfg.MTLSClientCertPEM +// + cfg.MTLSClientKeyPEM) or a custom CA bundle (cfg.TLSCABundlePEM), +// the transport's TLSClientConfig is populated accordingly. With +// neither set, TLSClientConfig stays nil and Go's net/http uses +// system defaults. BuildTLSConfig errors here are degraded to nil +// (see NewHTTPClient for the rationale). +func NewHTTPTransport(cfg Config) *http.Transport { + t := &http.Transport{ + DialContext: (&net.Dialer{ + Timeout: cfg.ConnectTimeout, + }).DialContext, + TLSHandshakeTimeout: cfg.ConnectTimeout, + ExpectContinueTimeout: time.Second, + IdleConnTimeout: idleConnectionTimeout, + MaxIdleConns: maxIdleConnections, + } + if tlsCfg, err := cfg.BuildTLSConfig(); err == nil && tlsCfg != nil { + t.TLSClientConfig = tlsCfg + } + return t +} + +// AuthHeader returns the canonical header name this connection's auth +// mode would set, so ValidateCustomHeaders can reject the model's +// attempts to spoof or override it. Empty string means no header-based +// auth (mode=none, mode=api_key with query placement). +func (c Config) AuthHeader() string { + switch c.AuthMode { + case AuthModeBearer: + return AuthorizationHeader + case AuthModeAPIKey: + if c.CredentialPlacement == CredentialPlacementHeader { + return c.APIKeyHeader + } + } + return "" +} + +// ValidateCustomHeaders refuses model-supplied headers that would +// collide with a credential or with an operator-pinned static header. +// The model never gets to set Authorization, the header its own +// connection's auth mode owns, or any name the operator fixed in +// static_headers: those are the operator's to decide, and a header the +// model set would either be silently overwritten at request build time +// or, worse, win. +func (c Config) ValidateCustomHeaders(headers map[string]string) error { + authHeader := c.AuthHeader() + for name := range headers { + if strings.EqualFold(name, AuthorizationHeader) { + return c.err("Authorization header is reserved; configure auth via connection") + } + if authHeader != "" && strings.EqualFold(name, authHeader) { + return c.errf("%s header is reserved by this connection's auth_mode", authHeader) + } + for staticName := range c.StaticHeaders { + if strings.EqualFold(name, staticName) { + return c.errf("%s header is reserved by this connection's static_headers", staticName) + } + } + } + return nil +} + +// ValidateStaticHeaders refuses operator config that would collide with +// the auth path or with hop-by-hop headers Go forbids on a request. A +// static header attempting to set Authorization (or the +// auth-mode-reserved header for api_key+header) would silently lose to +// the auth layer at request time — fail loudly here instead. +func (c Config) ValidateStaticHeaders() error { + if len(c.StaticHeaders) == 0 { + return nil + } + authHeader := c.AuthHeader() + for name, value := range c.StaticHeaders { + if err := c.checkStaticHeader(name, value, authHeader); err != nil { + return err + } + } + return nil +} + +// checkStaticHeader is one static header's worth of ValidateStaticHeaders, +// split out so the loop body's branches stay under the cognitive-complexity +// ceiling. +func (c Config) checkStaticHeader(name, value, authHeader string) error { + if name == "" { + return c.err("static_headers contains an empty header name") + } + if !isValidHeaderName(name) { + return c.errf("static_headers name %q contains characters not permitted in an HTTP header name", name) + } + if strings.ContainsAny(value, "\r\n\x00") { + return c.errf("static_headers[%q] contains CR/LF/NUL — header smuggling vector", name) + } + if strings.EqualFold(name, AuthorizationHeader) { + return c.err("static_headers must not set Authorization; configure auth via auth_mode") + } + if authHeader != "" && strings.EqualFold(name, authHeader) { + return c.errf("static_headers must not set %q — already managed by auth_mode", name) + } + if isReservedHopHeader(name) { + return c.errf("static_headers must not set hop-by-hop or net/http-managed header %q", name) + } + return nil +} + +// isValidHeaderName matches RFC 7230 token chars. Permissive enough for +// real-world headers (x-goog-user-project, X-Subscription-Key) and +// strict enough to refuse spaces / control chars that would let an +// operator inject CRLF via a header name. +func isValidHeaderName(name string) bool { + for i := 0; i < len(name); i++ { + c := name[i] + switch { + case c >= 'A' && c <= 'Z': + case c >= 'a' && c <= 'z': + case c >= '0' && c <= '9': + case strings.ContainsRune("!#$%&'*+-.^_`|~", rune(c)): + default: + return false + } + } + return name != "" +} + +// isReservedHopHeader names headers Go's net/http manages on the +// request itself (Host, Content-Length) or that are meaningless on a +// per-call basis (Connection, Transfer-Encoding, Upgrade). Setting +// these from operator config would either be silently overridden or +// break the transport. +func isReservedHopHeader(name string) bool { + switch strings.ToLower(name) { + case "host", "content-length", "connection", "transfer-encoding", + "upgrade", "keep-alive", "proxy-authenticate", + "proxy-authorization", "te", "trailer": + return true + } + return false +} + +// ReadLimit is the most of a response to read given a configured limit, +// falling back to the default cap when there is none. The one definition +// a buffered call, a page of a walk, and the memory reservation share, +// so the three cannot drift. +func ReadLimit(limit int64) int64 { + if limit > 0 { + return limit + } + return DefaultMaxResponseBytes +} + +// ReadBody reads at most maxBytes of an upstream response, reporting +// whether it was cut short. One extra byte is read so a body exactly at +// the cap is distinguishable from one that overruns it. +func ReadBody(prefix string, r io.Reader, maxBytes int64) (body []byte, truncated bool, err error) { + if maxBytes <= 0 { + maxBytes = DefaultMaxResponseBytes + } + limited := io.LimitReader(r, maxBytes+1) + read, rerr := io.ReadAll(limited) + if rerr != nil { + return nil, false, errf(prefix, "reading response body: %w", rerr) + } + if int64(len(read)) > maxBytes { + return read[:maxBytes], true, nil + } + return read, false, nil +} + +// ReserveBodyBudget computes the worst-case number of bytes a buffered +// read of this response could hold and tries to reserve them against +// the shared budget. It returns the amount reserved (to be released by +// the caller) and whether the reservation was granted. +// +// When the upstream declares a Content-Length below the read cap, only +// that many bytes are reserved so small (and empty) responses do not +// each tie up the full per-request cap and falsely exhaust the budget. +// This is safe because Go's HTTP client bounds resp.Body to the declared +// Content-Length — a server that writes more than it declared cannot +// make ReadBody buffer beyond it. Unknown/chunked responses +// (ContentLength < 0) and over-cap responses reserve the full cap, which +// is exactly what ReadBody may buffer. A nil/disabled budget always +// grants the reservation and Release is a no-op, so the buffered path is +// unchanged when no budget is configured. +func ReserveBodyBudget(b *membudget.Budget, contentLength, readCap int64) (reserved int64, ok bool) { + reserved = readCap + if contentLength >= 0 && contentLength < readCap { + reserved = contentLength + } + return reserved, b.Acquire(reserved) +} diff --git a/internal/upstreamauth/transport_test.go b/internal/upstreamauth/transport_test.go new file mode 100644 index 000000000..2161b14e1 --- /dev/null +++ b/internal/upstreamauth/transport_test.go @@ -0,0 +1,457 @@ +package upstreamauth + +import ( + "context" + "crypto/rand" + "crypto/rsa" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "io" + "math/big" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/txn2/mcp-data-platform/internal/membudget" +) + +func TestConfigAuthHeader(t *testing.T) { + cases := []struct { + name string + cfg Config + want string + }{ + {"none", Config{AuthMode: AuthModeNone}, ""}, + {"bearer", Config{AuthMode: AuthModeBearer}, "Authorization"}, + {"api_key header default", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: DefaultAPIKeyHeader}, DefaultAPIKeyHeader}, + {"api_key header custom", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: "X-My-Key"}, "X-My-Key"}, + {"api_key query", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementQuery, APIKeyParam: "key"}, ""}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + if got := tc.cfg.AuthHeader(); got != tc.want { + t.Errorf("AuthHeader = %q; want %q", got, tc.want) + } + }) + } +} + +func TestValidateCustomHeaders_RejectsAuthorization(t *testing.T) { + err := Config{}.ValidateCustomHeaders(map[string]string{"AUTHORIZATION": "anything"}) + if err == nil { + t.Error("Authorization header allowed") + } +} + +func TestValidateCustomHeaders_RejectsConfiguredAPIKeyHeader(t *testing.T) { + cfg := Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: "X-Custom-Key"} + err := cfg.ValidateCustomHeaders(map[string]string{"x-custom-key": "spoof"}) + if err == nil { + t.Error("configured api_key header allowed (case-insensitive check failed)") + } +} + +func TestValidateCustomHeaders_AllowsOtherHeaders(t *testing.T) { + cfg := Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: DefaultAPIKeyHeader} + err := cfg.ValidateCustomHeaders(map[string]string{"Accept-Language": "en"}) + if err != nil { + t.Errorf("unrelated header rejected: %v", err) + } +} + +func TestValidateCustomHeaders_RejectsStaticHeaderOverride(t *testing.T) { + cfg := Config{StaticHeaders: map[string]string{"X-Goog-User-Project": "secret"}} + err := cfg.ValidateCustomHeaders(map[string]string{"x-goog-user-project": "spoof"}) + if err == nil { + t.Error("model attempt to override static header allowed (case-insensitive check failed)") + } +} + +func TestIsValidHeaderName(t *testing.T) { + cases := []struct { + in string + want bool + }{ + {"X-API-Key", true}, + {"X-Api-Version2", true}, + {"X-Subscription-Key", true}, + {"x-goog-user-project", true}, + {"Content-Type", true}, + {"", false}, + {"Bad Name", false}, + {"With\rCR", false}, + {"Colon:Inside", false}, + } + for _, tc := range cases { + t.Run(tc.in, func(t *testing.T) { + if got := isValidHeaderName(tc.in); got != tc.want { + t.Errorf("isValidHeaderName(%q) = %v; want %v", tc.in, got, tc.want) + } + }) + } +} + +// TestNewTokenExchangeClient_BadBundleFallsBackQuietly is the +// resilience contract: a CA bundle that fails to parse at runtime +// (impossible if Validate ran but possible if a caller bypassed it) +// must NOT panic or block token fetches with a nil transport. The +// fallback is a plain http.Client without the bundle, matching the +// pre-feature behavior; the request will then fail with a TLS error +// against the IdP and the operator gets a normal error path. +func TestNewTokenExchangeClient_BadBundleFallsBackQuietly(t *testing.T) { + client := newTokenExchangeClient(Config{TLSCABundlePEM: "not pem"}) + require.NotNil(t, client) + assert.Nil(t, client.Transport, "fallback must not attach a half-built transport") +} + +// TestNewTokenExchangeClient_HonorsCABundle exercises the IdP-side CA +// trust plumbing for oauth2_client_credentials: when the IdP is +// signed by a private CA in tls_ca_bundle_pem, the token-fetch must +// succeed. The negative branch (no bundle) is implicit: without the +// trust the default RoundTripper would reject the IdP's cert. +func TestNewTokenExchangeClient_HonorsCABundle(t *testing.T) { + ca := newTestCA(t) + idpCert, idpKey := ca.issueServerCert(t, "127.0.0.1") + srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"access_token":"abc","token_type":"bearer","expires_in":3600}`) + })) + srv.TLS = &tls.Config{ + MinVersion: tls.VersionTLS12, + Certificates: []tls.Certificate{mustKeyPair(t, idpCert, idpKey)}, + } + srv.StartTLS() + defer srv.Close() + + cfg := Config{TLSCABundlePEM: ca.certPEM} + client := newTokenExchangeClient(cfg) + postReq, err := http.NewRequestWithContext(context.Background(), + http.MethodPost, srv.URL+"/token", + strings.NewReader("")) + require.NoError(t, err) + postReq.Header.Set("Content-Type", "application/x-www-form-urlencoded") + resp, err := client.Do(postReq) + require.NoError(t, err) + defer func() { _ = resp.Body.Close() }() + assert.Equal(t, http.StatusOK, resp.StatusCode) + body, _ := io.ReadAll(resp.Body) + assert.Contains(t, string(body), "access_token") +} + +// --- test CA helper -------------------------------------------------- + +// testCA is a minimal single-cert CA built per test. It exists so the +// token-exchange tests can stand up an IdP whose certificate chains to +// a bundle this package's config carries and to nothing the host +// trusts, which is the only way to prove the bundle is what made the +// handshake succeed. +type testCA struct { + certPEM string + cert *x509.Certificate + key *rsa.PrivateKey +} + +func newTestCA(t *testing.T) *testCA { + t.Helper() + key, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(1), + Subject: pkix.Name{CommonName: "upstreamauth-test-ca"}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature, + BasicConstraintsValid: true, + IsCA: true, + } + der, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key) + require.NoError(t, err) + cert, err := x509.ParseCertificate(der) + require.NoError(t, err) + return &testCA{ + certPEM: string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})), + cert: cert, + key: key, + } +} + +// issueServerCert mints a leaf certificate for the given IP SAN, +// signed by this CA. +func (ca *testCA) issueServerCert(t *testing.T, ip string) (certPEM, keyPEM string) { + t.Helper() + key, err := rsa.GenerateKey(rand.Reader, 2048) + require.NoError(t, err) + tmpl := &x509.Certificate{ + SerialNumber: big.NewInt(2), + Subject: pkix.Name{CommonName: ip}, + NotBefore: time.Now().Add(-time.Hour), + NotAfter: time.Now().Add(24 * time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + IPAddresses: []net.IP{net.ParseIP(ip)}, + } + der, err := x509.CreateCertificate(rand.Reader, tmpl, ca.cert, &key.PublicKey, ca.key) + require.NoError(t, err) + keyDER, err := x509.MarshalPKCS8PrivateKey(key) + require.NoError(t, err) + return string(pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})), + string(pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER})) +} + +func mustKeyPair(t *testing.T, certPEM, keyPEM string) tls.Certificate { + t.Helper() + pair, err := tls.X509KeyPair([]byte(certPEM), []byte(keyPEM)) + require.NoError(t, err) + return pair +} + +// --- transport policy ------------------------------------------------ + +// TestValidateStaticHeaders walks the operator-config rules. Each +// refusal is a header that would either be silently overridden by the +// auth layer, be managed by net/http, or smuggle a second header +// through a CR/LF payload — all cases where accepting the config and +// failing later would be worse than refusing the save. +func TestValidateStaticHeaders(t *testing.T) { + apiKeyCfg := func(headers map[string]string) Config { + return Config{ + ErrPrefix: "apigateway", + AuthMode: AuthModeAPIKey, + Credential: "k", + CredentialPlacement: CredentialPlacementHeader, + APIKeyHeader: "X-Custom-Key", + StaticHeaders: headers, + } + } + cases := []struct { + name string + cfg Config + wantMsg string // empty = must pass + }{ + {name: "no static headers", cfg: Config{}}, + {name: "ordinary header", cfg: apiKeyCfg(map[string]string{"X-Goog-User-Project": "proj"})}, + {name: "empty name", cfg: apiKeyCfg(map[string]string{"": "v"}), wantMsg: "empty header name"}, + {name: "invalid name", cfg: apiKeyCfg(map[string]string{"Bad Name": "v"}), wantMsg: "not permitted in an HTTP header name"}, + {name: "CRLF in value", cfg: apiKeyCfg(map[string]string{"X-Ok": "a\r\nX-Smuggled: 1"}), wantMsg: "header smuggling vector"}, + {name: "Authorization", cfg: apiKeyCfg(map[string]string{"authorization": "Bearer x"}), wantMsg: "must not set Authorization"}, + {name: "the auth mode's own header", cfg: apiKeyCfg(map[string]string{"x-custom-key": "spoof"}), wantMsg: "already managed by auth_mode"}, + {name: "hop-by-hop header", cfg: apiKeyCfg(map[string]string{"Connection": "close"}), wantMsg: "hop-by-hop"}, + {name: "net/http-managed header", cfg: apiKeyCfg(map[string]string{"Content-Length": "10"}), wantMsg: "net/http-managed"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + err := tc.cfg.ValidateStaticHeaders() + if tc.wantMsg == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + assert.Contains(t, err.Error(), tc.wantMsg) + }) + } +} + +func TestIsReservedHopHeader(t *testing.T) { + for _, name := range []string{ + "host", "Content-Length", "CONNECTION", "transfer-encoding", "Upgrade", + "keep-alive", "proxy-authenticate", "proxy-authorization", "te", "trailer", + } { + if !isReservedHopHeader(name) { + t.Errorf("isReservedHopHeader(%q) = false; want true", name) + } + } + for _, name := range []string{"X-API-Key", "Accept", "x-goog-user-project"} { + if isReservedHopHeader(name) { + t.Errorf("isReservedHopHeader(%q) = true; want false", name) + } + } +} + +func TestReadLimit(t *testing.T) { + assert.Equal(t, int64(4096), ReadLimit(4096)) + assert.Equal(t, DefaultMaxResponseBytes, ReadLimit(0), "unset limit takes the default cap") + assert.Equal(t, DefaultMaxResponseBytes, ReadLimit(-1), "a negative limit must not read nothing") +} + +// TestReadBody covers the cap's three outcomes: a body under it comes +// back whole, a body over it comes back cut and flagged, and a reader +// that fails surfaces the error in the calling kind's voice. +func TestReadBody(t *testing.T) { + t.Run("under the cap", func(t *testing.T) { + body, truncated, err := ReadBody("apigateway", strings.NewReader("hello"), 1024) + require.NoError(t, err) + assert.False(t, truncated) + assert.Equal(t, "hello", string(body)) + }) + t.Run("exactly at the cap is not truncated", func(t *testing.T) { + body, truncated, err := ReadBody("apigateway", strings.NewReader("hello"), 5) + require.NoError(t, err) + assert.False(t, truncated) + assert.Equal(t, "hello", string(body)) + }) + t.Run("over the cap", func(t *testing.T) { + body, truncated, err := ReadBody("apigateway", strings.NewReader("hello world"), 5) + require.NoError(t, err) + assert.True(t, truncated) + assert.Equal(t, "hello", string(body)) + }) + t.Run("zero cap takes the default", func(t *testing.T) { + body, truncated, err := ReadBody("apigateway", strings.NewReader("hello"), 0) + require.NoError(t, err) + assert.False(t, truncated) + assert.Equal(t, "hello", string(body)) + }) + t.Run("read error", func(t *testing.T) { + _, _, err := ReadBody("apigateway", failingReader{}, 1024) + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: reading response body") + }) +} + +// failingReader fails on the first Read so ReadBody's error arm is +// reachable without a live connection. +type failingReader struct{} + +func (failingReader) Read([]byte) (int, error) { return 0, io.ErrUnexpectedEOF } + +// TestReserveBodyBudget pins the reservation rule that keeps a small +// response from tying up the full per-request cap: a declared +// Content-Length under the cap reserves only that much, while an +// unknown length reserves the whole cap because that is what the read +// may buffer. +func TestReserveBodyBudget(t *testing.T) { + cases := []struct { + name string + contentLength int64 + readCap int64 + want int64 + }{ + {"declared length under the cap", 100, 1024, 100}, + {"declared length over the cap", 4096, 1024, 1024}, + {"declared empty body", 0, 1024, 0}, + {"unknown length reserves the cap", -1, 1024, 1024}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + budget := membudget.New(1 << 20) + reserved, ok := ReserveBodyBudget(budget, tc.contentLength, tc.readCap) + assert.True(t, ok) + assert.Equal(t, tc.want, reserved) + budget.Release(reserved) + }) + } + t.Run("refused when the budget is exhausted", func(t *testing.T) { + budget := membudget.New(512) + _, ok := ReserveBodyBudget(budget, -1, 1024) + assert.False(t, ok) + }) + t.Run("a nil budget always grants", func(t *testing.T) { + reserved, ok := ReserveBodyBudget(nil, -1, 1024) + assert.True(t, ok) + assert.Equal(t, int64(1024), reserved) + }) +} + +// TestNewHTTPClient_WiresTimeoutsAndRefusesRedirects covers the +// transport policy a kind gets for free: the call timeout as the client +// deadline, the connect timeout on the dial and handshake, and a +// redirect that is returned rather than followed — the rule that stops +// a 302 from carrying the connection's credential to a host the +// operator never configured. +func TestNewHTTPClient_WiresTimeoutsAndRefusesRedirects(t *testing.T) { + cfg := Config{ConnectTimeout: 1500 * time.Millisecond, CallTimeout: 30 * time.Second} + client := NewHTTPClient(cfg) + assert.Equal(t, cfg.CallTimeout, client.Timeout) + require.NotNil(t, client.Transport, "a nil transport would silently fall back to http.DefaultTransport") + + tr, ok := client.Transport.(*http.Transport) + require.True(t, ok) + assert.Equal(t, cfg.ConnectTimeout, tr.TLSHandshakeTimeout) + require.NotNil(t, tr.DialContext, "DialContext is nil; ConnectTimeout cannot be enforced") + + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/from" { + http.Redirect(w, r, "/to", http.StatusFound) + return + } + _, _ = io.WriteString(w, "followed") + })) + defer srv.Close() + + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, srv.URL+"/from", http.NoBody) + require.NoError(t, err) + resp, err := client.Do(req) + require.NoError(t, err) + defer func() { _ = resp.Body.Close() }() + assert.Equal(t, http.StatusFound, resp.StatusCode, "the redirect must be returned, not followed") +} + +// TestBuildTLSConfig covers the three shapes a connection's TLS +// material takes: none (so net/http keeps its defaults), a CA bundle +// alone, and a client keypair. +func TestBuildTLSConfig(t *testing.T) { + ca := newTestCA(t) + certPEM, keyPEM := ca.issueServerCert(t, "127.0.0.1") + + t.Run("no material yields no config", func(t *testing.T) { + got, err := Config{}.BuildTLSConfig() + require.NoError(t, err) + assert.Nil(t, got, "nil lets http.Transport use its defaults") + }) + t.Run("CA bundle only", func(t *testing.T) { + got, err := Config{TLSCABundlePEM: ca.certPEM}.BuildTLSConfig() + require.NoError(t, err) + require.NotNil(t, got) + assert.NotNil(t, got.RootCAs) + assert.Empty(t, got.Certificates) + }) + t.Run("client keypair", func(t *testing.T) { + got, err := Config{MTLSClientCertPEM: certPEM, MTLSClientKeyPEM: keyPEM}.BuildTLSConfig() + require.NoError(t, err) + require.NotNil(t, got) + assert.Len(t, got.Certificates, 1) + }) + t.Run("unusable keypair is refused in the kind's voice", func(t *testing.T) { + _, err := Config{ + ErrPrefix: "graphql", + MTLSClientCertPEM: "not pem", + MTLSClientKeyPEM: "not pem", + }.BuildTLSConfig() + require.Error(t, err) + assert.Contains(t, err.Error(), "graphql: building mtls keypair") + }) +} + +// TestValidateTLSMaterial_RequiresThePairUnderMTLS proves the auth mode +// reaches the TLS validator: under auth_mode=mtls the client keypair is +// the credential, so its absence is a refused connection rather than a +// handshake failure on the first call. +func TestValidateTLSMaterial_RequiresThePairUnderMTLS(t *testing.T) { + err := Config{ErrPrefix: "apigateway", AuthMode: AuthModeMTLS}.ValidateTLSMaterial() + require.Error(t, err) + assert.Contains(t, err.Error(), "apigateway: mtls_client_cert_pem and mtls_client_key_pem are required") + + require.NoError(t, Config{ErrPrefix: "apigateway", AuthMode: AuthModeNone}.ValidateTLSMaterial()) +} + +// TestNewHTTPTransport_AttachesTheConnectionsTLSConfig covers the arm +// that plumbs a connection's CA bundle into the transport. Without it +// an upstream behind a private CA would fail its handshake even though +// the operator configured the bundle. +func TestNewHTTPTransport_AttachesTheConnectionsTLSConfig(t *testing.T) { + ca := newTestCA(t) + tr := NewHTTPTransport(Config{TLSCABundlePEM: ca.certPEM}) + require.NotNil(t, tr.TLSClientConfig) + assert.NotNil(t, tr.TLSClientConfig.RootCAs) + + bare := NewHTTPTransport(Config{}) + assert.Nil(t, bare.TLSClientConfig, "with no material net/http must keep its own defaults") +} diff --git a/pkg/toolkits/apigateway/config.go b/pkg/toolkits/apigateway/config.go index 02d2b7095..2bef3238c 100644 --- a/pkg/toolkits/apigateway/config.go +++ b/pkg/toolkits/apigateway/config.go @@ -14,14 +14,10 @@ import ( "errors" "fmt" "log/slog" - "maps" - "strconv" - "strings" "time" - "golang.org/x/oauth2" - - "github.com/txn2/mcp-data-platform/pkg/connoauth" + "github.com/txn2/mcp-data-platform/internal/cfgmap" + "github.com/txn2/mcp-data-platform/internal/upstreamauth" ) const ( @@ -29,85 +25,6 @@ const ( // this in the admin UI's connection picker. Kind = "api" - // AuthModeNone disables outbound authentication. - AuthModeNone = "none" - // AuthModeBearer sends "Authorization: Bearer ". - AuthModeBearer = "bearer" - // AuthModeAPIKey sends the credential as a header (default - // "X-API-Key") or as a query parameter; placement and key name are - // per-connection so APIs that use non-standard schemes (e.g. an - // "api_key" query parameter, or a custom "X-Api-Token" header) can - // be onboarded without code changes. - AuthModeAPIKey = "api_key" - // AuthModeBasic sends "Authorization: Basic base64(username:password)" - // per RFC 7617. Required for the long tail of older REST APIs (Jenkins, - // on-prem Jira / Confluence Server / DC, internal apps) that never moved - // to bearer or OAuth. RFC 7617 §2 forbids ":" in the userid; password - // may be empty (some APIs accept "token:" as a bearer-token-in-username - // pattern). Password is encrypted at rest via the platform's - // FieldEncryptor (the "password" config key is already in the - // sensitive-keys list). - AuthModeBasic = "basic" - - // AuthModeOAuth is the canonical OAuth auth_mode shared across - // every toolkit kind (see connoauth.AuthModeOAuth). The specific - // flow is carried separately in OAuth2Config.Grant. ParseConfig - // normalizes the legacy api-only auth_mode values below to this - // form, so a parsed Config always reports AuthModeOAuth for an - // OAuth connection. - AuthModeOAuth = connoauth.AuthModeOAuth - - // AuthModeOAuth2ClientCredentials is the legacy api-only auth_mode - // that encoded the client_credentials grant in the mode string. - // Retained so raw config authored before the schema unified (and - // hand-built test Configs) still parse; ParseConfig normalizes it - // to AuthModeOAuth + Grant=client_credentials. - // - // client_credentials acquires a bearer token — server-to-server, - // no human in the loop. The platform exchanges the configured - // client_id + client_secret for a token at OAuth.TokenURL and - // applies it as "Authorization: Bearer " on outbound - // calls. Tokens are cached + refreshed automatically by the - // underlying golang.org/x/oauth2 library; no DB state is - // required because every restart can re-acquire from credentials. - // The authorization_code grant (which DOES require DB-persisted - // refresh tokens + a browser flow) is its own follow-up issue. - AuthModeOAuth2ClientCredentials = "oauth2_client_credentials" // #nosec G101 -- mode name, not a credential - - // AuthModeOAuth2AuthorizationCode runs the user-driven OAuth 2.1 - // authorization-code grant: an admin completes a one-time browser - // flow at connection setup; the resulting refresh token is - // persisted (encrypted) so subsequent platform restarts and - // background workloads keep working without further interaction. - // Tokens are refreshed automatically before expiry. Requires the - // platform's database (refresh-token state survives restarts). - AuthModeOAuth2AuthorizationCode = "oauth2_authorization_code" // #nosec G101 -- mode name, not a credential - - // AuthModeMTLS authenticates by presenting an X.509 client - // certificate during the TLS handshake per RFC 5246 / 8446. No - // Authorization header is sent: the cert IS the credential. - // Used by upstreams that map the cert's subject DN (or a SAN) to - // a user identity in their authorizer, including service-mesh - // peers, PKI-fronted internal APIs, healthcare integration - // engines, financial messaging endpoints, and FedRAMP / DoD- - // boundary services. Requires both mtls_client_cert_pem and - // mtls_client_key_pem on the connection config; the mTLS material - // can also be present alongside other auth modes (bearer + mTLS, - // etc.), but auth_mode=mtls is the explicit "no header - // credential" signal. - AuthModeMTLS = "mtls" - - // CredentialPlacementHeader (default) sends the credential as an HTTP - // header named by APIKeyHeader. - CredentialPlacementHeader = "header" - // CredentialPlacementQuery sends the credential as a URL query parameter - // named by APIKeyParam. - CredentialPlacementQuery = "query" - - // DefaultAPIKeyHeader is the conventional API-key header name when - // the connection does not specify one. - DefaultAPIKeyHeader = "X-API-Key" // #nosec G101 -- header name, not a credential - // TrustLevelUntrusted is the default. Advisory only in v1: the // field is parsed and validated but no platform code reads it. // Reserved for future response-shaping enforcement (see issue @@ -119,16 +36,16 @@ const ( // DefaultConnectTimeout caps the time spent establishing the // outbound connection (TCP + TLS handshake) on each invocation. - DefaultConnectTimeout = 10 * time.Second + DefaultConnectTimeout = upstreamauth.DefaultConnectTimeout // DefaultCallTimeout caps the total per-call time including // upstream processing and response read. - DefaultCallTimeout = 60 * time.Second + DefaultCallTimeout = upstreamauth.DefaultCallTimeout // DefaultMaxResponseBytes is the upstream read cap: the most the // gateway reads of any one response (a page of a walk, an inline // call). It bounds transfer and buffering, not what reaches the // model; that is MaxInlineBytes. - DefaultMaxResponseBytes = int64(10 * 1024 * 1024) + DefaultMaxResponseBytes = upstreamauth.DefaultMaxResponseBytes // DefaultMaxInlineBytes is the inline budget: the most a rendered // api_invoke_endpoint tool result may hold. It is a model-context @@ -153,32 +70,44 @@ const ( DefaultMaxInlineBytes = int64(32 * 1024) ) +// The credential vocabulary an operator configures on an api connection. +// The values are defined by internal/upstreamauth, which owns the +// outbound authentication policy for every HTTP-based connection kind; +// they are aliased here because they are part of this toolkit's public +// API and of what an operator types into auth_mode. +const ( + AuthModeNone = upstreamauth.AuthModeNone + AuthModeBearer = upstreamauth.AuthModeBearer + AuthModeAPIKey = upstreamauth.AuthModeAPIKey + AuthModeBasic = upstreamauth.AuthModeBasic + AuthModeOAuth = upstreamauth.AuthModeOAuth + AuthModeOAuth2ClientCredentials = upstreamauth.AuthModeOAuth2ClientCredentials + AuthModeOAuth2AuthorizationCode = upstreamauth.AuthModeOAuth2AuthorizationCode + AuthModeMTLS = upstreamauth.AuthModeMTLS + + CredentialPlacementHeader = upstreamauth.CredentialPlacementHeader + CredentialPlacementQuery = upstreamauth.CredentialPlacementQuery + + DefaultAPIKeyHeader = upstreamauth.DefaultAPIKeyHeader +) + // cfgKey* constants name the keys used to read a Config from a // map[string]any (the form connections take in the platform's generic // connection_instances store). const ( - cfgKeyBaseURL = "base_url" - cfgKeyAuthMode = "auth_mode" - cfgKeyCredential = "credential" // #nosec G101 -- map key, not a secret - cfgKeyAPIKeyHeader = "api_key_header" // #nosec G101 -- map key, not a credential - cfgKeyAPIKeyParam = "api_key_param" // #nosec G101 -- map key, not a credential - cfgKeyAPIKeyPlacement = "api_key_placement" - cfgKeyUsername = "username" - cfgKeyPassword = "password" // #nosec G101 -- map key, not a credential; encryption handled by platform FieldEncryptor sensitive-keys list - cfgKeyConnectTimeout = "connect_timeout" - cfgKeyCallTimeout = "call_timeout" - cfgKeyTrustLevel = "trust_level" - cfgKeyMaxResponseBytes = "max_response_bytes" - cfgKeyMaxInlineBytes = "max_inline_bytes" - // cfgKeyStaticHeaders holds operator-configured headers that the - // toolkit appends to every outbound request. Required for upstreams - // that demand BOTH an Authorization bearer AND a separate - // subscription/key header (Google Cloud's x-goog-user-project - // quota-billing header, vendor subscription keys, etc.). Stored - // as a map[string]any whose values are encrypted at rest by - // platform.FieldEncryptor (see CfgKeyStaticHeaders in - // pkg/platform/fieldcrypt.go). - cfgKeyStaticHeaders = "static_headers" + cfgKeyBaseURL = "base_url" + cfgKeyTrustLevel = "trust_level" + cfgKeyMaxInlineBytes = "max_inline_bytes" + + // The credential, timeout, response-cap, static-header and TLS + // keys are read by internal/upstreamauth from the same config map: + // auth_mode, credential, api_key_header, api_key_param, + // api_key_placement, username, password, connect_timeout, + // call_timeout, max_response_bytes, static_headers, + // mtls_client_cert_pem, mtls_client_key_pem, tls_ca_bundle_pem and + // identity_passthrough. They are not re-declared here so the two + // packages cannot come to disagree about a key's spelling. + // cfgKeyCatalogID names the api_catalogs row that supplies this // connection's OpenAPI specs. Empty = connection has no spec // surface (api_discover answers with a note and no operations, @@ -188,25 +117,6 @@ const ( // of duplicating the documentation. cfgKeyCatalogID = "catalog_id" - // OAuth config keys are owned by pkg/connoauth (the canonical - // oauth_* vocabulary plus the legacy oauth2_* fallback). This - // toolkit delegates OAuth parsing to connoauth.ParseConfig rather - // than declaring its own key constants. The keys remain top-level - // (not nested) so the platform's FieldEncryptor, which walks only - // the top level of the config map, encrypts oauth_client_secret at - // rest without changes to the encryptor. - - // mTLS material config keys. Cert and CA bundle are public - // material (plain text at rest); the private key is in the - // platform's sensitive-keys list (see pkg/platform/fieldcrypt.go) - // and encrypted via FieldEncryptor like every other secret on a - // connection. - cfgKeyMTLSClientCertPEM = "mtls_client_cert_pem" // #nosec G101 -- map key, not a credential - cfgKeyMTLSClientKeyPEM = "mtls_client_key_pem" // #nosec G101 -- map key, not a credential - cfgKeyTLSCABundlePEM = "tls_ca_bundle_pem" - - cfgKeyIdentityPassthrough = "identity_passthrough" - // cfgKeyDescription is an optional human-readable description of the // connection, surfaced via ListConnections (and thus the admin UI and // the list_connections MCP tool). When unset, ListConnections falls @@ -402,10 +312,11 @@ type OAuth2Config struct { Prompt string } -// EndpointAuthStyle values. +// EndpointAuthStyle values, owned by internal/upstreamauth and aliased +// here because they are part of this toolkit's public API. const ( - OAuth2AuthStyleHeader = "header" - OAuth2AuthStyleParams = "params" + OAuth2AuthStyleHeader = upstreamauth.OAuth2AuthStyleHeader + OAuth2AuthStyleParams = upstreamauth.OAuth2AuthStyleParams ) // MultiConfig holds parsed per-connection configs plus the aggregate @@ -438,53 +349,23 @@ func ParseMultiConfig(defaultName string, raw map[string]map[string]any) (MultiC // ParseConfig parses a Config from a generic map (the form admin-saved // connections take in the connection_instances table) and applies -// defaults. The returned Config is fully validated. +// defaults. The credential, timeout, response-cap, static-header and +// TLS keys are read by internal/upstreamauth from the same map; the +// keys below are this toolkit's own. The returned Config is fully +// validated. func ParseConfig(cfg map[string]any) (Config, error) { - c := Config{ - AuthMode: AuthModeNone, - CredentialPlacement: CredentialPlacementHeader, - APIKeyHeader: DefaultAPIKeyHeader, - ConnectTimeout: DefaultConnectTimeout, - CallTimeout: DefaultCallTimeout, - TrustLevel: TrustLevelUntrusted, - MaxResponseBytes: DefaultMaxResponseBytes, - MaxInlineBytes: DefaultMaxInlineBytes, - } - - c.BaseURL = trimTrailingSlash(getString(cfg, cfgKeyBaseURL)) - c.AuthMode = getStringDefault(cfg, cfgKeyAuthMode, c.AuthMode) - c.Credential = getString(cfg, cfgKeyCredential) - c.CredentialPlacement = getStringDefault(cfg, cfgKeyAPIKeyPlacement, c.CredentialPlacement) - c.APIKeyHeader = getStringDefault(cfg, cfgKeyAPIKeyHeader, c.APIKeyHeader) - c.APIKeyParam = getString(cfg, cfgKeyAPIKeyParam) - c.Username = getString(cfg, cfgKeyUsername) - c.Password = getString(cfg, cfgKeyPassword) - c.ConnectTimeout = getDuration(cfg, cfgKeyConnectTimeout, c.ConnectTimeout) - c.CallTimeout = getDuration(cfg, cfgKeyCallTimeout, c.CallTimeout) - c.TrustLevel = getStringDefault(cfg, cfgKeyTrustLevel, c.TrustLevel) - c.MaxResponseBytes = getInt64(cfg, cfgKeyMaxResponseBytes, c.MaxResponseBytes) - c.MaxInlineBytes = getInt64(cfg, cfgKeyMaxInlineBytes, c.MaxInlineBytes) - c.CatalogID = getString(cfg, cfgKeyCatalogID) - if isOAuthAuthMode(c.AuthMode) { - // Delegate OAuth parsing to the shared connoauth.ParseConfig - // (canonical oauth_* keys, legacy oauth2_* fallback, grant - // derivation) and normalize the auth_mode to the canonical - // AuthModeOAuth so the authenticator and validation dispatch on - // the grant rather than on three divergent mode strings. - parsed, err := connoauth.ParseConfig(Kind, getString(cfg, cfgKeyBaseURL), cfg) - if err != nil { - return Config{}, fmt.Errorf("apigateway: %w", err) - } - c.AuthMode = AuthModeOAuth - c.OAuth2 = oauth2ConfigFromConnoauth(parsed) + up, err := upstreamauth.Parse(Kind, upstreamErrPrefix, cfgmap.String(cfg, cfgKeyBaseURL), cfg) + if err != nil { + //nolint:wrapcheck // the seam builds its messages with this toolkit's prefix; wrapping would state it twice + return Config{}, err } - c.StaticHeaders = getStringMap(cfg, cfgKeyStaticHeaders) - c.MTLSClientCertPEM = getString(cfg, cfgKeyMTLSClientCertPEM) - c.MTLSClientKeyPEM = getString(cfg, cfgKeyMTLSClientKeyPEM) - c.TLSCABundlePEM = getString(cfg, cfgKeyTLSCABundlePEM) - c.IdentityPassthrough = getBool(cfg, cfgKeyIdentityPassthrough) - c.Description = getString(cfg, cfgKeyDescription) - c.Handler = getString(cfg, cfgKeyHandler) + c := configFromUpstream(up) + c.BaseURL = trimTrailingSlash(cfgmap.String(cfg, cfgKeyBaseURL)) + c.TrustLevel = cfgmap.StringDefault(cfg, cfgKeyTrustLevel, TrustLevelUntrusted) + c.MaxInlineBytes = cfgmap.Int64(cfg, cfgKeyMaxInlineBytes, DefaultMaxInlineBytes) + c.CatalogID = cfgmap.String(cfg, cfgKeyCatalogID) + c.Description = cfgmap.String(cfg, cfgKeyDescription) + c.Handler = cfgmap.String(cfg, cfgKeyHandler) if c.Handler == HandlerInternal && c.BaseURL == "" { c.BaseURL = internalBaseURL } @@ -496,12 +377,18 @@ func ParseConfig(cfg map[string]any) (Config, error) { } // Validate returns an error if the configuration is missing required -// fields or contains invalid values. +// fields or contains invalid values. The auth, timeout, static-header, +// passthrough and TLS rules belong to internal/upstreamauth; they are +// called individually rather than through its Validate so this +// toolkit's own checks stay interleaved in the order an operator has +// always seen them. func (c Config) Validate() error { + up := c.upstream() if c.BaseURL == "" { return errors.New("apigateway: base_url is required") } - if err := c.validateAuth(); err != nil { + if err := up.ValidateAuth(); err != nil { + //nolint:wrapcheck // as above: already an operator-facing "apigateway: ..." message return err } switch c.TrustLevel { @@ -509,23 +396,18 @@ func (c Config) Validate() error { default: return fmt.Errorf("apigateway: invalid trust_level %q (want untrusted or trusted)", c.TrustLevel) } - if c.ConnectTimeout <= 0 { - return errors.New("apigateway: connect_timeout must be positive") - } - if c.CallTimeout <= 0 { - return errors.New("apigateway: call_timeout must be positive") - } - if c.MaxResponseBytes <= 0 { - return errors.New("apigateway: max_response_bytes must be positive") + if err := up.ValidateTransport(); err != nil { + //nolint:wrapcheck // as above + return err } if c.MaxInlineBytes <= 0 { return errors.New("apigateway: max_inline_bytes must be positive") } return firstConfigError( - c.validateStaticHeaders, - c.validateIdentityPassthrough, + up.ValidateStaticHeaders, + up.ValidateIdentityPassthrough, c.validateHandler, - c.validateTLSMaterial, + up.ValidateTLSMaterial, ) } @@ -566,278 +448,13 @@ func (c Config) validateHandler() error { return nil } -// validateIdentityPassthrough enforces that a passthrough connection -// carries no shared credential. Passthrough forwards the caller's inbound -// token as the Authorization header, so a configured auth_mode would -// either be ignored (confusing) or fight for the same header. Requiring -// auth_mode=none keeps the single-credential-source invariant explicit. -func (c Config) validateIdentityPassthrough() error { - if c.IdentityPassthrough && c.AuthMode != AuthModeNone { - return fmt.Errorf("apigateway: identity_passthrough requires auth_mode=none, got %q", c.AuthMode) - } - return nil -} - -// validateStaticHeaders refuses operator config that would collide with -// the toolkit's auth path or with hop-by-hop headers Go forbids on a -// request. A static header attempting to set Authorization (or the -// auth-mode-reserved header for api_key+header) would silently lose to -// the auth layer at request time — fail loudly here instead. -func (c Config) validateStaticHeaders() error { - if len(c.StaticHeaders) == 0 { - return nil - } - authHeader := authHeaderForConfig(c) - for name, value := range c.StaticHeaders { - if name == "" { - return errors.New("apigateway: static_headers contains an empty header name") - } - if !isValidHeaderName(name) { - return fmt.Errorf("apigateway: static_headers name %q contains characters not permitted in an HTTP header name", name) - } - if strings.ContainsAny(value, "\r\n\x00") { - return fmt.Errorf("apigateway: static_headers[%q] contains CR/LF/NUL — header smuggling vector", name) - } - if strings.EqualFold(name, authorizationHeader) { - return errors.New("apigateway: static_headers must not set Authorization; configure auth via auth_mode") - } - if authHeader != "" && strings.EqualFold(name, authHeader) { - return fmt.Errorf("apigateway: static_headers must not set %q — already managed by auth_mode", name) - } - if isReservedHopHeader(name) { - return fmt.Errorf("apigateway: static_headers must not set hop-by-hop or net/http-managed header %q", name) - } - } - return nil -} - -// isValidHeaderName matches RFC 7230 token chars. Permissive enough for -// real-world headers (x-goog-user-project, X-Subscription-Key) and -// strict enough to refuse spaces / control chars that would let an -// operator inject CRLF via a header name. -func isValidHeaderName(name string) bool { - for i := 0; i < len(name); i++ { - c := name[i] - switch { - case c >= 'A' && c <= 'Z': - case c >= 'a' && c <= 'z': - case c >= '0' && c <= '9': - case strings.ContainsRune("!#$%&'*+-.^_`|~", rune(c)): - default: - return false - } - } - return name != "" -} - -// isReservedHopHeader names headers Go's net/http manages on the -// request itself (Host, Content-Length) or that are meaningless on a -// per-call basis (Connection, Transfer-Encoding, Upgrade). Setting -// these from operator config would either be silently overridden or -// break the transport. -func isReservedHopHeader(name string) bool { - switch strings.ToLower(name) { - case "host", "content-length", "connection", "transfer-encoding", - "upgrade", "keep-alive", "proxy-authenticate", - "proxy-authorization", "te", "trailer": - return true - } - return false -} - -func (c Config) validateAuth() error { - switch c.AuthMode { - case AuthModeNone: - return nil - case AuthModeBearer: - if c.Credential == "" { - return errors.New("apigateway: credential is required when auth_mode is \"bearer\"") - } - return nil - case AuthModeAPIKey: - return c.validateAPIKeyAuth() - case AuthModeBasic: - return c.validateBasicAuth() - case AuthModeOAuth: - if c.OAuth2.Grant == connoauth.GrantAuthorizationCode { - return c.validateOAuth2AuthCode() - } - return c.validateOAuth2() - case AuthModeOAuth2ClientCredentials: - return c.validateOAuth2() - case AuthModeOAuth2AuthorizationCode: - return c.validateOAuth2AuthCode() - case AuthModeMTLS: - // The mTLS material is validated centrally by - // validateTLSMaterial (called from Validate) so the same - // rules apply whether mTLS is the credential or layered on - // top of bearer/api_key/basic/oauth2_*. The mode-specific - // requirement (cert + key MUST be present) is enforced - // there via Config.AuthMode inspection. - return nil - default: - return fmt.Errorf("apigateway: invalid auth_mode %q (want none, bearer, api_key, basic, oauth2_client_credentials, oauth2_authorization_code, or mtls)", c.AuthMode) - } -} - -// validateBasicAuth enforces RFC 7617 + the platform's smuggling -// defenses for the "basic" auth mode. The userid (username) must be -// non-empty and contain no ":" (RFC 7617 §2 forbids it because the -// decoder splits on the first colon). Both fields must be free of -// CR/LF/NUL because neither RFC 7617 nor base64 stops an operator from -// pasting a "username\r\nX-Smuggled: 1" string that would inject -// extra headers after the toolkit's Authorization line. The password -// may be empty: some legacy APIs accept a bearer token in the userid -// slot with an empty password (the "token:" pattern), so refusing -// empty here would block a real use case. -func (c Config) validateBasicAuth() error { - if c.Username == "" { - return errors.New("apigateway: username is required when auth_mode is \"basic\"") - } - // Smuggling defenses run before the colon check: a payload like - // "alice\r\nX-Smuggled: 1" contains both CRLF and ":" and we want - // the security-relevant error to surface, not the RFC compliance - // one. - if strings.ContainsAny(c.Username, "\r\n\x00") { - return errors.New("apigateway: username contains CR/LF/NUL header smuggling vector") - } - if strings.ContainsAny(c.Password, "\r\n\x00") { - return errors.New("apigateway: password contains CR/LF/NUL header smuggling vector") - } - if strings.Contains(c.Username, ":") { - return errors.New("apigateway: username must not contain \":\" (RFC 7617 §2 forbids it in the userid)") - } - return nil -} - -// validateOAuth2AuthCode adds the authorization_code-specific -// requirement (AuthorizationURL) on top of the client_credentials -// validation. ClientSecret is still required because OAuth 2.1 -// authorization-code with confidential clients exchanges -// (client_id, client_secret, code) for tokens. -func (c Config) validateOAuth2AuthCode() error { - if err := c.validateOAuth2(); err != nil { - return err - } - if c.OAuth2.AuthorizationURL == "" { - return errors.New("apigateway: oauth2.authorization_url is required when auth_mode is \"oauth2_authorization_code\"") - } - return nil -} - -func (c Config) validateOAuth2() error { - if c.OAuth2.TokenURL == "" { - return errors.New("apigateway: oauth2.token_url is required when auth_mode is \"oauth2_client_credentials\"") - } - if c.OAuth2.ClientID == "" { - return errors.New("apigateway: oauth2.client_id is required when auth_mode is \"oauth2_client_credentials\"") - } - if c.OAuth2.ClientSecret == "" { - return errors.New("apigateway: oauth2.client_secret is required when auth_mode is \"oauth2_client_credentials\"") - } - switch c.OAuth2.EndpointAuthStyle { - case OAuth2AuthStyleHeader, OAuth2AuthStyleParams: - return nil - default: - return fmt.Errorf("apigateway: invalid oauth2.endpoint_auth_style %q (want %q or %q)", - c.OAuth2.EndpointAuthStyle, OAuth2AuthStyleHeader, OAuth2AuthStyleParams) - } -} - -func (c Config) validateAPIKeyAuth() error { - if c.Credential == "" { - return errors.New("apigateway: credential is required when auth_mode is \"api_key\"") - } - switch c.CredentialPlacement { - case CredentialPlacementHeader: - if c.APIKeyHeader == "" { - return errors.New("apigateway: api_key_header must not be empty") - } - case CredentialPlacementQuery: - if c.APIKeyParam == "" { - return errors.New("apigateway: api_key_param is required when api_key_placement is \"query\"") - } - default: - return fmt.Errorf("apigateway: invalid api_key_placement %q (want header or query)", c.CredentialPlacement) - } - return nil -} - // IsOAuthAuthorizationCode reports whether the connection uses the // OAuth authorization_code grant (canonical AuthModeOAuth plus that // grant). The admin redirect handler and the kind handler gate the // one-time browser flow on this, so they do not depend on the raw // auth_mode string shape. func (c Config) IsOAuthAuthorizationCode() bool { - return c.AuthMode == AuthModeOAuth && c.OAuth2.Grant == connoauth.GrantAuthorizationCode -} - -// isOAuthAuthMode reports whether mode names an OAuth connection in any -// of the recognized input shapes: the canonical AuthModeOAuth or either -// legacy api-only mode that encoded the grant. ParseConfig uses this to -// decide when to delegate to connoauth.ParseConfig and normalize. -func isOAuthAuthMode(mode string) bool { - switch mode { - case AuthModeOAuth, AuthModeOAuth2ClientCredentials, AuthModeOAuth2AuthorizationCode: - return true - default: - return false - } -} - -// oauth2ConfigFromConnoauth projects the shared connoauth.Config onto -// the toolkit's OAuth2Config. The endpoint auth style is mapped back to -// the operator-facing string form the authenticators and validation -// expect. -func oauth2ConfigFromConnoauth(c connoauth.Config) OAuth2Config { - style := OAuth2AuthStyleHeader - if c.EndpointAuthStyle == oauth2.AuthStyleInParams { - style = OAuth2AuthStyleParams - } - return OAuth2Config{ - Grant: c.Grant, - TokenURL: c.TokenURL, - ClientID: c.ClientID, - ClientSecret: c.ClientSecret, - Scopes: c.Scopes, - EndpointAuthStyle: style, - AuthorizationURL: c.AuthorizationURL, - Prompt: c.Prompt, - } -} - -// getStringMap reads a map[string]string from the config map. Accepts -// map[string]string (programmatic construction) or map[string]any (YAML/JSON -// unmarshaling). Non-string values are skipped. Empty/missing returns nil. -func getStringMap(cfg map[string]any, key string) map[string]string { - raw, ok := cfg[key] - if !ok { - return nil - } - switch v := raw.(type) { - case map[string]string: - if len(v) == 0 { - return nil - } - out := make(map[string]string, len(v)) - maps.Copy(out, v) - return out - case map[string]any: - if len(v) == 0 { - return nil - } - out := make(map[string]string, len(v)) - for k, val := range v { - if s, isStr := val.(string); isStr { - out[k] = s - } - } - if len(out) == 0 { - return nil - } - return out - } - return nil + return c.upstream().IsOAuthAuthorizationCode() } func trimTrailingSlash(s string) string { @@ -846,74 +463,3 @@ func trimTrailingSlash(s string) string { } return s } - -func getString(cfg map[string]any, key string) string { - if v, ok := cfg[key].(string); ok { - return v - } - return "" -} - -func getStringDefault(cfg map[string]any, key, defaultVal string) string { - if v, ok := cfg[key].(string); ok && v != "" { - return v - } - return defaultVal -} - -func getDuration(cfg map[string]any, key string, defaultVal time.Duration) time.Duration { - raw, ok := cfg[key] - if !ok { - return defaultVal - } - switch v := raw.(type) { - case string: - if d, err := time.ParseDuration(v); err == nil { - return d - } - case time.Duration: - return v - case int: - return time.Duration(v) * time.Second - case int64: - return time.Duration(v) * time.Second - case float64: - return time.Duration(v) * time.Second - } - return defaultVal -} - -func getInt64(cfg map[string]any, key string, defaultVal int64) int64 { - raw, ok := cfg[key] - if !ok { - return defaultVal - } - switch v := raw.(type) { - case int: - return int64(v) - case int64: - return v - case float64: - return int64(v) - } - return defaultVal -} - -// getBool reads a boolean flag from the config map. Absent or -// unrecognized values default to false. A string value is parsed -// leniently (strconv.ParseBool) so YAML/JSON that round-trips a flag as -// "true"/"false" is honored alongside a native bool. -func getBool(cfg map[string]any, key string) bool { - raw, ok := cfg[key] - if !ok { - return false - } - switch v := raw.(type) { - case bool: - return v - case string: - b, err := strconv.ParseBool(v) - return err == nil && b - } - return false -} diff --git a/pkg/toolkits/apigateway/config_test.go b/pkg/toolkits/apigateway/config_test.go index 908ad15bb..b211d1bee 100644 --- a/pkg/toolkits/apigateway/config_test.go +++ b/pkg/toolkits/apigateway/config_test.go @@ -607,26 +607,3 @@ func TestParseConfig_StaticHeaders_EmptyMapNotPersisted(t *testing.T) { t.Errorf("empty static_headers stored as %#v; want nil", c.StaticHeaders) } } - -func TestIsValidHeaderName(t *testing.T) { - cases := []struct { - in string - want bool - }{ - {"X-API-Key", true}, - {"X-Subscription-Key", true}, - {"x-goog-user-project", true}, - {"Content-Type", true}, - {"", false}, - {"Bad Name", false}, - {"With\rCR", false}, - {"Colon:Inside", false}, - } - for _, tc := range cases { - t.Run(tc.in, func(t *testing.T) { - if got := isValidHeaderName(tc.in); got != tc.want { - t.Errorf("isValidHeaderName(%q) = %v; want %v", tc.in, got, tc.want) - } - }) - } -} diff --git a/pkg/toolkits/apigateway/identity_passthrough_test.go b/pkg/toolkits/apigateway/identity_passthrough_test.go index b9b9fc5ee..7617ee6cd 100644 --- a/pkg/toolkits/apigateway/identity_passthrough_test.go +++ b/pkg/toolkits/apigateway/identity_passthrough_test.go @@ -63,34 +63,6 @@ func TestValidate_IdentityPassthroughRequiresAuthModeNone(t *testing.T) { } } -func TestGetBool(t *testing.T) { - tests := []struct { - name string - val any - want bool - }{ - {"native true", true, true}, - {"native false", false, false}, - {"string true", "true", true}, - {"string 1", "1", true}, - {"string false", "false", false}, - {"string garbage", "nope", false}, - {"absent", nil, false}, - {"wrong type", 42, false}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - cfg := map[string]any{} - if tt.val != nil { - cfg["k"] = tt.val - } - if got := getBool(cfg, "k"); got != tt.want { - t.Errorf("getBool(%v) = %v; want %v", tt.val, got, tt.want) - } - }) - } -} - // TestHandleInvoke_IdentityPassthroughForwardsCallerToken verifies the // invoke path replaces the (absent) shared credential with the caller's // inbound token read from the context. diff --git a/pkg/toolkits/apigateway/invoke.go b/pkg/toolkits/apigateway/invoke.go index 2f6b60cff..6e85881c9 100644 --- a/pkg/toolkits/apigateway/invoke.go +++ b/pkg/toolkits/apigateway/invoke.go @@ -21,6 +21,7 @@ import ( "github.com/txn2/mcp-data-platform/internal/inlinefit" "github.com/txn2/mcp-data-platform/internal/pagewalk" + "github.com/txn2/mcp-data-platform/internal/upstreamauth" "github.com/txn2/mcp-data-platform/pkg/mcpcontext" "github.com/txn2/mcp-data-platform/pkg/observability" ) @@ -237,17 +238,6 @@ func inlineBudgetFor(ctx context.Context, cfg Config) int64 { return inlineBudget(cfg) } -// readLimit is the most of a response to read given a configured limit, -// falling back to the default cap when there is none. The one definition -// the buffered call, a page of a walk, and the memory reservation share, -// so the three cannot drift. -func readLimit(limit int64) int64 { - if limit > 0 { - return limit - } - return DefaultMaxResponseBytes -} - // inlineBudget is the connection's effective inline budget: the most // of a response returned through a tool result. Unset values take the // defaults, and the read cap bounds it. @@ -394,8 +384,7 @@ func buildUpstreamRequest(ctx context.Context, cfg Config, auth Authenticator, c if err := validatePath(in.Path); err != nil { return nil, err } - authHeader := authHeaderForConfig(cfg) - if err := validateCustomHeaders(in.Headers, authHeader, cfg.StaticHeaders); err != nil { + if err := cfg.upstream().ValidateCustomHeaders(in.Headers); err != nil { return nil, err } reqURL, err := buildURL(cfg.BaseURL, in.Path, in.Query) @@ -440,14 +429,14 @@ func buildUpstreamRequest(ctx context.Context, cfg Config, auth Authenticator, c // passthrough connection (the built-in platform-admin self-connection) // must act as the calling admin, so an anonymous loopback call to the // admin API would be wrong, not merely unauthenticated. The Authorization -// header is reserved from model input by validateCustomHeaders, so the +// header is reserved from model input by ValidateCustomHeaders, so the // value set here cannot be overridden by a tool argument. func applyIdentityPassthrough(ctx context.Context, req *http.Request) error { token := mcpcontext.GetAuthToken(ctx) if token == "" { return errors.New("apigateway: identity passthrough requires an authenticated caller token, but none was present on the request") } - req.Header.Set(authorizationHeader, "Bearer "+token) + req.Header.Set(upstreamauth.AuthorizationHeader, "Bearer "+token) return nil } @@ -543,45 +532,6 @@ func checkPathSegments(p string) error { return nil } -// authorizationHeader is the HTTP header bearer-mode auth populates. -// Extracted as a named constant so the same literal isn't repeated -// across the auth dispatch switch and the header-spoof rejection. -const authorizationHeader = "Authorization" - -// authHeaderForConfig returns the canonical header name the -// connection's auth mode would set, so validateCustomHeaders can -// reject the model's attempts to spoof or override it. Empty string -// means no header-based auth (mode=none, mode=api_key with query -// placement). -func authHeaderForConfig(c Config) string { - switch c.AuthMode { - case AuthModeBearer: - return authorizationHeader - case AuthModeAPIKey: - if c.CredentialPlacement == CredentialPlacementHeader { - return c.APIKeyHeader - } - } - return "" -} - -func validateCustomHeaders(headers map[string]string, authHeader string, staticHeaders map[string]string) error { - for name := range headers { - if strings.EqualFold(name, authorizationHeader) { - return errors.New("apigateway: Authorization header is reserved; configure auth via connection") - } - if authHeader != "" && strings.EqualFold(name, authHeader) { - return fmt.Errorf("apigateway: %s header is reserved by this connection's auth_mode", authHeader) - } - for staticName := range staticHeaders { - if strings.EqualFold(name, staticName) { - return fmt.Errorf("apigateway: %s header is reserved by this connection's static_headers", staticName) - } - } - } - return nil -} - // buildURL composes the upstream URL from the connection's base // URL and the model-supplied path + query. Defense against SSRF: // the base URL is parsed independently of the path; the path is @@ -1077,7 +1027,7 @@ func buildRequest(ctx context.Context, spec requestSpec) (*http.Request, error) // Static (operator-configured) headers override per-call (model) // headers so a connection's mandatory subscription/quota header // (e.g. Google's x-goog-user-project) is authoritative. - // validateCustomHeaders also rejects model attempts at the same + // ValidateCustomHeaders also rejects model attempts at the same // header names, so this is belt-and-suspenders. for name, value := range spec.staticHeaders { req.Header.Set(name, value) @@ -1287,21 +1237,6 @@ func executeRequest(p execParams) (InvokeOutput, error) { return out, nil } -func readBody(r io.Reader, maxBytes int64) (body []byte, truncated bool, err error) { - if maxBytes <= 0 { - maxBytes = DefaultMaxResponseBytes - } - limited := io.LimitReader(r, maxBytes+1) - read, rerr := io.ReadAll(limited) - if rerr != nil { - return nil, false, fmt.Errorf("apigateway: reading response body: %w", rerr) - } - if int64(len(read)) > maxBytes { - return read[:maxBytes], true, nil - } - return read, false, nil -} - // decodeBody parses a JSON response into a Go value when the // Content-Type indicates JSON; otherwise returns the body as a // string. Decoding failure on a JSON-typed response falls back to diff --git a/pkg/toolkits/apigateway/invoke_test.go b/pkg/toolkits/apigateway/invoke_test.go index 77b94f27e..52922f2e3 100644 --- a/pkg/toolkits/apigateway/invoke_test.go +++ b/pkg/toolkits/apigateway/invoke_test.go @@ -145,56 +145,6 @@ func TestBuildURL_RejectsBaseWithoutSchemeOrHost(t *testing.T) { } } -func TestAuthHeaderForConfig(t *testing.T) { - cases := []struct { - name string - cfg Config - want string - }{ - {"none", Config{AuthMode: AuthModeNone}, ""}, - {"bearer", Config{AuthMode: AuthModeBearer}, "Authorization"}, - {"api_key header default", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: DefaultAPIKeyHeader}, DefaultAPIKeyHeader}, - {"api_key header custom", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementHeader, APIKeyHeader: "X-My-Key"}, "X-My-Key"}, - {"api_key query", Config{AuthMode: AuthModeAPIKey, CredentialPlacement: CredentialPlacementQuery, APIKeyParam: "key"}, ""}, - } - for _, tc := range cases { - t.Run(tc.name, func(t *testing.T) { - if got := authHeaderForConfig(tc.cfg); got != tc.want { - t.Errorf("authHeaderForConfig = %q; want %q", got, tc.want) - } - }) - } -} - -func TestValidateCustomHeaders_RejectsAuthorization(t *testing.T) { - err := validateCustomHeaders(map[string]string{"AUTHORIZATION": "anything"}, "", nil) - if err == nil { - t.Error("Authorization header allowed") - } -} - -func TestValidateCustomHeaders_RejectsConfiguredAPIKeyHeader(t *testing.T) { - err := validateCustomHeaders(map[string]string{"x-custom-key": "spoof"}, "X-Custom-Key", nil) - if err == nil { - t.Error("configured api_key header allowed (case-insensitive check failed)") - } -} - -func TestValidateCustomHeaders_AllowsOtherHeaders(t *testing.T) { - err := validateCustomHeaders(map[string]string{"Accept-Language": "en"}, "X-API-Key", nil) - if err != nil { - t.Errorf("unrelated header rejected: %v", err) - } -} - -func TestValidateCustomHeaders_RejectsStaticHeaderOverride(t *testing.T) { - staticHeaders := map[string]string{"X-Goog-User-Project": "secret"} - err := validateCustomHeaders(map[string]string{"x-goog-user-project": "spoof"}, "", staticHeaders) - if err == nil { - t.Error("model attempt to override static header allowed (case-insensitive check failed)") - } -} - func TestBuildURL_NoQuery(t *testing.T) { got, err := buildURL("https://api.example.com", "/v1/items", nil) if err != nil { diff --git a/pkg/toolkits/apigateway/memguard.go b/pkg/toolkits/apigateway/memguard.go index 72c139e5b..2146516a4 100644 --- a/pkg/toolkits/apigateway/memguard.go +++ b/pkg/toolkits/apigateway/memguard.go @@ -55,29 +55,6 @@ const ( ErrCodeBodyNotInlineable = "upstream_body_not_inlineable" ) -// reserveBodyBudget computes the worst-case number of bytes a buffered -// read of this response could hold and tries to reserve them against -// the shared budget. It returns the amount reserved (to be released by -// the caller) and whether the reservation was granted. -// -// When the upstream declares a Content-Length below the read cap, only -// that many bytes are reserved so small (and empty) responses do not -// each tie up the full per-request cap and falsely exhaust the budget. -// This is safe because Go's HTTP client bounds resp.Body to the declared -// Content-Length — a server that writes more than it declared cannot -// make readBody buffer beyond it. Unknown/chunked responses -// (ContentLength < 0) and over-cap responses reserve the full cap, which -// is exactly what readBody may buffer. A nil/disabled budget always -// grants the reservation and Release is a no-op, so the buffered path is -// unchanged when no budget is configured. -func reserveBodyBudget(b *MemBudget, contentLength, readCap int64) (reserved int64, ok bool) { - reserved = readCap - if contentLength >= 0 && contentLength < readCap { - reserved = contentLength - } - return reserved, b.Acquire(reserved) -} - // budgetError is the typed error the buffered tools return when a body // buffer reservation is refused. handleInvoke / handleExport detect it // (errors.As) and render the structured 429 envelope; the REST shim diff --git a/pkg/toolkits/apigateway/oauth_kind_handler.go b/pkg/toolkits/apigateway/oauth_kind_handler.go index e9db9c621..e9bcc0776 100644 --- a/pkg/toolkits/apigateway/oauth_kind_handler.go +++ b/pkg/toolkits/apigateway/oauth_kind_handler.go @@ -33,7 +33,7 @@ func NewOAuthKindHandler(_ *Toolkit) *OAuthKindHandler { // authorization_code grant: the unified handler maps that to HTTP // 409 Conflict, matching the prior per-kind handler's response code. // -// The mapping is delegated to connoauthConfigFromOAuth2 so the initial +// The mapping is delegated to upstreamauth's ConnOAuthConfig so the initial // code-exchange (this path) and the per-call silent-refresh (the // authenticator path) read every field through the same translator. // A regression here would otherwise drop CABundlePEM, Prompt, or @@ -46,7 +46,7 @@ func (*OAuthKindHandler) ParseOAuthConfig(connConfig map[string]any) (connoauth. if !cfg.IsOAuthAuthorizationCode() { return connoauth.Config{}, errors.New("connection is not configured for authorization_code OAuth") } - return connoauthConfigFromOAuth2(cfg), nil + return cfg.upstream().ConnOAuthConfig(), nil } // AfterConnect is a no-op. The API gateway's Authenticator reads the diff --git a/pkg/toolkits/apigateway/tls_test.go b/pkg/toolkits/apigateway/tls_test.go index af74caf5a..38e0759a8 100644 --- a/pkg/toolkits/apigateway/tls_test.go +++ b/pkg/toolkits/apigateway/tls_test.go @@ -19,7 +19,6 @@ import ( "net/http" "net/http/httptest" "net/url" - "strings" "testing" "time" @@ -244,53 +243,6 @@ func TestNewHTTPTransport_RejectsHandshakeWithoutClientCert(t *testing.T) { assert.True(t, errors.As(err, &urlErr)) } -// TestNewTokenExchangeClient_BadBundleFallsBackQuietly is the -// resilience contract: a CA bundle that fails to parse at runtime -// (impossible if Validate ran but possible if a caller bypassed it) -// must NOT panic or block token fetches with a nil transport. The -// fallback is a plain http.Client without the bundle, matching the -// pre-feature behavior; the request will then fail with a TLS error -// against the IdP and the operator gets a normal error path. -func TestNewTokenExchangeClient_BadBundleFallsBackQuietly(t *testing.T) { - client := newTokenExchangeClient(Config{TLSCABundlePEM: "not pem"}) - require.NotNil(t, client) - assert.Nil(t, client.Transport, "fallback must not attach a half-built transport") -} - -// TestNewTokenExchangeClient_HonorsCABundle exercises the IdP-side CA -// trust plumbing for oauth2_client_credentials: when the IdP is -// signed by a private CA in tls_ca_bundle_pem, the token-fetch must -// succeed. The negative branch (no bundle) is implicit: without the -// trust the default RoundTripper would reject the IdP's cert. -func TestNewTokenExchangeClient_HonorsCABundle(t *testing.T) { - ca := newTestCA(t) - idpCert, idpKey := ca.issueServerCert(t, "127.0.0.1") - srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.Header().Set("Content-Type", "application/json") - _, _ = io.WriteString(w, `{"access_token":"abc","token_type":"bearer","expires_in":3600}`) - })) - srv.TLS = &tls.Config{ - MinVersion: tls.VersionTLS12, - Certificates: []tls.Certificate{mustKeyPair(t, idpCert, idpKey)}, - } - srv.StartTLS() - defer srv.Close() - - cfg := Config{TLSCABundlePEM: ca.certPEM} - client := newTokenExchangeClient(cfg) - postReq, err := http.NewRequestWithContext(context.Background(), - http.MethodPost, srv.URL+"/token", - strings.NewReader("")) - require.NoError(t, err) - postReq.Header.Set("Content-Type", "application/x-www-form-urlencoded") - resp, err := client.Do(postReq) - require.NoError(t, err) - defer func() { _ = resp.Body.Close() }() - assert.Equal(t, http.StatusOK, resp.StatusCode) - body, _ := io.ReadAll(resp.Body) - assert.Contains(t, string(body), "access_token") -} - // --- test CA helpers ------------------------------------------------- // testCA is a minimal CA built per-test. It mints leaf certs for use diff --git a/pkg/toolkits/apigateway/toolkit.go b/pkg/toolkits/apigateway/toolkit.go index 3883945dd..63ea5b73e 100644 --- a/pkg/toolkits/apigateway/toolkit.go +++ b/pkg/toolkits/apigateway/toolkit.go @@ -5,7 +5,6 @@ import ( "errors" "fmt" "log/slog" - "net" "net/http" "sort" "strings" @@ -18,6 +17,7 @@ import ( "github.com/txn2/mcp-data-platform/internal/apigwmetrics" "github.com/txn2/mcp-data-platform/internal/logsan" + "github.com/txn2/mcp-data-platform/internal/upstreamauth" "github.com/txn2/mcp-data-platform/pkg/authevents" "github.com/txn2/mcp-data-platform/pkg/connoauth" "github.com/txn2/mcp-data-platform/pkg/embedding" @@ -278,9 +278,7 @@ func (t *Toolkit) SetConnOAuthStore(s connoauth.Store) { defer t.mu.Unlock() t.connOAuthStore = s for _, c := range t.connections { - if ac, ok := c.auth.(*oauth2AuthorizationCodeAuth); ok { - ac.SetConnOAuthStore(s) - } + upstreamauth.SetConnOAuthStore(c.auth, s) } } @@ -323,9 +321,7 @@ func (t *Toolkit) SetAuthEvents(w *authevents.Writer) { defer t.mu.Unlock() t.authEvents = w for _, c := range t.connections { - if ac, ok := c.auth.(*oauth2AuthorizationCodeAuth); ok { - ac.SetAuthEvents(w) - } + upstreamauth.SetAuthEvents(c.auth, w) } } @@ -806,13 +802,11 @@ func (t *Toolkit) addParsedConnection(name string, cfg Config) error { // SetConnOAuthStore still becomes functional once that wire step // runs (which re-threads the store across all connections). // Either ordering works. - if ac, ok := auth.(*oauth2AuthorizationCodeAuth); ok { - if t.connOAuthStore != nil { - ac.SetConnOAuthStore(t.connOAuthStore) - } - if t.authEvents != nil { - ac.SetAuthEvents(t.authEvents) - } + if t.connOAuthStore != nil { + upstreamauth.SetConnOAuthStore(auth, t.connOAuthStore) + } + if t.authEvents != nil { + upstreamauth.SetAuthEvents(auth, t.authEvents) } t.connections[name] = c return nil @@ -1274,76 +1268,6 @@ func routePolicyError(ctx context.Context, policy RoutePolicy, in InvokeInput, t return errors.New(msg) } -// newHTTPClient builds the per-connection HTTP client. Redirects -// are explicitly disallowed so the toolkit does not blindly -// re-issue a request (and re-attach the connection's credential) -// to a host the operator did not authorize. The model can follow -// redirects manually by reading the upstream Location header from -// the response and issuing a new api_invoke_endpoint call with the -// redirected URL. -// -// Metrics wrapping is applied by the caller (see -// apigwmetrics.Instrument) rather than here so test helpers can construct -// a bare client without threading a metrics handle through every -// call site. -// -// TLS-config build errors are intentionally not surfaced from this -// constructor. ParseConfig has already validated cert + key + CA -// bundle, so buildTLSConfig only fails here if a caller has -// constructed a Config by hand and bypassed Validate. The fallback -// returns a transport with the system default tls.Config and the -// first outbound call will fail loudly with the underlying tls -// error, which is the same surface a misconfigured transport would -// produce on any other auth mode. -func newHTTPClient(cfg Config) *http.Client { - return &http.Client{ - Timeout: cfg.CallTimeout, - Transport: newHTTPTransport(cfg), - CheckRedirect: func(_ *http.Request, _ []*http.Request) error { - return http.ErrUseLastResponse - }, - } -} - -// idleConnectionTimeout caps how long an idle keep-alive connection -// can sit in the pool before being closed. Independent of the -// per-call timeouts; a generous default reduces reconnect churn for -// chatty connections. -const idleConnectionTimeout = 90 * time.Second - -// maxIdleConnections caps the per-host pool of reusable keep-alive -// sockets. Modest because each connection's typical workload is -// occasional fan-out from MCP tool calls, not high-throughput. -const maxIdleConnections = 10 - -// newHTTPTransport builds the per-connection http.Transport. The -// dial step (TCP + TLS handshake) is bound by cfg.ConnectTimeout so -// an unreachable upstream fails fast instead of consuming the full -// CallTimeout budget. Exposed as a separate function so unit tests -// can verify the wiring without standing up a network listener. -// -// When the connection carries mTLS material (cfg.MTLSClientCertPEM -// + cfg.MTLSClientKeyPEM) or a custom CA bundle (cfg.TLSCABundlePEM), -// the transport's TLSClientConfig is populated accordingly. With -// neither set, TLSClientConfig stays nil and Go's net/http uses -// system defaults. buildTLSConfig errors here are degraded to nil -// (see newHTTPClient for the rationale). -func newHTTPTransport(cfg Config) *http.Transport { - t := &http.Transport{ - DialContext: (&net.Dialer{ - Timeout: cfg.ConnectTimeout, - }).DialContext, - TLSHandshakeTimeout: cfg.ConnectTimeout, - ExpectContinueTimeout: time.Second, - IdleConnTimeout: idleConnectionTimeout, - MaxIdleConns: maxIdleConnections, - } - if tlsCfg, err := buildTLSConfig(cfg); err == nil && tlsCfg != nil { - t.TLSClientConfig = tlsCfg - } - return t -} - // Verify interface compliance at compile time. The registry.Toolkit // shape is inlined to avoid an import cycle (pkg/registry imports // this package via factories.go). diff --git a/pkg/toolkits/apigateway/toolkit_test.go b/pkg/toolkits/apigateway/toolkit_test.go index 5eb171302..7c00be35e 100644 --- a/pkg/toolkits/apigateway/toolkit_test.go +++ b/pkg/toolkits/apigateway/toolkit_test.go @@ -14,6 +14,7 @@ import ( "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/txn2/mcp-data-platform/pkg/authevents" "github.com/txn2/mcp-data-platform/pkg/connoauth" "github.com/txn2/mcp-data-platform/pkg/toolkit" ) @@ -673,17 +674,6 @@ func TestNewHTTPClient_BlocksRedirects(t *testing.T) { } } -func TestNewHTTPTransport_AppliesConnectTimeout(t *testing.T) { - cfg := Config{ConnectTimeout: 1500 * time.Millisecond, CallTimeout: 30 * time.Second} - tr := newHTTPTransport(cfg) - if tr.TLSHandshakeTimeout != cfg.ConnectTimeout { - t.Errorf("TLSHandshakeTimeout = %v; want %v", tr.TLSHandshakeTimeout, cfg.ConnectTimeout) - } - if tr.DialContext == nil { - t.Fatal("DialContext is nil; ConnectTimeout cannot be enforced") - } -} - // TestNewHTTPClient_HasTransport prevents a regression where the // transport falls back to http.DefaultTransport (which would silently // drop cfg.ConnectTimeout). A nil Transport on the returned Client @@ -763,6 +753,57 @@ func TestJSONResult_EmbedsPayload(t *testing.T) { } } +// TestToolkit_ConnectionAddedAfterWiring_PicksUpStoreAndEvents is the +// other half of the "either ordering works" contract the wiring carries: +// a connection registered AFTER SetConnOAuthStore and SetAuthEvents must +// come up with both already attached, without waiting for a re-thread +// that will never run again. +func TestToolkit_ConnectionAddedAfterWiring_PicksUpStoreAndEvents(t *testing.T) { + tk := New("primary") + t.Cleanup(func() { _ = tk.Close() }) + + tk.SetConnOAuthStore(connoauth.NewMemoryStore()) + tk.SetAuthEvents(authevents.NewWriter(nil, slog.Default())) + + err := tk.AddConnection("acme", map[string]any{ + "base_url": "https://api.example.com", + "auth_mode": AuthModeOAuth2AuthorizationCode, + "oauth2_token_url": "https://idp.example/token", + "oauth2_authorization_url": "https://idp.example/auth", + "oauth2_client_id": "id", + "oauth2_client_secret": "sec", + }) + if err != nil { + t.Fatalf("AddConnection: %v", err) + } + + tk.mu.RLock() + c := tk.connections["acme"] + tk.mu.RUnlock() + if c == nil { + t.Fatal("connection vanished") + } + // As in the re-thread test above, the store's arrival is asserted by + // behavior: with no store the authenticator refuses with "token + // store not wired", and with the (empty) store it gets as far as + // reporting that no token is persisted. + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://api.example.com/probe", http.NoBody) + if err != nil { + t.Fatalf("NewRequest: %v", err) + } + if applyErr := c.auth.Apply(req); !errors.Is(applyErr, ErrNeedsReauth) { + t.Errorf("connection added after wiring did not receive the store: Apply = %v", applyErr) + } + + // Re-threading the events writer over an already-registered + // connection is the same contract as the store's, and it must not + // disturb the wiring the connection already has. + tk.SetAuthEvents(authevents.NewWriter(nil, slog.Default())) + if applyErr := c.auth.Apply(req); !errors.Is(applyErr, ErrNeedsReauth) { + t.Errorf("re-threading the events writer broke the connection: Apply = %v", applyErr) + } +} + // TestToolkit_SetTokenStore_RethreadsAuthorizationCodeAuth proves the // wiring contract that platform.WireAPIGatewayTokenStore depends on: // when SetTokenStore is called AFTER addParsedConnection has already @@ -806,12 +847,19 @@ func TestToolkit_SetConnOAuthStore_RethreadsAuthorizationCodeAuth(t *testing.T) if c == nil { t.Fatal("connection vanished") } - ac, ok := c.auth.(*oauth2AuthorizationCodeAuth) - if !ok { - t.Fatalf("expected *oauth2AuthorizationCodeAuth, got %T", c.auth) + // Assert the re-thread by what the Authenticator now does rather + // than by reaching into it: an authorization_code authenticator + // with no store refuses with "token store not wired", while one + // that has the (empty) store gets as far as looking for a token and + // reports ErrNeedsReauth. Only the second is reachable if the new + // store actually landed on the already-materialized authenticator. + req, err := http.NewRequestWithContext(context.Background(), http.MethodGet, "https://api.example.com/probe", http.NoBody) + if err != nil { + t.Fatalf("NewRequest: %v", err) } - if ac.store != want { - t.Error("re-thread did not deliver the new store to the existing Authenticator") + applyErr := c.auth.Apply(req) + if !errors.Is(applyErr, ErrNeedsReauth) { + t.Errorf("re-thread did not deliver the new store to the existing Authenticator: Apply = %v", applyErr) } } diff --git a/pkg/toolkits/apigateway/upstream.go b/pkg/toolkits/apigateway/upstream.go new file mode 100644 index 000000000..3093eaa65 --- /dev/null +++ b/pkg/toolkits/apigateway/upstream.go @@ -0,0 +1,138 @@ +package apigateway + +import ( + "io" + "net/http" + + "github.com/txn2/mcp-data-platform/internal/membudget" + "github.com/txn2/mcp-data-platform/internal/upstreamauth" +) + +// upstreamErrPrefix names this toolkit in every message the shared +// upstream auth and transport policy produces, so an operator whose +// connection save is refused reads "apigateway: ..." and not the name +// of a seam that appears in no configuration file. +const upstreamErrPrefix = "apigateway" + +// Authenticator applies a connection's authentication scheme to an +// outbound HTTP request. The implementations and every auth mode live +// in internal/upstreamauth, shared with the platform's other +// HTTP-based connection kinds; this alias keeps the toolkit's own +// vocabulary intact at its call sites. +type Authenticator = upstreamauth.Authenticator + +// ErrNeedsReauth is the structured error api_invoke_endpoint surfaces +// when an authorization_code connection's stored refresh token is +// missing, expired beyond refresh_expires_at, or definitively rejected +// by the IdP (RFC 6749 §5.2 invalid_grant on the refresh_token grant). +// Transient failures (network, 5xx, request cancellation) DO NOT +// produce this error. +var ErrNeedsReauth = upstreamauth.ErrNeedsReauth + +// NewAuthenticator returns the Authenticator implementation for a +// validated Config. +func NewAuthenticator(c Config) (Authenticator, error) { + //nolint:wrapcheck // the message already carries this toolkit's prefix; wrapping would state it twice + return upstreamauth.NewAuthenticator(c.upstream()) +} + +// upstream projects the toolkit's Config onto the authentication and +// transport slice internal/upstreamauth owns. The toolkit keeps its own +// exported Config — its keys, its defaults, its public API — and this +// is the single place the two are related, so a field added to either +// side has exactly one place to be joined up. +func (c Config) upstream() upstreamauth.Config { + return upstreamauth.Config{ + Kind: Kind, + ErrPrefix: upstreamErrPrefix, + ConnectionName: c.ConnectionName, + AuthMode: c.AuthMode, + Credential: c.Credential, + CredentialPlacement: c.CredentialPlacement, + APIKeyHeader: c.APIKeyHeader, + APIKeyParam: c.APIKeyParam, + Username: c.Username, + Password: c.Password, + OAuth2: upstreamauth.OAuth2Config{ + Grant: c.OAuth2.Grant, + TokenURL: c.OAuth2.TokenURL, + ClientID: c.OAuth2.ClientID, + ClientSecret: c.OAuth2.ClientSecret, + Scopes: c.OAuth2.Scopes, + EndpointAuthStyle: c.OAuth2.EndpointAuthStyle, + AuthorizationURL: c.OAuth2.AuthorizationURL, + Prompt: c.OAuth2.Prompt, + }, + ConnectTimeout: c.ConnectTimeout, + CallTimeout: c.CallTimeout, + MaxResponseBytes: c.MaxResponseBytes, + StaticHeaders: c.StaticHeaders, + MTLSClientCertPEM: c.MTLSClientCertPEM, + MTLSClientKeyPEM: c.MTLSClientKeyPEM, + TLSCABundlePEM: c.TLSCABundlePEM, + IdentityPassthrough: c.IdentityPassthrough, + } +} + +// configFromUpstream is the inverse of Config.upstream: it seeds a +// Config with the authentication and transport values Parse read, so +// ParseConfig only has to fill in the keys this toolkit owns. Kept +// beside upstream() because the two must be changed together. +func configFromUpstream(up upstreamauth.Config) Config { + return Config{ + AuthMode: up.AuthMode, + Credential: up.Credential, + CredentialPlacement: up.CredentialPlacement, + APIKeyHeader: up.APIKeyHeader, + APIKeyParam: up.APIKeyParam, + Username: up.Username, + Password: up.Password, + OAuth2: OAuth2Config{ + Grant: up.OAuth2.Grant, + TokenURL: up.OAuth2.TokenURL, + ClientID: up.OAuth2.ClientID, + ClientSecret: up.OAuth2.ClientSecret, + Scopes: up.OAuth2.Scopes, + EndpointAuthStyle: up.OAuth2.EndpointAuthStyle, + AuthorizationURL: up.OAuth2.AuthorizationURL, + Prompt: up.OAuth2.Prompt, + }, + ConnectTimeout: up.ConnectTimeout, + CallTimeout: up.CallTimeout, + MaxResponseBytes: up.MaxResponseBytes, + StaticHeaders: up.StaticHeaders, + MTLSClientCertPEM: up.MTLSClientCertPEM, + MTLSClientKeyPEM: up.MTLSClientKeyPEM, + TLSCABundlePEM: up.TLSCABundlePEM, + IdentityPassthrough: up.IdentityPassthrough, + } +} + +// newHTTPClient builds the per-connection HTTP client from the shared +// transport policy: the call timeout, the connection's TLS material, +// and a CheckRedirect that refuses 3xx. +func newHTTPClient(cfg Config) *http.Client { + return upstreamauth.NewHTTPClient(cfg.upstream()) +} + +// readLimit is the most of a response to read given a configured limit, +// falling back to the default cap when there is none. The one definition +// the buffered call, a page of a walk, and the memory reservation share, +// so the three cannot drift. +func readLimit(limit int64) int64 { + return upstreamauth.ReadLimit(limit) +} + +// readBody reads at most maxBytes of an upstream response, reporting +// whether it was cut short. +func readBody(r io.Reader, maxBytes int64) (body []byte, truncated bool, err error) { + //nolint:wrapcheck // the message already carries this toolkit's prefix + return upstreamauth.ReadBody(upstreamErrPrefix, r, maxBytes) +} + +// reserveBodyBudget reserves the worst-case buffer for one response +// against the shared in-flight budget, returning the amount to release +// and whether the reservation was granted. +func reserveBodyBudget(b *membudget.Budget, contentLength, readCap int64) (reserved int64, ok bool) { + return upstreamauth.ReserveBodyBudget(b, contentLength, readCap) +} diff --git a/test/acceptance/issue_1647_test.go b/test/acceptance/issue_1647_test.go new file mode 100644 index 000000000..12ad45fdf --- /dev/null +++ b/test/acceptance/issue_1647_test.go @@ -0,0 +1,383 @@ +//go:build integration + +package acceptance + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" + "testing" + "time" +) + +// issue1647JSONReader encodes a request body for the admin REST routes. The +// suite's restJSON returns only the status, and these criteria read the +// refusal text, so the body is built here and handed to rest directly. +func issue1647JSONReader(t *testing.T, v any) io.Reader { + t.Helper() + raw, err := json.Marshal(v) + if err != nil { + t.Fatalf("marshal request body: %v", err) + } + return bytes.NewReader(raw) +} + +// Issue #1647: the API gateway's upstream authentication and transport policy +// moved into internal/upstreamauth so a second HTTP-based connection kind +// reuses them instead of copying them. Nothing an operator or a model sees may +// change, which is what this suite executes: for every auth mode, what the +// platform actually puts on the wire, read back off the api-test fixture's own +// echo of the request it received. +// +// The fixture (dev/docker-compose.yml, api-test) is the upstream rather than a +// stand-in: /v1/echo returns the method, path, query and headers of the +// request as it arrived, so a credential the platform attached is observed +// where it lands rather than where it is built. Credential VALUES on +// Authorization and X-API-Key come back "[redacted]" — the fixture's own +// policy — so those modes assert the header's presence and the upstream's +// verdict, and the custom-header mode, whose name the fixture does not treat +// as a credential, asserts the value too. +// +// Wire forms: api_invoke_endpoint's `body` is untyped and admits an object and +// a string of JSON. Both are sent as literal tools/call params and asserted to +// reach the upstream identically. + +const ( + issue1647Tool = "api_invoke_endpoint" + issue1647FixtureURL = "http://localhost:9282" + issue1647FixtureKey = "apitest-dev-key-2024" + issue1647CustomKey = "X-Acc-1647-Key" + issue1647StaticName = "X-Acc-1647-Static" + issue1647StaticValue = "pinned-by-the-operator" + issue1647Purpose = "Acceptance for #1647: the shared upstream auth and transport seam puts the same bytes on the wire." +) + +// issue1647Connect registers one catalog-less connection to the api-test +// fixture with the given auth settings and returns its name. Every connection +// is removed when the test ends. +func issue1647Connect(t *testing.T, c *client, label string, auth map[string]any) string { + t.Helper() + name := fmt.Sprintf("acc-1647-%s-%d", label, time.Now().UnixNano()) + cfg := map[string]any{ + "base_url": issue1647FixtureURL, + "connection_name": name, + "connect_timeout": "5s", + "call_timeout": "10s", + "trust_level": "untrusted", + } + for k, v := range auth { + cfg[k] = v + } + status := c.restJSON(http.MethodPut, "/api/v1/admin/connection-instances/api/"+name, map[string]any{ + "config": cfg, + "description": "Acceptance 1647: " + label, + }) + if status != http.StatusCreated && status != http.StatusOK { + t.Fatalf("register %s connection: HTTP %d", label, status) + } + t.Cleanup(func() { + c.rest(http.MethodDelete, "/api/v1/admin/connection-instances/api/"+name, http.NoBody) + }) + return name +} + +// issue1647Echo calls the fixture's echo endpoint through the connection and +// returns the request as the fixture saw it. +func issue1647Echo(t *testing.T, c *client, connection string, extra map[string]any) map[string]any { + t.Helper() + args := map[string]any{ + "connection": connection, + "method": http.MethodGet, + "path": "/v1/echo", + "purpose": issue1647Purpose, + } + for k, v := range extra { + args[k] = v + } + out := c.call(issue1647Tool, args) + if got := number(t, out, "status"); got != http.StatusOK { + t.Fatalf("%s through %s: upstream status %v, body %v", issue1647Tool, connection, got, out["body"]) + } + body, ok := out["body"].(map[string]any) + if !ok { + t.Fatalf("echo body is not an object: %v", out["body"]) + } + return body +} + +// issue1647Header reads one header out of the fixture's echo. The fixture +// reports every header as a list of values. +func issue1647Header(body map[string]any, name string) []string { + headers, _ := body["headers"].(map[string]any) + raw, ok := headers[name].([]any) + if !ok { + return nil + } + out := make([]string, 0, len(raw)) + for _, v := range raw { + s, _ := v.(string) + out = append(out, s) + } + return out +} + +// TestIssue1647_AuthModesPutTheSameBytesOnTheWire is the criterion: for each +// auth mode the fixture can observe, the credential arrives where that mode +// says it should and nowhere else. +func TestIssue1647_AuthModesPutTheSameBytesOnTheWire(t *testing.T) { + c := connect(t) + + t.Run("none attaches no credential", func(t *testing.T) { + conn := issue1647Connect(t, c, "none", map[string]any{"auth_mode": "none"}) + body := issue1647Echo(t, c, conn, nil) + if got := issue1647Header(body, "Authorization"); got != nil { + t.Errorf("auth_mode=none sent an Authorization header: %v", got) + } + if got := issue1647Header(body, "X-Api-Key"); got != nil { + t.Errorf("auth_mode=none sent an X-API-Key header: %v", got) + } + }) + + t.Run("api_key in a custom header carries name and value", func(t *testing.T) { + conn := issue1647Connect(t, c, "apikey-header", map[string]any{ + "auth_mode": "api_key", + "credential": issue1647FixtureKey, + "api_key_placement": "header", + "api_key_header": issue1647CustomKey, + }) + body := issue1647Echo(t, c, conn, nil) + got := issue1647Header(body, issue1647CustomKey) + if len(got) != 1 || got[0] != issue1647FixtureKey { + t.Errorf("%s = %v; want the configured credential", issue1647CustomKey, got) + } + }) + + t.Run("api_key in the query carries name and value", func(t *testing.T) { + conn := issue1647Connect(t, c, "apikey-query", map[string]any{ + "auth_mode": "api_key", + "credential": issue1647FixtureKey, + "api_key_placement": "query", + "api_key_param": "acc_1647_key", + }) + body := issue1647Echo(t, c, conn, nil) + query, _ := body["query"].(map[string]any) + values, _ := query["acc_1647_key"].([]any) + if len(values) != 1 || values[0] != issue1647FixtureKey { + t.Errorf("query acc_1647_key = %v; want the configured credential", values) + } + if got := issue1647Header(body, issue1647CustomKey); got != nil { + t.Errorf("query placement also sent a header: %v", got) + } + }) + + t.Run("api_key in the default header authenticates the caller", func(t *testing.T) { + conn := issue1647Connect(t, c, "apikey-default", map[string]any{ + "auth_mode": "api_key", + "credential": issue1647FixtureKey, + }) + out := c.call(issue1647Tool, map[string]any{ + "connection": conn, "method": http.MethodGet, "path": "/v1/whoami", + "purpose": issue1647Purpose, + }) + body, _ := out["body"].(map[string]any) + if body["auth_type"] != "apikey" || body["key_name"] != "dev-fixture" { + t.Errorf("whoami = %v; want the fixture to recognize the api key", body) + } + }) + + t.Run("basic attaches an Authorization header", func(t *testing.T) { + conn := issue1647Connect(t, c, "basic", map[string]any{ + "auth_mode": "basic", + "username": "alice", + "password": "s3cret", + }) + body := issue1647Echo(t, c, conn, nil) + if got := issue1647Header(body, "Authorization"); len(got) != 1 { + t.Errorf("auth_mode=basic sent Authorization %v; want exactly one value", got) + } + }) + + // The fixture refuses a bearer token it does not know, and answers the + // same path 200 with no credential at all. The contrast is the proof the + // header was attached: its value comes back redacted, so presence is + // asserted by the upstream's verdict rather than by reading it. + t.Run("bearer attaches the token the upstream then judges", func(t *testing.T) { + conn := issue1647Connect(t, c, "bearer", map[string]any{ + "auth_mode": "bearer", + "credential": "acc-1647-token-the-fixture-does-not-know", + }) + out := c.call(issue1647Tool, map[string]any{ + "connection": conn, "method": http.MethodGet, "path": "/v1/echo", + "purpose": issue1647Purpose, + }) + if got := number(t, out, "status"); got != http.StatusUnauthorized { + t.Fatalf("bearer call status = %v; want 401 from the fixture, which means the header arrived", got) + } + body, _ := out["body"].(map[string]any) + if !strings.Contains(fmt.Sprint(body), "invalid credential") { + t.Errorf("fixture refusal = %v; want its invalid-credential answer", body) + } + }) +} + +// TestIssue1647_OperatorHeadersAndTheirReservations covers the other half of +// the policy the seam owns: headers the operator pins, and the model's +// inability to reach any header a credential or the operator already claims. +func TestIssue1647_OperatorHeadersAndTheirReservations(t *testing.T) { + c := connect(t) + conn := issue1647Connect(t, c, "static", map[string]any{ + "auth_mode": "api_key", + "credential": issue1647FixtureKey, + "api_key_placement": "header", + "api_key_header": issue1647CustomKey, + "static_headers": map[string]any{issue1647StaticName: issue1647StaticValue}, + }) + + t.Run("a pinned header reaches the upstream", func(t *testing.T) { + body := issue1647Echo(t, c, conn, nil) + got := issue1647Header(body, issue1647StaticName) + if len(got) != 1 || got[0] != issue1647StaticValue { + t.Errorf("%s = %v; want the operator's value", issue1647StaticName, got) + } + }) + + t.Run("a model header the connection does not claim is forwarded", func(t *testing.T) { + body := issue1647Echo(t, c, conn, map[string]any{ + "headers": map[string]any{"X-Acc-1647-Free": "from-the-model"}, + }) + got := issue1647Header(body, "X-Acc-1647-Free") + if len(got) != 1 || got[0] != "from-the-model" { + t.Errorf("X-Acc-1647-Free = %v; want the model's value", got) + } + }) + + reserved := []struct { + name string + header string + wantMsg string + }{ + {"Authorization", "Authorization", "Authorization header is reserved"}, + {"the auth mode's own header", strings.ToLower(issue1647CustomKey), "reserved by this connection's auth_mode"}, + {"a pinned header", strings.ToLower(issue1647StaticName), "reserved by this connection's static_headers"}, + } + for _, tc := range reserved { + t.Run("the model may not set "+tc.name, func(t *testing.T) { + res, text, err := c.callRaw(issue1647Tool, map[string]any{ + "connection": conn, "method": http.MethodGet, "path": "/v1/echo", + "headers": map[string]any{tc.header: "spoofed"}, + "purpose": issue1647Purpose, + }) + if err != nil { + t.Fatalf("transport error: %v", err) + } + if !res.IsError { + t.Fatalf("setting %s was allowed: %s", tc.header, text) + } + if !strings.Contains(text, tc.wantMsg) { + t.Errorf("refusal = %q; want it to name %q", text, tc.wantMsg) + } + }) + } +} + +// TestIssue1647_BodyReachesTheUpstreamInEveryWireForm sends the one untyped +// parameter this tool takes in both forms its schema admits, as literal +// tools/call params, and asserts the upstream received the same document. +func TestIssue1647_BodyReachesTheUpstreamInEveryWireForm(t *testing.T) { + c := connect(t) + conn := issue1647Connect(t, c, "body", map[string]any{ + "auth_mode": "api_key", + "credential": issue1647FixtureKey, + }) + + forms := []struct { + name string + body any + }{ + {"object", map[string]any{"issue": float64(1647), "shape": "object"}}, + {"string of JSON", `{"issue":1647,"shape":"object"}`}, + } + for _, form := range forms { + t.Run(form.name, func(t *testing.T) { + out := c.call(issue1647Tool, map[string]any{ + "connection": conn, "method": http.MethodPost, "path": "/v1/echo", + "body": form.body, + "purpose": issue1647Purpose, + }) + if got := number(t, out, "status"); got != http.StatusOK { + t.Fatalf("status = %v, body %v", got, out["body"]) + } + echoed, _ := out["body"].(map[string]any) + received, _ := echoed["body"].(map[string]any) + if received["issue"] != float64(1647) || received["shape"] != "object" { + t.Errorf("the upstream received %v; want the document the call carried", received) + } + }) + } +} + +// TestIssue1647_ConfigRefusalsKeepTheToolkitsVoice is the operator-facing half +// of the move: the auth and transport rules now live behind a shared seam, and +// a refused connection save must still name the surface the operator +// configured rather than the seam behind it. +func TestIssue1647_ConfigRefusalsKeepTheToolkitsVoice(t *testing.T) { + c := connect(t) + + cases := []struct { + name string + cfg map[string]any + wantMsg string + }{ + { + name: "bearer without a credential", + cfg: map[string]any{"auth_mode": "bearer"}, + wantMsg: `apigateway: credential is required when auth_mode is "bearer"`, + }, + { + name: "api_key in the query without a parameter name", + cfg: map[string]any{"auth_mode": "api_key", "credential": "k", "api_key_placement": "query"}, + wantMsg: `apigateway: api_key_param is required when api_key_placement is "query"`, + }, + { + name: "basic with a colon in the userid", + cfg: map[string]any{"auth_mode": "basic", "username": "a:b"}, + wantMsg: "apigateway: username must not contain", + }, + { + name: "a pinned header that would fight the auth layer", + cfg: map[string]any{"auth_mode": "none", "static_headers": map[string]any{"Authorization": "Bearer x"}}, + wantMsg: "apigateway: static_headers must not set Authorization", + }, + { + name: "an unusable timeout", + cfg: map[string]any{"auth_mode": "none", "call_timeout": "0s"}, + wantMsg: "apigateway: call_timeout must be positive", + }, + { + name: "mtls without the keypair that is its credential", + cfg: map[string]any{"auth_mode": "mtls"}, + wantMsg: "apigateway: mtls_client_cert_pem and mtls_client_key_pem are required", + }, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + name := fmt.Sprintf("acc-1647-bad-%d", time.Now().UnixNano()) + cfg := map[string]any{"base_url": issue1647FixtureURL, "connection_name": name} + for k, v := range tc.cfg { + cfg[k] = v + } + status, body := c.rest(http.MethodPut, "/api/v1/admin/connection-instances/api/"+name, + issue1647JSONReader(t, map[string]any{"config": cfg, "description": "Acceptance 1647: refused"})) + if status == http.StatusCreated || status == http.StatusOK { + c.rest(http.MethodDelete, "/api/v1/admin/connection-instances/api/"+name, http.NoBody) + t.Fatalf("an invalid connection was accepted: HTTP %d", status) + } + if !strings.Contains(fmt.Sprint(body), tc.wantMsg) { + t.Errorf("refusal body = %v; want it to carry %q", body, tc.wantMsg) + } + }) + } +} diff --git a/test/structure/testdata/allowed_internal_imports.txt b/test/structure/testdata/allowed_internal_imports.txt index 18b0ac4c8..23aa6d206 100644 --- a/test/structure/testdata/allowed_internal_imports.txt +++ b/test/structure/testdata/allowed_internal_imports.txt @@ -488,6 +488,11 @@ internal/producedview -> pkg/script internal/server -> internal/platform/knowledgebuiltin internal/server -> pkg/platform internal/tableavail -> pkg/query +internal/upstreamauth -> internal/apigwtls +internal/upstreamauth -> internal/cfgmap +internal/upstreamauth -> internal/membudget +internal/upstreamauth -> pkg/authevents +internal/upstreamauth -> pkg/connoauth pkg/admin -> internal/admin/apiroutesapi pkg/admin -> internal/admin/auditapi pkg/admin -> internal/admin/callapi @@ -795,11 +800,12 @@ pkg/textpatch -> pkg/contenttype pkg/textpatch/patchmcp -> pkg/middleware pkg/textpatch/patchmcp -> pkg/textpatch pkg/toolkits/apigateway -> internal/apigwmetrics -pkg/toolkits/apigateway -> internal/apigwtls +pkg/toolkits/apigateway -> internal/cfgmap pkg/toolkits/apigateway -> internal/inlinefit pkg/toolkits/apigateway -> internal/logsan pkg/toolkits/apigateway -> internal/membudget pkg/toolkits/apigateway -> internal/pagewalk +pkg/toolkits/apigateway -> internal/upstreamauth pkg/toolkits/apigateway -> pkg/authevents pkg/toolkits/apigateway -> pkg/blobserve pkg/toolkits/apigateway -> pkg/connoauth