diff --git a/go.mod b/go.mod index 5f49d165..77a8f3ec 100644 --- a/go.mod +++ b/go.mod @@ -49,6 +49,7 @@ require ( go.opentelemetry.io/otel/sdk v1.43.0 go.opentelemetry.io/otel/trace v1.44.0 go.uber.org/goleak v1.3.0 + golang.org/x/net v0.56.0 golang.org/x/term v0.44.0 golang.org/x/text v0.39.0 google.golang.org/genproto/googleapis/api v0.0.0-20260420184626-e10c466a9529 @@ -176,7 +177,6 @@ require ( go.yaml.in/yaml/v3 v3.0.4 // indirect golang.org/x/crypto v0.53.0 // indirect golang.org/x/exp v0.0.0-20260410095643-746e56fc9e2f // indirect - golang.org/x/net v0.56.0 // indirect golang.org/x/oauth2 v0.36.0 // indirect golang.org/x/sys v0.47.0 // indirect golang.org/x/time v0.15.0 // indirect diff --git a/network/urlguard/admit.go b/network/urlguard/admit.go new file mode 100644 index 00000000..fc9ff323 --- /dev/null +++ b/network/urlguard/admit.go @@ -0,0 +1,60 @@ +package urlguard + +import ( + "fmt" + "slices" + "strings" +) + +// Ceiling is the manifest-declared origin ceiling. A concrete origin is +// admitted only if it stays inside every dimension of the ceiling. Host +// patterns are bounded: either an exact host or a single "*." wildcard that +// matches proper subdomains only. +type Ceiling struct { + Schemes []string + HostPatterns []string + Ports []uint32 + AllowedClasses []NetworkClass +} + +// Admit normalizes raw and confirms it stays within the ceiling. It returns +// the normalized origin on success. Admission is purely structural: it does +// not resolve DNS. The caller pins a peer with Resolve and dials with Client. +func (c Ceiling) Admit(raw string) (Origin, error) { + origin, err := NormalizeOrigin(raw) + if err != nil { + return Origin{}, err + } + return origin, c.AdmitOrigin(origin) +} + +// AdmitOrigin confirms an already-normalized origin stays within the ceiling. +func (c Ceiling) AdmitOrigin(origin Origin) error { + if !slices.Contains(c.Schemes, origin.Scheme) { + return fmt.Errorf("origin scheme %q is outside the ceiling", origin.Scheme) + } + if !slices.Contains(c.Ports, origin.Port) { + return fmt.Errorf("origin port %d is outside the ceiling", origin.Port) + } + if !hostMatchesAny(origin.Host, c.HostPatterns) { + return fmt.Errorf("origin host %q is outside the ceiling", origin.Host) + } + return nil +} + +// hostMatchesAny reports whether host is an exact match for a pattern or a +// proper subdomain of a "*." pattern. The apex of a wildcard pattern is not +// matched by the wildcard itself. +func hostMatchesAny(host string, patterns []string) bool { + for _, pattern := range patterns { + if host == pattern { + return true + } + if suffix, ok := strings.CutPrefix(pattern, "*."); ok { + if strings.HasSuffix(host, "."+suffix) { + return true + } + } + } + return false +} diff --git a/network/urlguard/fuzz_test.go b/network/urlguard/fuzz_test.go new file mode 100644 index 00000000..b88cbc62 --- /dev/null +++ b/network/urlguard/fuzz_test.go @@ -0,0 +1,54 @@ +package urlguard_test + +import ( + "testing" + + "github.com/codefly-dev/core/network/urlguard" +) + +// FuzzNormalizeOrigin proves URL normalization never panics and that an +// admitted origin is always in canonical form (round-trips through parse). +func FuzzNormalizeOrigin(f *testing.F) { + for _, seed := range []string{ + "https://api.example.com", + "https://API.Example.com.:8443", + "http://user:pass@host/path?q#f", + "//scheme-relative", + "https://[::1]", + "ftp://host", + "https://xn--pple-43d.com", + } { + f.Add(seed) + } + f.Fuzz(func(t *testing.T, raw string) { + origin, err := urlguard.NormalizeOrigin(raw) + if err != nil { + return + } + // A normalized origin must re-normalize to itself. + again, err := urlguard.NormalizeOrigin(origin.String()) + if err != nil { + t.Fatalf("normalized origin %q did not re-normalize: %v", origin.String(), err) + } + if again != origin { + t.Fatalf("normalization is not idempotent: %v -> %v", origin, again) + } + }) +} + +// FuzzSafePath proves path validation never panics and never admits traversal +// or a query/fragment. +func FuzzSafePath(f *testing.F) { + for _, seed := range []string{"/", "/v1/x", "/a/../b", "/a%2e%2e/b", "/a?x=1", "//b"} { + f.Add(seed) + } + f.Fuzz(func(t *testing.T, path string) { + safe, err := urlguard.SafePath(path) + if err != nil { + return + } + if safe == "" || safe[0] != '/' { + t.Fatalf("safe path %q is not absolute", safe) + } + }) +} diff --git a/network/urlguard/urlguard.go b/network/urlguard/urlguard.go new file mode 100644 index 00000000..5b9124ca --- /dev/null +++ b/network/urlguard/urlguard.go @@ -0,0 +1,300 @@ +// Package urlguard provides host-owned URL normalization, origin admission, +// and SSRF-hardened transport construction for outbound requests that a +// codefly host makes on behalf of untrusted provider code. +// +// Provider code never dials the network directly. It hands the host a +// descriptor; the host normalizes and admits the origin here, resolves and +// pins a single peer IP, and dials only that IP. This closes the classic +// bypasses: URL userinfo, encoded path traversal, scheme-relative URLs, +// proxy environment influence, credentialed redirects, and DNS rebinding +// between validation and dial. +package urlguard + +import ( + "context" + "crypto/tls" + "fmt" + "net" + "net/http" + "net/url" + "slices" + "strconv" + "strings" + "time" + + "golang.org/x/net/idna" +) + +// NetworkClass classifies a concrete resolved address. It is independent of +// the provider proto so this package stays reusable; callers map it to their +// own vocabulary. +type NetworkClass string + +const ( + ClassPublic NetworkClass = "public" + ClassLoopback NetworkClass = "loopback" + ClassLinkLocal NetworkClass = "link-local" + ClassPrivate NetworkClass = "private" +) + +// Origin is a normalized, host-attested base origin. Host is a canonical +// lowercase ASCII hostname or IP literal with no trailing dot; Port is always +// explicit. +type Origin struct { + Scheme string + Host string + Port uint32 +} + +// String renders the canonical scheme://host:port form. IPv6 hosts are +// bracketed. +func (o Origin) String() string { + return o.Scheme + "://" + o.HostPort() +} + +// HostPort renders host:port, bracketing IPv6 hosts, suitable for a URL +// authority. +func (o Origin) HostPort() string { + host := o.Host + if strings.Contains(host, ":") { + host = "[" + host + "]" + } + return fmt.Sprintf("%s:%d", host, o.Port) +} + +// defaultPort returns the scheme default. Callers only reach it for the two +// admitted schemes. +func defaultPort(scheme string) uint32 { + if scheme == "http" { + return 80 + } + return 443 +} + +// NormalizeOrigin parses and canonicalizes a base-origin string. It rejects +// anything that could smuggle a second interpretation past admission: URL +// userinfo, a path/query/fragment, opaque or scheme-relative forms, non-HTTP(S) +// schemes, and IDNA-ambiguous hosts. +func NormalizeOrigin(raw string) (Origin, error) { + if raw == "" { + return Origin{}, fmt.Errorf("origin is required") + } + if strings.HasPrefix(raw, "//") { + return Origin{}, fmt.Errorf("origin must not be scheme-relative") + } + parsed, err := url.Parse(raw) + if err != nil { + return Origin{}, fmt.Errorf("origin is not a valid URL: %w", err) + } + if parsed.Opaque != "" { + return Origin{}, fmt.Errorf("origin must not be opaque") + } + if parsed.User != nil { + return Origin{}, fmt.Errorf("origin must not carry userinfo") + } + if parsed.Path != "" || parsed.RawQuery != "" || parsed.Fragment != "" || parsed.ForceQuery { + return Origin{}, fmt.Errorf("base origin must not carry a path, query, or fragment") + } + scheme := strings.ToLower(parsed.Scheme) + if scheme != "https" && scheme != "http" { + return Origin{}, fmt.Errorf("origin scheme %q is not admitted", parsed.Scheme) + } + host, err := canonicalHost(parsed.Hostname()) + if err != nil { + return Origin{}, err + } + port := defaultPort(scheme) + if raw := parsed.Port(); raw != "" { + n, err := strconv.ParseUint(raw, 10, 16) + if err != nil || n == 0 { + return Origin{}, fmt.Errorf("origin port %q is invalid", raw) + } + port = uint32(n) + } + return Origin{Scheme: scheme, Host: host, Port: port}, nil +} + +// canonicalHost lowercases, strips a single trailing dot, canonicalizes IDNA +// to ASCII, and canonicalizes IP literals. It rejects hosts that are not +// already in their canonical ASCII form so a provider cannot admit one origin +// and later dial a homograph. +func canonicalHost(host string) (string, error) { + if host == "" { + return "", fmt.Errorf("origin host is required") + } + host = strings.TrimSuffix(host, ".") + if host == "" { + return "", fmt.Errorf("origin host is required") + } + if ip := net.ParseIP(host); ip != nil { + return ip.String(), nil + } + lowered := strings.ToLower(host) + ascii, err := idna.Lookup.ToASCII(lowered) + if err != nil { + return "", fmt.Errorf("origin host is not a valid domain: %w", err) + } + if ascii != lowered { + return "", fmt.Errorf("origin host is not in canonical form") + } + if !validHostname(ascii) { + return "", fmt.Errorf("origin host is not a valid domain") + } + return ascii, nil +} + +func validHostname(host string) bool { + if len(host) == 0 || len(host) > 253 || strings.Contains(host, "..") { + return false + } + for _, label := range strings.Split(host, ".") { + if len(label) == 0 || len(label) > 63 || label[0] == '-' || label[len(label)-1] == '-' { + return false + } + for _, r := range label { + if (r < 'a' || r > 'z') && (r < '0' || r > '9') && r != '-' { + return false + } + } + } + return true +} + +// SafePath validates that a request path is already in canonical, fully +// unescaped-safe form. It rejects encoded traversal and any character that +// would let the path smuggle a query, fragment, or second path segment past +// admission. An empty path is treated as "/". +func SafePath(path string) (string, error) { + if path == "" { + return "/", nil + } + if path[0] != '/' { + return "", fmt.Errorf("path must be absolute") + } + if strings.ContainsAny(path, "?#\\") { + return "", fmt.Errorf("path must not contain a query, fragment, or backslash") + } + if strings.Contains(path, "..") { + return "", fmt.Errorf("path must not contain traversal") + } + // %2e (.) and %2f (/) reintroduce traversal after the server decodes; any + // percent-encoding of a path separator or dot is rejected outright. + lowered := strings.ToLower(path) + if strings.Contains(lowered, "%2e") || strings.Contains(lowered, "%2f") || strings.Contains(lowered, "%5c") { + return "", fmt.Errorf("path must not contain encoded traversal") + } + if strings.Contains(path, "//") { + return "", fmt.Errorf("path must not contain empty segments") + } + return path, nil +} + +// Classify maps a concrete IP to its network class. Cloud metadata endpoints +// (169.254.169.254) fall in link-local and are blocked with the rest of that +// class unless a rule explicitly admits link-local. +func Classify(ip net.IP) NetworkClass { + switch { + case ip.IsLoopback(): + return ClassLoopback + case ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast(): + return ClassLinkLocal + case ip.IsUnspecified() || ip.IsPrivate() || ip.IsMulticast() || ip.IsInterfaceLocalMulticast() || isCGNAT(ip): + return ClassPrivate + default: + return ClassPublic + } +} + +// isCGNAT reports whether ip is in the RFC 6598 shared-address space +// (100.64.0.0/10), which Go's IsPrivate does not cover. +func isCGNAT(ip net.IP) bool { + v4 := ip.To4() + if v4 == nil { + return false + } + return v4[0] == 100 && v4[1]&0xc0 == 64 +} + +// admitClass reports whether class may be dialed. Public is always admitted; +// any private class requires explicit allowance. +func admitClass(class NetworkClass, allowed []NetworkClass) bool { + return class == ClassPublic || slices.Contains(allowed, class) +} + +// Resolution is the single pinned peer for one operation. +type Resolution struct { + IP net.IP + Class NetworkClass +} + +// Resolve looks up the origin host exactly once, classifies every returned +// address, and pins the first admitted one. If any returned address is a +// disallowed class the resolution fails closed — a host that returns one +// public and one private answer must not be dialable. An IP-literal host is +// classified without a lookup. +func Resolve(ctx context.Context, resolver *net.Resolver, origin Origin, allowed []NetworkClass) (Resolution, error) { + if resolver == nil { + resolver = net.DefaultResolver + } + if literal := net.ParseIP(origin.Host); literal != nil { + class := Classify(literal) + if !admitClass(class, allowed) { + return Resolution{}, fmt.Errorf("origin address class %q is not admitted", class) + } + return Resolution{IP: literal, Class: class}, nil + } + addrs, err := resolver.LookupIPAddr(ctx, origin.Host) + if err != nil { + return Resolution{}, fmt.Errorf("resolve origin host: %w", err) + } + if len(addrs) == 0 { + return Resolution{}, fmt.Errorf("origin host %q did not resolve", origin.Host) + } + for _, addr := range addrs { + if class := Classify(addr.IP); !admitClass(class, allowed) { + return Resolution{}, fmt.Errorf("origin host resolved to disallowed class %q", class) + } + } + first := addrs[0] + return Resolution{IP: first.IP, Class: Classify(first.IP)}, nil +} + +// Client builds an SSRF-hardened HTTP client bound to one pinned resolution. +// The transport has no proxy, dials only the pinned IP regardless of the +// address the URL carries (defeating DNS rebinding), verifies TLS against the +// origin host, and the client refuses every redirect. +func Client(origin Origin, pinned Resolution, deadlines Deadlines) *http.Client { + dialer := &net.Dialer{Timeout: deadlines.Connect} + pinnedAddr := net.JoinHostPort(pinned.IP.String(), strconv.FormatUint(uint64(origin.Port), 10)) + transport := &http.Transport{ + Proxy: nil, + DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) { + return dialer.DialContext(ctx, network, pinnedAddr) + }, + ForceAttemptHTTP2: true, + TLSClientConfig: &tls.Config{ServerName: origin.Host, MinVersion: tls.VersionTLS12}, + TLSHandshakeTimeout: deadlines.TLS, + ResponseHeaderTimeout: deadlines.ResponseHeader, + ExpectContinueTimeout: time.Second, + MaxIdleConns: 1, + IdleConnTimeout: deadlines.Connect, + } + return &http.Client{ + Transport: transport, + CheckRedirect: func(_ *http.Request, _ []*http.Request) error { + return http.ErrUseLastResponse + }, + } +} + +// Deadlines bounds each phase of an outbound request. +type Deadlines struct { + Connect time.Duration + TLS time.Duration + ResponseHeader time.Duration +} + +// DefaultDeadlines are conservative bounds suitable for vendor management APIs. +func DefaultDeadlines() Deadlines { + return Deadlines{Connect: 5 * time.Second, TLS: 5 * time.Second, ResponseHeader: 15 * time.Second} +} diff --git a/network/urlguard/urlguard_test.go b/network/urlguard/urlguard_test.go new file mode 100644 index 00000000..6a5dbb80 --- /dev/null +++ b/network/urlguard/urlguard_test.go @@ -0,0 +1,214 @@ +package urlguard_test + +import ( + "context" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/codefly-dev/core/network/urlguard" + "github.com/stretchr/testify/require" +) + +func TestNormalizeOrigin_Canonicalization(t *testing.T) { + cases := []struct { + name string + raw string + want urlguard.Origin + }{ + {"lowercases host", "https://API.Example.COM", urlguard.Origin{Scheme: "https", Host: "api.example.com", Port: 443}}, + {"strips trailing dot", "https://api.example.com.", urlguard.Origin{Scheme: "https", Host: "api.example.com", Port: 443}}, + {"explicit port", "https://api.example.com:8443", urlguard.Origin{Scheme: "https", Host: "api.example.com", Port: 8443}}, + {"http default port", "http://api.example.com", urlguard.Origin{Scheme: "http", Host: "api.example.com", Port: 80}}, + {"ipv6 literal", "https://[2606:4700:4700::1111]", urlguard.Origin{Scheme: "https", Host: "2606:4700:4700::1111", Port: 443}}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + got, err := urlguard.NormalizeOrigin(tc.raw) + require.NoError(t, err) + require.Equal(t, tc.want, got) + }) + } +} + +func TestNormalizeOrigin_Rejections(t *testing.T) { + cases := map[string]string{ + "userinfo": "https://user:pass@api.example.com", + "userinfo no pass": "https://user@api.example.com", + "path": "https://api.example.com/v1", + "query": "https://api.example.com?x=1", + "fragment": "https://api.example.com#frag", + "scheme relative": "//api.example.com", + "non http scheme": "file://api.example.com", + "ftp scheme": "ftp://api.example.com", + "empty": "", + "opaque": "mailto:x@example.com", + "unicode host": "https://exämple.com", // must be submitted as punycode, not raw unicode + } + for name, raw := range cases { + t.Run(name, func(t *testing.T) { + _, err := urlguard.NormalizeOrigin(raw) + require.Error(t, err, "expected %q to be rejected", raw) + }) + } +} + +func TestSafePath(t *testing.T) { + ok := []string{"/", "/v1/resources", "/v1/resources/abc-123"} + for _, p := range ok { + got, err := urlguard.SafePath(p) + require.NoError(t, err, p) + require.NotEmpty(t, got) + } + bad := []string{"relative", "/a/../b", "/a/%2e%2e/b", "/a%2fb", "/a?x=1", "/a#f", "/a//b", "/a\\b"} + for _, p := range bad { + _, err := urlguard.SafePath(p) + require.Error(t, err, p) + } +} + +func TestClassify(t *testing.T) { + cases := map[string]urlguard.NetworkClass{ + "8.8.8.8": urlguard.ClassPublic, + "127.0.0.1": urlguard.ClassLoopback, + "::1": urlguard.ClassLoopback, + "169.254.169.254": urlguard.ClassLinkLocal, // cloud metadata + "fe80::1": urlguard.ClassLinkLocal, + "10.1.2.3": urlguard.ClassPrivate, + "192.168.0.1": urlguard.ClassPrivate, + "172.16.0.1": urlguard.ClassPrivate, + "100.64.0.1": urlguard.ClassPrivate, // CGNAT + "fc00::1": urlguard.ClassPrivate, // ULA + "0.0.0.0": urlguard.ClassPrivate, // unspecified + "239.0.0.1": urlguard.ClassPrivate, // administratively-scoped multicast + "224.0.0.1": urlguard.ClassLinkLocal, // link-local multicast + } + for raw, want := range cases { + require.Equal(t, want, urlguard.Classify(net.ParseIP(raw)), raw) + } +} + +func TestResolve_LiteralPrivateRejectedUnlessAllowed(t *testing.T) { + origin, err := urlguard.NormalizeOrigin("https://10.0.0.5") + require.NoError(t, err) + + _, err = urlguard.Resolve(context.Background(), nil, origin, nil) + require.Error(t, err) + + res, err := urlguard.Resolve(context.Background(), nil, origin, []urlguard.NetworkClass{urlguard.ClassPrivate}) + require.NoError(t, err) + require.Equal(t, urlguard.ClassPrivate, res.Class) +} + +func TestResolve_MetadataRejected(t *testing.T) { + origin, err := urlguard.NormalizeOrigin("http://169.254.169.254") + require.NoError(t, err) + // Not admitted even with private allowed: metadata is link-local. + _, err = urlguard.Resolve(context.Background(), nil, origin, []urlguard.NetworkClass{urlguard.ClassPrivate}) + require.Error(t, err) +} + +func TestResolve_LocalhostRequiresLoopbackAdmission(t *testing.T) { + origin, err := urlguard.NormalizeOrigin("http://localhost") + require.NoError(t, err) + + _, err = urlguard.Resolve(context.Background(), nil, origin, nil) + require.Error(t, err) + + res, err := urlguard.Resolve(context.Background(), nil, origin, []urlguard.NetworkClass{urlguard.ClassLoopback}) + require.NoError(t, err) + require.Equal(t, urlguard.ClassLoopback, res.Class) +} + +// TestClient_PinsPeerIP proves the client dials the pinned IP regardless of +// the host in the request URL — the DNS-rebinding defense — and that no proxy +// environment variable can divert it. +func TestClient_PinsPeerIP(t *testing.T) { + t.Setenv("HTTP_PROXY", "http://127.0.0.1:9") // canary: must be ignored + t.Setenv("HTTPS_PROXY", "http://127.0.0.1:9") + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + })) + defer server.Close() + + host, portStr, err := net.SplitHostPort(strings.TrimPrefix(server.URL, "http://")) + require.NoError(t, err) + port := mustPort(t, portStr) + + origin := urlguard.Origin{Scheme: "http", Host: "pinned.invalid", Port: port} + pinned := urlguard.Resolution{IP: net.ParseIP(host), Class: urlguard.ClassLoopback} + client := urlguard.Client(origin, pinned, urlguard.DefaultDeadlines()) + + transport, ok := client.Transport.(*http.Transport) + require.True(t, ok) + require.Nil(t, transport.Proxy, "transport must have no proxy") + + // URL host is a name that does not resolve to the server; the pinned dial + // reaches it anyway. + resp, err := client.Get(origin.String() + "/anything") + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) +} + +func TestClient_RefusesRedirects(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/start" { + http.Redirect(w, r, "/other", http.StatusFound) + return + } + w.WriteHeader(http.StatusTeapot) + })) + defer server.Close() + + host, portStr, err := net.SplitHostPort(strings.TrimPrefix(server.URL, "http://")) + require.NoError(t, err) + origin := urlguard.Origin{Scheme: "http", Host: "pinned.invalid", Port: mustPort(t, portStr)} + pinned := urlguard.Resolution{IP: net.ParseIP(host), Class: urlguard.ClassLoopback} + client := urlguard.Client(origin, pinned, urlguard.DefaultDeadlines()) + + resp, err := client.Get(origin.String() + "/start") + require.NoError(t, err) + defer resp.Body.Close() + // The redirect is surfaced, never followed. + require.Equal(t, http.StatusFound, resp.StatusCode) +} + +func TestCeiling_Admit(t *testing.T) { + ceiling := urlguard.Ceiling{ + Schemes: []string{"https"}, + HostPatterns: []string{"api.stripe.com", "*.sentry.io"}, + Ports: []uint32{443}, + } + ok := []string{"https://api.stripe.com", "https://o1.ingest.sentry.io", "https://API.Stripe.com."} + for _, raw := range ok { + _, err := ceiling.Admit(raw) + require.NoError(t, err, raw) + } + bad := []string{ + "http://api.stripe.com", // scheme + "https://api.stripe.com:8443", // port + "https://evil.com", // host + "https://sentry.io", // wildcard apex not matched + "https://api.stripe.com.evil.com", + "https://user@api.stripe.com", // userinfo + } + for _, raw := range bad { + _, err := ceiling.Admit(raw) + require.Error(t, err, raw) + } +} + +func mustPort(t *testing.T, s string) uint32 { + t.Helper() + var n uint32 + for _, r := range s { + require.True(t, r >= '0' && r <= '9') + n = n*10 + uint32(r-'0') + } + return n +} diff --git a/provider/broker/broker.go b/provider/broker/broker.go new file mode 100644 index 00000000..79a5c52d --- /dev/null +++ b/provider/broker/broker.go @@ -0,0 +1,417 @@ +// Package broker is the trusted host-mediated HTTP boundary for provider +// agents. Provider code never dials the network; it hands the host a structured +// request from an already-admitted plan action, and the broker re-derives the +// exact action, resource, origin, headers, and credentials from host-owned +// context before any bytes leave the machine. Provider descriptors are treated +// as adversarial: the action, resource, read-only intent, and origin the +// provider claims are all ignored in favor of what the host planned. +// +// The broker never retries. It reports one of NOT_SENT, SENT_OUTCOME_UNKNOWN, +// or RESPONSE_RECEIVED so the coordinator can recover honestly, and it filters +// every response through the manifest response policy so secret-bearing fields +// are captured or suppressed rather than forwarded. +package broker + +import ( + "context" + "fmt" + "net" + "net/http" + "net/http/httptrace" + "sync" + "time" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" + "github.com/codefly-dev/core/provider/canonical" + "github.com/codefly-dev/core/provider/cassette" + "github.com/codefly-dev/core/provider/credentials" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" +) + +// Checkpointer reports the durable recovery checkpoint the coordinator persisted +// for an operation. The durable store itself is out of scope (F3); the broker +// depends only on this read-side to enforce checkpoint-before-send ordering. +type Checkpointer interface { + Latest(ctx context.Context, operation *providerv0.OperationIdentity) (*providerv0.ActionCheckpoint, error) +} + +// Config binds one broker session to exactly one admitted plan action. Every +// field is host-owned; the provider cannot override any of it. +type Config struct { + Manifest *manifest.Manifest + Action *providerv0.PlanAction + Binding *providerv0.BindingAddress + Budget *providerv0.RequestBudget + ReadOnly bool + + Vault *credentials.Vault + Sink responsepolicy.Sink + Checkpoints Checkpointer + Resolver *net.Resolver + Deadlines urlguard.Deadlines + UserAgent string + + // Cassette, when set, records or replays this session's responses. Replay + // never touches the network; record stores only filtered safe responses. + Cassette *cassette.Cassette + + // ClientFor overrides guarded client construction. Production leaves it nil + // to use the SSRF-hardened urlguard client; tests inject a client bound to a + // local server. + ClientFor func(urlguard.Origin, urlguard.Resolution) *http.Client + // Now overrides the clock for deterministic tests. + Now func() time.Time +} + +// Session is the per-action broker context. A session serializes its requests: +// budget, capture-gate, and delivery are inherently sequential, so Execute holds +// a session lock for the duration of each call. +type Session struct { + mu sync.Mutex + + action *providerv0.PlanAction + binding *providerv0.BindingAddress + readOnly bool + userAgent string + + requests map[string]manifest.RequestDescriptor + originRules map[string]manifest.OriginRule + responseSchemas map[string]manifest.ResponseSchema + purposes map[string]manifest.CredentialPurpose + + budget *providerv0.RequestBudget + limits responsepolicy.Limits + remaining uint32 + + vault *credentials.Vault + sink responsepolicy.Sink + checkpoints Checkpointer + cassette *cassette.Cassette + resolver *net.Resolver + clientFor func(urlguard.Origin, urlguard.Resolution) *http.Client + now func() time.Time + + // captureGate holds the checkpoint id present when the last capture became + // durable. A later external request is refused until a newer checkpoint + // acknowledges the capture. + captureGate string +} + +// New builds a session and validates that the action belongs to the manifest. +func New(cfg Config) (*Session, error) { + if cfg.Manifest == nil || cfg.Action == nil || cfg.Budget == nil { + return nil, fmt.Errorf("broker requires a manifest, action, and budget") + } + if cfg.Vault == nil || cfg.Sink == nil || cfg.Checkpoints == nil { + return nil, fmt.Errorf("broker requires a vault, sink, and checkpointer") + } + if err := canonical.ValidatePlanAction(cfg.Action); err != nil { + return nil, err + } + session := &Session{ + action: cfg.Action, + binding: cfg.Binding, + readOnly: cfg.ReadOnly, + userAgent: cfg.UserAgent, + requests: make(map[string]manifest.RequestDescriptor), + originRules: make(map[string]manifest.OriginRule), + responseSchemas: make(map[string]manifest.ResponseSchema), + purposes: make(map[string]manifest.CredentialPurpose), + budget: cfg.Budget, + remaining: cfg.Budget.GetRequestCount(), + vault: cfg.Vault, + sink: cfg.Sink, + checkpoints: cfg.Checkpoints, + cassette: cfg.Cassette, + resolver: cfg.Resolver, + clientFor: cfg.ClientFor, + now: cfg.Now, + } + if session.userAgent == "" { + session.userAgent = "codefly-provider-broker" + } + if session.now == nil { + session.now = time.Now + } + if session.clientFor == nil { + session.clientFor = func(origin urlguard.Origin, pinned urlguard.Resolution) *http.Client { + return urlguard.Client(origin, pinned, cfg.Deadlines) + } + } + for _, descriptor := range cfg.Manifest.Requests { + session.requests[descriptor.ID] = descriptor + } + for _, rule := range cfg.Manifest.OriginRules { + session.originRules[rule.ID] = rule + } + for _, schema := range cfg.Manifest.ResponseSchemas { + session.responseSchemas[schema.ID] = schema + } + for _, purpose := range cfg.Manifest.CredentialPurposes { + session.purposes[purpose.ID] = purpose + } + session.limits = responsePolicyLimits(cfg.Budget) + return session, nil +} + +func responsePolicyLimits(budget *providerv0.RequestBudget) responsepolicy.Limits { + limits := responsepolicy.DefaultLimits() + if bytes := int64(budget.GetResponseBytes()); bytes > 0 { + limits.MaxCompressedBytes = bytes + limits.MaxDecompressedBytes = bytes + } + return limits +} + +// Execute runs one admitted broker request. It returns a filtered response and, +// on any admission or filtering failure, a non-nil error the coordinator treats +// as a hard stop. Original response bytes are never forwarded. +func (s *Session) Execute(ctx context.Context, request *providerv0.ExecuteRequestRequest) (*providerv0.ExecuteRequestResponse, error) { + s.mu.Lock() + defer s.mu.Unlock() + + if err := canonical.ValidateExecuteRequest(s.action, request); err != nil { + return nil, fmt.Errorf("admit request: %w", err) + } + planned := request.GetRequest() + descriptor, err := s.admitDescriptor(planned) + if err != nil { + return nil, err + } + if s.readOnly && !descriptor.ReadOnly { + return nil, fmt.Errorf("read-only context cannot execute a mutating request") + } + if err := s.checkBudget(); err != nil { + return nil, err + } + origin, ceiling, err := s.admitOrigin(descriptor, request.GetOrigin()) + if err != nil { + return nil, err + } + // SSRF, DNS-rebinding, and private-address checks must complete before any + // credential is resolved. Replay is enforcement-identical but performs no + // network I/O, so it resolves and pins nothing. + replay := s.replaying() + var pinned urlguard.Resolution + if !replay { + if pinned, err = s.resolveAndPin(ctx, origin, ceiling, request.GetOrigin()); err != nil { + return nil, err + } + } + remoteID, err := plannedRemoteID(s.action) + if err != nil { + return nil, err + } + httpRequest, bodySize, err := s.buildRequest(ctx, descriptor, planned, origin, remoteID) + if err != nil { + return nil, err + } + checkpointID, err := s.requireCheckpoint(ctx, request.GetContext().GetOperation(), planned) + if err != nil { + return nil, err + } + if err := s.injectCredentials(httpRequest, request, origin); err != nil { + return nil, err + } + // The whole outbound message — request line, host-owned headers including the + // resolved credential, and body — is bounded by the request byte budget. + if limit := s.budget.GetRequestBytes(); limit != 0 && outboundSize(httpRequest, bodySize) > int64(limit) { + return nil, fmt.Errorf("request exceeds byte budget") + } + + // The callback is consumed once the request is fully admitted and about to be + // delivered. + s.remaining-- + + delivery, response, err := s.deliver(ctx, descriptor, planned, request.GetOrigin(), origin, pinned, httpRequest) + if err != nil { + return nil, err + } + // Arm the capture gate: a durable capture must be checkpointed before the + // next external request, whether it was produced live or replayed. + if len(response.GetCaptures()) > 0 { + s.captureGate = checkpointID + } + response.RequestId = request.GetRequestId() + response.Delivery = delivery + return response, nil +} + +// replaying reports whether this session serves responses from a cassette. +func (s *Session) replaying() bool { + return s.cassette != nil && s.cassette.Mode() == cassette.ModeReplay +} + +// deliver performs the request either live or through a cassette. Replay never +// touches the network; record stores the filtered response. All admission +// checks have already run, so both paths share identical validation. +func (s *Session) deliver( + ctx context.Context, + descriptor manifest.RequestDescriptor, + planned *providerv0.PlannedRequest, + attested *providerv0.AdmittedOrigin, + origin urlguard.Origin, + pinned urlguard.Resolution, + httpRequest *http.Request, +) (providerv0.DeliveryState, *providerv0.ExecuteRequestResponse, error) { + if s.cassette != nil && s.cassette.Mode() == cassette.ModeReplay { + status, response, err := s.cassette.Replay(cassette.NewKey(planned, attested)) + if err != nil { + return 0, nil, err + } + response.StatusCode = status + return providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, response, nil + } + client := s.clientFor(origin, pinned) + delivery, response, err := s.roundTrip(ctx, client, descriptor, httpRequest) + if err != nil { + return 0, nil, err + } + if s.cassette != nil && s.cassette.Mode() == cassette.ModeRecord && delivery == providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED { + if err := s.cassette.Record(cassette.NewKey(planned, attested), response.GetStatusCode(), nil, response); err != nil { + return 0, nil, fmt.Errorf("record cassette: %w", err) + } + } + return delivery, response, nil +} + +func (s *Session) admitDescriptor(planned *providerv0.PlannedRequest) (manifest.RequestDescriptor, error) { + descriptor, ok := s.requests[planned.GetRequestDescriptorId()] + if !ok { + return manifest.RequestDescriptor{}, fmt.Errorf("request descriptor %q is not packaged", planned.GetRequestDescriptorId()) + } + digest, err := manifest.RequestDescriptorDigest(descriptor) + if err != nil { + return manifest.RequestDescriptor{}, err + } + if digest != planned.GetRequestDescriptorDigest() { + return manifest.RequestDescriptor{}, fmt.Errorf("request descriptor digest mismatch") + } + return descriptor, nil +} + +func (s *Session) checkBudget() error { + if s.remaining == 0 { + return fmt.Errorf("request budget is exhausted") + } + if deadline := s.budget.GetDeadline(); deadline != nil && s.now().After(deadline.AsTime()) { + return fmt.Errorf("request budget deadline has passed") + } + return nil +} + +// admitOrigin re-admits the host-attested origin against the descriptor's +// ceiling, resolves and pins a single peer IP, and confirms the attested +// network class matches the resolved one. +// admitOrigin performs the structural, network-free admission: it binds the +// attested origin to the descriptor's rule and confirms it stays within the +// manifest ceiling. It never resolves DNS, so it runs identically for live and +// replay. +func (s *Session) admitOrigin(descriptor manifest.RequestDescriptor, attested *providerv0.AdmittedOrigin) (urlguard.Origin, urlguard.Ceiling, error) { + if attested.GetOriginRuleId() != descriptor.OriginRule { + return urlguard.Origin{}, urlguard.Ceiling{}, fmt.Errorf("origin rule does not match descriptor") + } + rule, ok := s.originRules[descriptor.OriginRule] + if !ok { + return urlguard.Origin{}, urlguard.Ceiling{}, fmt.Errorf("origin rule %q is not packaged", descriptor.OriginRule) + } + origin := urlguard.Origin{Scheme: attested.GetScheme(), Host: attested.GetHost(), Port: attested.GetPort()} + ceiling := ceilingFor(rule) + if err := ceiling.AdmitOrigin(origin); err != nil { + return urlguard.Origin{}, urlguard.Ceiling{}, err + } + return origin, ceiling, nil +} + +// resolveAndPin performs the single DNS resolution in the request path, pins one +// admitted peer IP, and confirms the resolved network class matches the +// host-attested class. It runs only for live delivery — never during replay — +// and always before any credential is resolved. +func (s *Session) resolveAndPin(ctx context.Context, origin urlguard.Origin, ceiling urlguard.Ceiling, attested *providerv0.AdmittedOrigin) (urlguard.Resolution, error) { + pinned, err := urlguard.Resolve(ctx, s.resolver, origin, ceiling.AllowedClasses) + if err != nil { + return urlguard.Resolution{}, err + } + if classEnum(pinned.Class) != attested.GetPrivateNetworkClass() { + return urlguard.Resolution{}, fmt.Errorf("resolved network class does not match attested origin") + } + return pinned, nil +} + +// requireCheckpoint enforces that a durable pre-send checkpoint exists before +// any bytes leave, and that a prior durable capture has been acknowledged by a +// newer checkpoint before the next external request. +func (s *Session) requireCheckpoint(ctx context.Context, operation *providerv0.OperationIdentity, planned *providerv0.PlannedRequest) (string, error) { + checkpoint, err := s.checkpoints.Latest(ctx, operation) + if err != nil { + return "", fmt.Errorf("read checkpoint: %w", err) + } + if checkpoint == nil { + return "", fmt.Errorf("a durable checkpoint is required before sending") + } + if checkpoint.GetOperation().GetActionId() != s.action.GetActionId() { + return "", fmt.Errorf("checkpoint does not bind this action") + } + if isMutating(planned.GetMethod()) && checkpoint.GetIdempotencyKey() != planned.GetIdempotencyKey() { + return "", fmt.Errorf("checkpoint idempotency key does not match request") + } + if s.captureGate != "" && checkpoint.GetCheckpointId() == s.captureGate { + return "", fmt.Errorf("a durable capture must be checkpointed before the next request") + } + return checkpoint.GetCheckpointId(), nil +} + +func (s *Session) injectCredentials(httpRequest *http.Request, request *providerv0.ExecuteRequestRequest, origin urlguard.Origin) error { + use := credentials.Use{ + Binding: s.binding, + ActionID: s.action.GetActionId(), + RequestDigest: request.GetRequest().GetRequestDigest(), + Origin: origin, + Method: request.GetRequest().GetMethod(), + } + for _, handle := range request.GetCredentialHandles() { + use.Purpose = handle.GetPurpose() + if err := s.vault.Inject(httpRequest, handle.GetHandle(), use); err != nil { + return fmt.Errorf("resolve credential: %w", err) + } + } + return nil +} + +// roundTrip performs exactly one request with no retry. Request tracing tells +// whether request bytes crossed the boundary so a lost response is reported as +// SENT_OUTCOME_UNKNOWN rather than NOT_SENT. +func (s *Session) roundTrip(ctx context.Context, client *http.Client, descriptor manifest.RequestDescriptor, httpRequest *http.Request) (providerv0.DeliveryState, *providerv0.ExecuteRequestResponse, error) { + // The absolute budget deadline bounds the whole round trip, not just the + // pre-send admission checks. + if deadline := s.budget.GetDeadline(); deadline != nil { + var cancel context.CancelFunc + ctx, cancel = context.WithDeadline(ctx, deadline.AsTime()) + defer cancel() + } + var wrote bool + trace := &httptrace.ClientTrace{WroteRequest: func(httptrace.WroteRequestInfo) { wrote = true }} + traced := httpRequest.WithContext(httptrace.WithClientTrace(ctx, trace)) + + resp, err := client.Do(traced) + if err != nil { + if wrote { + return providerv0.DeliveryState_DELIVERY_STATE_SENT_OUTCOME_UNKNOWN, + &providerv0.ExecuteRequestResponse{Certainty: providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_UNCERTAIN}, nil + } + certainty := providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE + if descriptor.ReadOnly { + certainty = providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_UNCERTAIN + } + return providerv0.DeliveryState_DELIVERY_STATE_NOT_SENT, + &providerv0.ExecuteRequestResponse{Certainty: certainty}, nil + } + defer resp.Body.Close() + response, err := s.handleResponse(descriptor, resp) + if err != nil { + return providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, nil, err + } + return providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, response, nil +} diff --git a/provider/broker/broker_test.go b/provider/broker/broker_test.go new file mode 100644 index 00000000..8ff8df78 --- /dev/null +++ b/provider/broker/broker_test.go @@ -0,0 +1,263 @@ +package broker_test + +import ( + "context" + "encoding/json" + "fmt" + "net" + "net/http" + "net/http/httptest" + "testing" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/broker" + "github.com/codefly-dev/core/provider/manifest" + "github.com/stretchr/testify/require" +) + +func accountServer(t *testing.T, captured *http.Request) *httptest.Server { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if captured != nil { + *captured = *r.Clone(context.Background()) + } + w.Header().Set("Content-Type", "application/json") + _, _ = fmt.Fprintf(w, `{"id":%q,"secret":%q,"metadata":{"internal":"private"}}`, remoteID, poisonSecret) + })) + t.Cleanup(server.Close) + return server +} + +func serverAddr(t *testing.T, server *httptest.Server) string { + t.Helper() + return server.Listener.Addr().String() +} + +func TestExecute_CreateFiltersAndCapturesSecret(t *testing.T) { + h := newHarness(t) + var sent http.Request + server := accountServer(t, &sent) + + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + response, err := session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.NoError(t, err) + + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, response.GetDelivery()) + require.Equal(t, uint32(http.StatusOK), response.GetStatusCode()) + + // The secret was captured to the sink and never forwarded. + require.Equal(t, []string{poisonSecret}, h.sink.stored) + require.Len(t, response.GetCaptures(), 1) + require.Len(t, response.GetSuppressedPresence(), 1) + + // $.id is the only forwarded field. + require.Len(t, response.GetForwarded(), 1) + require.Equal(t, remoteID, response.GetForwarded()[0].GetValue().GetStringValue()) + + // The poison secret is nowhere in the provider-facing response. + encoded, err := json.Marshal(response) + require.NoError(t, err) + require.NotContains(t, string(encoded), poisonSecret) + + // The host owned every header: Authorization is the injected bearer and no + // caller/proxy headers exist. + require.Equal(t, "Bearer "+poisonSecret, sent.Header.Get("Authorization")) + require.Equal(t, "identity", sent.Header.Get("Accept-Encoding")) + require.Equal(t, "idem-1", sent.Header.Get("Idempotency-Key")) + require.Equal(t, "codefly-provider-broker", sent.Header.Get("User-Agent")) +} + +func TestExecute_RejectsMutatingRequestInReadOnlyContext(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(true) // read-only context + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.ErrorContains(t, err, "read-only") + require.Empty(t, h.sink.stored) +} + +func TestExecute_WrongResourceIDInPathRejected(t *testing.T) { + m := mustLoad(t) + origin := admittedOrigin(t) + // Observe request whose path id is not the planned prospective id. + observe := boundObserveRequest(t, m, origin, "acct_999") + action := createAction(t, observe) + + h := newHarness(t) + h.manifest, h.action, h.admitted = m, action, origin + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, observe, providerv0.HTTPMethod_HTTP_METHOD_GET) + _, err = session.Execute(context.Background(), h.executeRequest(handle, observe)) + require.ErrorContains(t, err, "planned remote id") + require.Empty(t, h.sink.stored) +} + +func TestExecute_MissingCheckpointRefusesSend(t *testing.T) { + h := newHarness(t) + h.checkpoints.checkpoint = nil // no durable checkpoint + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.ErrorContains(t, err, "durable checkpoint is required") +} + +func TestExecute_BudgetExhausted(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.Budget = &providerv0.RequestBudget{RequestCount: 1, RequestBytes: 8192, ResponseBytes: 65536} + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.NoError(t, err) + + // Budget is now exhausted; a second call is refused before send. + handle2 := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle2, h.create)) + require.ErrorContains(t, err, "budget is exhausted") +} + +func TestExecute_TimeoutBeforeSendIsNotSent(t *testing.T) { + h := newHarness(t) + // Dial a port with no listener: connection refused before any bytes are sent. + closed := reservedClosedAddr(t) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(closed) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + response, err := session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.NoError(t, err) + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_NOT_SENT, response.GetDelivery()) +} + +func TestExecute_LostResponseIsSentOutcomeUnknown(t *testing.T) { + // A server that reads the full request then closes without responding: the + // request bytes crossed the boundary but no response returned. + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = listener.Close() }) + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + go func() { + // Read the request bytes off the wire (a single read is enough + // for a small request) and close without responding, so the + // request crosses the boundary but no response returns. + buffer := make([]byte, 4096) + _, _ = conn.Read(buffer) + _ = conn.Close() + }() + } + }() + + h := newHarness(t) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(listener.Addr().String()) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + response, err := session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.NoError(t, err) + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_SENT_OUTCOME_UNKNOWN, response.GetDelivery()) + require.Equal(t, providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_UNCERTAIN, response.GetCertainty()) +} + +func TestExecute_CaptureMustBeCheckpointedBeforeNext(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + handle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle, h.create)) + require.NoError(t, err) + require.Equal(t, []string{poisonSecret}, h.sink.stored) + + // The capture gate is armed: a second request against the same checkpoint is + // refused until a newer checkpoint acknowledges the capture. + handle2 := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle2, h.create)) + require.ErrorContains(t, err, "capture must be checkpointed") + + // A newer checkpoint releases the gate. + h.checkpoints.checkpoint = checkpoint("cp2", "idem-1") + handle3 := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + _, err = session.Execute(context.Background(), h.executeRequest(handle3, h.create)) + require.NoError(t, err) +} + +func TestExecute_CredentialCannotBeReusedOnAnotherRequest(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + // Mint a handle for the create request, then try to spend it on a different + // (observe) request: the request digest no longer matches the handle. + m := mustLoad(t) + observe := boundObserveRequest(t, m, h.admitted, remoteID) + h.action = createAction(t, h.create, observe) + cfg = h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err = broker.New(cfg) + require.NoError(t, err) + + createHandle := h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + // Present the create handle while executing the observe request. + request := h.executeRequest(createHandle, observe) + request.Request = observe + _, err = session.Execute(context.Background(), request) + require.Error(t, err) + require.Empty(t, h.sink.stored) +} + +func mustLoad(t *testing.T) *manifest.Manifest { + t.Helper() + m, err := manifest.Load([]byte(brokerManifestYAML(loopbackConfig()))) + require.NoError(t, err) + return m +} + +// reservedClosedAddr returns a loopback address that is bound then immediately +// closed, so a dial to it is refused. +func reservedClosedAddr(t *testing.T) string { + t.Helper() + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + addr := listener.Addr().String() + require.NoError(t, listener.Close()) + return addr +} diff --git a/provider/broker/cassette_test.go b/provider/broker/cassette_test.go new file mode 100644 index 00000000..ad9ed224 --- /dev/null +++ b/provider/broker/cassette_test.go @@ -0,0 +1,72 @@ +package broker_test + +import ( + "context" + "testing" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/broker" + "github.com/codefly-dev/core/provider/cassette" + "github.com/stretchr/testify/require" +) + +// TestExecute_RecordThenReplayIsEnforcementIdentical records a live create +// response, then replays it through the same broker admission path without any +// network — the replayed filtered response is identical to the recorded one and +// replay never falls back to live. +func TestExecute_RecordThenReplayIsEnforcementIdentical(t *testing.T) { + // Record phase: live server, cassette in record mode. + rec := newHarness(t) + server := accountServer(t, nil) + recordCassette := cassette.New(cassette.ModeRecord, "1.2.3") + recCfg := rec.config(false) + recCfg.ClientFor = dialClientFor(serverAddr(t, server)) + recCfg.Cassette = recordCassette + recSession, err := broker.New(recCfg) + require.NoError(t, err) + + recorded, err := recSession.Execute(context.Background(), rec.executeRequest(rec.mintHandle(t, rec.create, providerv0.HTTPMethod_HTTP_METHOD_POST), rec.create)) + require.NoError(t, err) + require.Equal(t, []string{poisonSecret}, rec.sink.stored) + + data, err := recordCassette.Marshal() + require.NoError(t, err) + require.NotContains(t, string(data), poisonSecret, "cassette must not persist the secret") + + // Replay phase: no live server reachable (closed addr), cassette in replay. + replayCassette, err := cassette.Load(data, "1.2.3") + require.NoError(t, err) + rep := newHarness(t) + repCfg := rep.config(false) + repCfg.ClientFor = dialClientFor(reservedClosedAddr(t)) // proves no live fallback + repCfg.Cassette = replayCassette + repSession, err := broker.New(repCfg) + require.NoError(t, err) + + replayed, err := repSession.Execute(context.Background(), rep.executeRequest(rep.mintHandle(t, rep.create, providerv0.HTTPMethod_HTTP_METHOD_POST), rep.create)) + require.NoError(t, err) + + // The replay produced no live capture (the sink was untouched) yet returned + // the identical filtered forwarded fields and capture references. + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, replayed.GetDelivery()) + require.Equal(t, recorded.GetStatusCode(), replayed.GetStatusCode()) + require.Equal(t, len(recorded.GetForwarded()), len(replayed.GetForwarded())) + require.Equal(t, recorded.GetForwarded()[0].GetValue().GetStringValue(), replayed.GetForwarded()[0].GetValue().GetStringValue()) + require.Equal(t, len(recorded.GetCaptures()), len(replayed.GetCaptures())) + require.Equal(t, recorded.GetSuppressedPresence(), replayed.GetSuppressedPresence()) +} + +// TestExecute_ReplayMissingEntryDoesNotHitNetwork proves an unrecorded request +// fails closed rather than reaching the network. +func TestExecute_ReplayMissingEntryDoesNotHitNetwork(t *testing.T) { + empty := cassette.New(cassette.ModeReplay, "1.2.3") + h := newHarness(t) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(reservedClosedAddr(t)) + cfg.Cassette = empty + session, err := broker.New(cfg) + require.NoError(t, err) + + _, err = session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.ErrorContains(t, err, "does not fall back to live") +} diff --git a/provider/broker/derive.go b/provider/broker/derive.go new file mode 100644 index 00000000..2ec83dab --- /dev/null +++ b/provider/broker/derive.go @@ -0,0 +1,117 @@ +package broker + +import ( + "fmt" + "strings" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" +) + +// purposeEnum maps a manifest credential-purpose consumer to the proto enum. +func purposeEnum(consumer string) (providerv0.CredentialPurpose, error) { + switch consumer { + case "management": + return providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, nil + case "runtime": + return providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_RUNTIME, nil + case "build": + return providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_BUILD, nil + case "webhook-verification": + return providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_WEBHOOK_VERIFICATION, nil + default: + return providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_UNSPECIFIED, fmt.Errorf("unknown credential purpose consumer %q", consumer) + } +} + +// classEnum maps a manifest private-network-class token to the proto enum. +func classEnum(class urlguard.NetworkClass) providerv0.PrivateNetworkClass { + switch class { + case urlguard.ClassLoopback: + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_LOOPBACK + case urlguard.ClassLinkLocal: + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_LINK_LOCAL + case urlguard.ClassPrivate: + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_PRIVATE + default: + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_PUBLIC + } +} + +// networkClasses maps manifest class tokens to urlguard classes. +func networkClasses(tokens []string) []urlguard.NetworkClass { + classes := make([]urlguard.NetworkClass, 0, len(tokens)) + for _, token := range tokens { + switch token { + case "loopback": + classes = append(classes, urlguard.ClassLoopback) + case "link-local": + classes = append(classes, urlguard.ClassLinkLocal) + case "private": + classes = append(classes, urlguard.ClassPrivate) + } + } + return classes +} + +// ceilingFor builds the urlguard ceiling for a manifest origin rule. +func ceilingFor(rule manifest.OriginRule) urlguard.Ceiling { + return urlguard.Ceiling{ + Schemes: append([]string(nil), rule.Schemes...), + HostPatterns: append([]string(nil), rule.HostPatterns...), + Ports: append([]uint32(nil), rule.Ports...), + AllowedClasses: networkClasses(rule.PrivateNetworkClasses), + } +} + +// responsePolicyFor derives the host response policy from the manifest response +// schema a descriptor references. Capture selectors that address an exact path +// are required (a planned single secret that vanishes is drift); wildcard +// captures are optional because an empty collection legitimately has none. +func (s *Session) responsePolicyFor(descriptor manifest.RequestDescriptor) (responsepolicy.Policy, error) { + schema, ok := s.responseSchemas[descriptor.ResponseSchema] + if !ok { + return responsepolicy.Policy{}, fmt.Errorf("response schema %q is not packaged", descriptor.ResponseSchema) + } + fields := make([]responsepolicy.Field, 0, len(schema.Fields)) + for _, field := range schema.Fields { + policyField := responsepolicy.Field{ + Selector: field.Selector, + Disposition: field.Disposition, + } + if field.Disposition == manifest.ResponseCaptureToSink { + purpose, ok := s.purposes[field.Purpose] + if !ok { + return responsepolicy.Policy{}, fmt.Errorf("capture purpose %q is not packaged", field.Purpose) + } + enum, err := purposeEnum(purpose.PermittedConsumer) + if err != nil { + return responsepolicy.Policy{}, err + } + policyField.Purpose = enum + policyField.Required = !strings.Contains(field.Selector.Path, "[*]") + policyField.SinkKey = fmt.Sprintf("%s/%s/%s", s.binding.GetBindingId(), descriptor.ID, field.Selector.Path) + } + fields = append(fields, policyField) + } + return responsepolicy.Policy{Fields: fields, Limits: s.limits}, nil +} + +// plannedRemoteID returns the exact remote identity the action authorizes. Path +// parameters and ownership body fields must equal this value; a valid template +// carrying any other id is rejected. +func plannedRemoteID(action *providerv0.PlanAction) (string, error) { + switch action.GetType() { + case providerv0.ActionType_ACTION_TYPE_CREATE, providerv0.ActionType_ACTION_TYPE_REPLACE: + if id := action.GetProspectiveRemoteId(); id != "" { + return id, nil + } + default: + if id := action.GetRemoteIdentity().GetRemoteId(); id != "" { + return id, nil + } + } + return "", fmt.Errorf("action has no bound remote identity") +} diff --git a/provider/broker/harness_test.go b/provider/broker/harness_test.go new file mode 100644 index 00000000..842d8f20 --- /dev/null +++ b/provider/broker/harness_test.go @@ -0,0 +1,458 @@ +package broker_test + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "fmt" + "net" + "net/http" + "testing" + "time" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" + "github.com/codefly-dev/core/provider/broker" + "github.com/codefly-dev/core/provider/canonical" + "github.com/codefly-dev/core/provider/credentials" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" + "github.com/stretchr/testify/require" +) + +// originConfig parameterizes the manifest origin rule and matching admitted +// origin so tests can target a loopback server or an unresolvable host. +type originConfig struct { + scheme string + host string + port uint32 + class string +} + +func loopbackConfig() originConfig { + return originConfig{scheme: "http", host: "localhost", port: 8080, class: "loopback"} +} + +func (o originConfig) urlguardOrigin() urlguard.Origin { + return urlguard.Origin{Scheme: o.scheme, Host: o.host, Port: o.port} +} + +func (o originConfig) networkClass() providerv0.PrivateNetworkClass { + switch o.class { + case "loopback": + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_LOOPBACK + case "link-local": + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_LINK_LOCAL + case "private": + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_PRIVATE + default: + return providerv0.PrivateNetworkClass_PRIVATE_NETWORK_CLASS_PUBLIC + } +} + +// brokerManifestFmt is a Stripe-like provider manifest whose origin rule is +// templated. It declares a read-only observe request and a mutating create +// request that returns a capturable secret. +const brokerManifestFmt = ` +schema_version: codefly.provider-manifest/v0 +protocol_version: codefly.provider/v0 +state_schema_versions: [1] +agent: + kind: codefly:provider + publisher: codefly.dev + name: fixture + version: 1.2.3 +default_deletion_policy: retain +permissions: + required: + - id: account-create + action: account.manage + resource: "provider:fixture/${workspace}/${environment}/${binding}/account" + resource_type: account + reason: Reconcile the declared account. + risk: high + credential_purpose: management + optional: + - id: account-observe + action: account.observe + resource: "provider:fixture/${workspace}/${environment}/${binding}/account" + resource_type: account + reason: Observe the declared account. + risk: low + credential_purpose: management + - id: account-delete + action: account.delete + resource: "provider:fixture/${workspace}/${environment}/${binding}/account" + resource_type: account + reason: Delete the declared account. + risk: critical + credential_purpose: management +resource_types: + - id: account + actions: [create, update, replace, delete, import, manual, blocked, no-op, project-output, observe] + import_identity: [account_id] + supports_replace: true + supports_delete: true +requests: + - id: account.observe + permissions: [account-observe] + resource_type: account + action: observe + origin_rule: api + operation: observe + method: GET + path_template: /v1/accounts/{account_id} + remote_id_parameters: [account_id] + allowed_query_fields: [expand] + request_byte_budget: 4096 + response_byte_budget: 65536 + read_only: true + response_schema: account + credential_purposes: [management] + - id: account.create + permissions: [account-create] + resource_type: account + action: create + origin_rule: api + operation: create + method: POST + path_template: /v1/accounts + remote_id_parameters: [] + allowed_body_fields: [name, enabled] + ownership_body_fields: [name] + request_byte_budget: 8192 + response_byte_budget: 65536 + read_only: false + response_schema: account + credential_purposes: [management] + - id: account.delete + permissions: [account-delete] + resource_type: account + action: delete + origin_rule: api + operation: delete + method: DELETE + path_template: /v1/accounts/{account_id} + remote_id_parameters: [account_id] + request_byte_budget: 4096 + response_byte_budget: 65536 + read_only: false + response_schema: deleted + credential_purposes: [management] +origin_rules: + - id: api + defaults: [%s://%s:%d] + schemes: [%s] + host_patterns: [%s] + ports: [%d] + binding_override: within-rule + private_network_classes: [%s] +credential_purposes: + - id: management + minimum_scope: Manage only the bound account. + permitted_consumer: management + - id: runtime + minimum_scope: Use the bound runtime account. + permitted_consumer: runtime +response_schemas: + - id: account + fields: + - selector: {version: v1, path: "$.id"} + disposition: FORWARD_SAFE + - selector: {version: v1, path: "$.secret"} + disposition: CAPTURE_TO_SINK + purpose: runtime + - selector: {version: v1, path: "$.metadata.internal"} + disposition: SUPPRESS_REPORT_PRESENCE + - id: deleted + fields: + - selector: {version: v1, path: "$.deleted"} + disposition: FORWARD_SAFE +sandbox: + network: deny +state: + schema_versions: [1] + import_identity: true + replace: true + delete: true + stepwise_upgrade: true +diagnostic_namespace: provider.fixture. +` + +func brokerManifestYAML(o originConfig) string { + return fmt.Sprintf(brokerManifestFmt, o.scheme, o.host, o.port, o.scheme, o.host, o.port, o.class) +} + +const ( + remoteID = "acct_123" + poisonSecret = "sk_live_1234567890abcdef" +) + +func admittedOriginFor(t *testing.T, o originConfig) *providerv0.AdmittedOrigin { + t.Helper() + origin := &providerv0.AdmittedOrigin{ + OriginRuleId: "api", + Scheme: o.scheme, + Host: o.host, + Port: o.port, + PrivateNetworkClass: o.networkClass(), + } + digest, err := canonical.AdmittedOriginDigest(origin) + require.NoError(t, err) + origin.AdmissionDigest = digest + return origin +} + +func admittedOrigin(t *testing.T) *providerv0.AdmittedOrigin { + t.Helper() + return admittedOriginFor(t, loopbackConfig()) +} + +func fakeDigest(seed string) string { + sum := sha256.Sum256([]byte(seed)) + return "sha256:" + hex.EncodeToString(sum[:]) +} + +func descriptorDigest(t *testing.T, m *manifest.Manifest, id string) string { + t.Helper() + for _, descriptor := range m.Requests { + if descriptor.ID == id { + digest, err := manifest.RequestDescriptorDigest(descriptor) + require.NoError(t, err) + return digest + } + } + t.Fatalf("descriptor %q not found", id) + return "" +} + +func pubString(value string) *providerv0.PublicValue { + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_StringValue{StringValue: value}} +} + +func boundCreateRequest(t *testing.T, m *manifest.Manifest, origin *providerv0.AdmittedOrigin) *providerv0.PlannedRequest { + t.Helper() + request := &providerv0.PlannedRequest{ + RequestDescriptorId: "account.create", + RequestDescriptorDigest: descriptorDigest(t, m, "account.create"), + Method: providerv0.HTTPMethod_HTTP_METHOD_POST, + AdmittedOriginDigest: origin.GetAdmissionDigest(), + Body: map[string]*providerv0.PublicValue{"name": pubString(remoteID)}, + CredentialPurposes: []providerv0.CredentialPurpose{providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT}, + ResponsePolicyDigest: fakeDigest("account-response-policy"), + IdempotencyKey: "idem-1", + } + bound, err := canonical.BindPlannedRequestDigest(request) + require.NoError(t, err) + return bound +} + +func boundObserveRequest(t *testing.T, m *manifest.Manifest, origin *providerv0.AdmittedOrigin, accountID string) *providerv0.PlannedRequest { + t.Helper() + request := &providerv0.PlannedRequest{ + RequestDescriptorId: "account.observe", + RequestDescriptorDigest: descriptorDigest(t, m, "account.observe"), + Method: providerv0.HTTPMethod_HTTP_METHOD_GET, + AdmittedOriginDigest: origin.GetAdmissionDigest(), + PathParameters: map[string]*providerv0.PublicValue{"account_id": pubString(accountID)}, + CredentialPurposes: []providerv0.CredentialPurpose{providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT}, + ResponsePolicyDigest: fakeDigest("account-response-policy"), + } + bound, err := canonical.BindPlannedRequestDigest(request) + require.NoError(t, err) + return bound +} + +func createAction(t *testing.T, requests ...*providerv0.PlannedRequest) *providerv0.PlanAction { + t.Helper() + action := &providerv0.PlanAction{ + ActionId: "a1", + Position: 0, + Type: providerv0.ActionType_ACTION_TYPE_CREATE, + ResourceType: "account", + ProspectiveRemoteId: remoteID, + Ownership: providerv0.Ownership_OWNERSHIP_OWNED, + Requests: requests, + } + require.NoError(t, canonical.ValidatePlanAction(action)) + return action +} + +func boundDeleteRequest(t *testing.T, m *manifest.Manifest, origin *providerv0.AdmittedOrigin) *providerv0.PlannedRequest { + t.Helper() + request := &providerv0.PlannedRequest{ + RequestDescriptorId: "account.delete", + RequestDescriptorDigest: descriptorDigest(t, m, "account.delete"), + Method: providerv0.HTTPMethod_HTTP_METHOD_DELETE, + AdmittedOriginDigest: origin.GetAdmissionDigest(), + PathParameters: map[string]*providerv0.PublicValue{"account_id": pubString(remoteID)}, + CredentialPurposes: []providerv0.CredentialPurpose{providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT}, + ResponsePolicyDigest: fakeDigest("deleted-response-policy"), + IdempotencyKey: "idem-del", + } + bound, err := canonical.BindPlannedRequestDigest(request) + require.NoError(t, err) + return bound +} + +func deleteAction(t *testing.T, requests ...*providerv0.PlannedRequest) *providerv0.PlanAction { + t.Helper() + action := &providerv0.PlanAction{ + ActionId: "a1", + Position: 0, + Type: providerv0.ActionType_ACTION_TYPE_DELETE, + ResourceType: "account", + RemoteIdentity: &providerv0.RemoteIdentity{Provider: "fixture", ResourceType: "account", RemoteId: remoteID}, + Ownership: providerv0.Ownership_OWNERSHIP_OWNED, + Requests: requests, + } + require.NoError(t, canonical.ValidatePlanAction(action)) + return action +} + +func binding() *providerv0.BindingAddress { + return &providerv0.BindingAddress{WorkspaceId: "ws", EnvironmentId: "env", BindingId: "bind"} +} + +func operation() *providerv0.OperationIdentity { + return &providerv0.OperationIdentity{OperationId: "op1", AttemptId: "att1", ActionId: "a1", PlanId: "plan1"} +} + +func budget() *providerv0.RequestBudget { + return &providerv0.RequestBudget{RequestCount: 4, RequestBytes: 8192, ResponseBytes: 65536} +} + +// fakeCheckpointer returns a preset durable checkpoint. +type fakeCheckpointer struct { + checkpoint *providerv0.ActionCheckpoint +} + +func (c *fakeCheckpointer) Latest(_ context.Context, _ *providerv0.OperationIdentity) (*providerv0.ActionCheckpoint, error) { + return c.checkpoint, nil +} + +func checkpoint(id, idempotencyKey string) *providerv0.ActionCheckpoint { + return &providerv0.ActionCheckpoint{ + CheckpointId: id, + Operation: operation(), + Delivery: providerv0.DeliveryState_DELIVERY_STATE_NOT_SENT, + IdempotencyKey: idempotencyKey, + } +} + +// dialClientFor returns a ClientFor that dials the given address regardless of +// the request URL, so tests reach a loopback server through the real transport +// path (tracing included) without depending on DNS or the manifest port. +func dialClientFor(addr string) func(urlguard.Origin, urlguard.Resolution) *http.Client { + return func(_ urlguard.Origin, _ urlguard.Resolution) *http.Client { + return &http.Client{ + Transport: &http.Transport{ + Proxy: nil, + DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) { + return (&net.Dialer{}).DialContext(ctx, network, addr) + }, + }, + CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, + } + } +} + +// harness wires a broker session with a mint credential handle and a durable +// checkpoint for the create request. +type harness struct { + origin originConfig + manifest *manifest.Manifest + action *providerv0.PlanAction + vault *credentials.Vault + sink *recordingSink + admitted *providerv0.AdmittedOrigin + create *providerv0.PlannedRequest + checkpoints *fakeCheckpointer +} + +func newHarness(t *testing.T) *harness { + t.Helper() + return newHarnessOn(t, loopbackConfig()) +} + +func newHarnessOn(t *testing.T, origin originConfig) *harness { + t.Helper() + m, err := manifest.Load([]byte(brokerManifestYAML(origin))) + require.NoError(t, err) + admitted := admittedOriginFor(t, origin) + create := boundCreateRequest(t, m, admitted) + action := createAction(t, create) + return &harness{ + origin: origin, + manifest: m, + action: action, + vault: credentials.NewVault(), + sink: &recordingSink{}, + admitted: admitted, + create: create, + checkpoints: &fakeCheckpointer{checkpoint: checkpoint("cp1", "idem-1")}, + } +} + +func (h *harness) mintHandle(t *testing.T, planned *providerv0.PlannedRequest, method providerv0.HTTPMethod) *providerv0.CredentialHandle { + t.Helper() + handle, err := h.vault.Mint(poisonSecret, credentials.Scope{ + Principal: "user", + Organization: "org", + ArtifactDigest: "sha256:aa", + Binding: binding(), + PlanID: "plan1", + ActionID: "a1", + RequestDigest: planned.GetRequestDigest(), + Purpose: providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, + Origin: h.origin.urlguardOrigin(), + Method: method, + Injection: credentials.Injection{Kind: credentials.InjectBearer}, + MaxUses: 1, + TTL: time.Minute, + }) + require.NoError(t, err) + return handle +} + +func (h *harness) executeRequest(handle *providerv0.CredentialHandle, planned *providerv0.PlannedRequest) *providerv0.ExecuteRequestRequest { + return &providerv0.ExecuteRequestRequest{ + Context: &providerv0.ProviderContext{ + Offline: &providerv0.OfflineProviderContext{Binding: binding()}, + Credentials: []*providerv0.CredentialHandle{handle}, + Operation: operation(), + Budget: budget(), + }, + RequestId: "req-1", + Request: planned, + Origin: h.admitted, + CredentialHandles: []*providerv0.CredentialHandle{handle}, + } +} + +func (h *harness) config(readOnly bool) broker.Config { + return broker.Config{ + Manifest: h.manifest, + Action: h.action, + Binding: binding(), + Budget: budget(), + ReadOnly: readOnly, + Vault: h.vault, + Sink: h.sink, + Checkpoints: h.checkpoints, + Deadlines: urlguard.DefaultDeadlines(), + } +} + +// recordingSink records every captured secret. +type recordingSink struct { + stored []string +} + +func (s *recordingSink) Put(_ context.Context, target responsepolicy.SinkTarget, secret string) (*providerv0.OpaqueReference, error) { + s.stored = append(s.stored, secret) + return &providerv0.OpaqueReference{ + Reference: "capture://" + target.Key, + Purpose: target.Purpose, + }, nil +} diff --git a/provider/broker/request.go b/provider/broker/request.go new file mode 100644 index 00000000..b9835b68 --- /dev/null +++ b/provider/broker/request.go @@ -0,0 +1,185 @@ +package broker + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "net/http" + "net/url" + "slices" + "strings" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" + "github.com/codefly-dev/core/provider/manifest" +) + +var methodNames = map[providerv0.HTTPMethod]string{ + providerv0.HTTPMethod_HTTP_METHOD_GET: http.MethodGet, + providerv0.HTTPMethod_HTTP_METHOD_HEAD: http.MethodHead, + providerv0.HTTPMethod_HTTP_METHOD_POST: http.MethodPost, + providerv0.HTTPMethod_HTTP_METHOD_PUT: http.MethodPut, + providerv0.HTTPMethod_HTTP_METHOD_PATCH: http.MethodPatch, + providerv0.HTTPMethod_HTTP_METHOD_DELETE: http.MethodDelete, +} + +// buildRequest constructs the exact outbound request. The provider supplies +// only descriptor-allowed structured fields; the host owns the URL, every +// security-relevant header, and the encoding. Path parameters and ownership +// body fields are bound to the action's planned remote id, so a valid template +// carrying another id fails before any credential is resolved. +func (s *Session) buildRequest( + ctx context.Context, + descriptor manifest.RequestDescriptor, + planned *providerv0.PlannedRequest, + origin urlguard.Origin, + remoteID string, +) (*http.Request, int64, error) { + method, ok := methodNames[planned.GetMethod()] + if !ok || method != descriptor.Method { + return nil, 0, fmt.Errorf("request method does not match descriptor") + } + + path, err := s.bindPath(descriptor, planned, remoteID) + if err != nil { + return nil, 0, err + } + query, err := bindQuery(descriptor, planned) + if err != nil { + return nil, 0, err + } + body, size, err := bindBody(descriptor, planned, remoteID) + if err != nil { + return nil, 0, err + } + + target := url.URL{Scheme: origin.Scheme, Host: origin.HostPort(), Path: path, RawQuery: query} + request, err := http.NewRequestWithContext(ctx, method, target.String(), bytes.NewReader(body)) + if err != nil { + return nil, 0, err + } + s.applyOwnedHeaders(request, planned, len(body) > 0) + return request, int64(size), nil +} + +func (s *Session) bindPath(descriptor manifest.RequestDescriptor, planned *providerv0.PlannedRequest, remoteID string) (string, error) { + template := descriptor.PathTemplate + params := planned.GetPathParameters() + if len(params) != len(descriptor.RemoteIDParameters) { + return "", fmt.Errorf("path parameters do not match descriptor") + } + for _, name := range descriptor.RemoteIDParameters { + value, ok := params[name] + if !ok { + return "", fmt.Errorf("path parameter %q is missing", name) + } + scalar, err := scalarString(value) + if err != nil { + return "", fmt.Errorf("path parameter %q: %w", name, err) + } + // Every remote-id path parameter must equal the action's planned id. + if scalar != remoteID { + return "", fmt.Errorf("path parameter %q is not the planned remote id", name) + } + template = strings.ReplaceAll(template, "{"+name+"}", url.PathEscape(scalar)) + } + if strings.ContainsAny(template, "{}") { + return "", fmt.Errorf("path template has unbound placeholders") + } + return urlguard.SafePath(template) +} + +func bindQuery(descriptor manifest.RequestDescriptor, planned *providerv0.PlannedRequest) (string, error) { + values := url.Values{} + for key, value := range planned.GetQuery() { + if !slices.Contains(descriptor.AllowedQueryFields, key) { + return "", fmt.Errorf("query field %q is not allowed", key) + } + scalar, err := scalarString(value) + if err != nil { + return "", fmt.Errorf("query field %q: %w", key, err) + } + values.Set(key, scalar) + } + return values.Encode(), nil +} + +func bindBody(descriptor manifest.RequestDescriptor, planned *providerv0.PlannedRequest, remoteID string) ([]byte, int, error) { + body := planned.GetBody() + // Every declared ownership field must be present and bound to the planned + // remote id. Checking only the fields the provider chose to include would let + // it omit an ownership field to escape identity binding entirely. + for _, field := range descriptor.OwnershipBodyFields { + value, ok := body[field] + if !ok { + return nil, 0, fmt.Errorf("ownership body field %q is missing", field) + } + scalar, err := scalarString(value) + if err != nil || scalar != remoteID { + return nil, 0, fmt.Errorf("ownership body field %q is not the planned remote id", field) + } + } + if len(body) == 0 { + return nil, 0, nil + } + object := make(map[string]any, len(body)) + for key, value := range body { + if !slices.Contains(descriptor.AllowedBodyFields, key) { + return nil, 0, fmt.Errorf("body field %q is not allowed", key) + } + converted, err := nativeValue(value) + if err != nil { + return nil, 0, fmt.Errorf("body field %q: %w", key, err) + } + object[key] = converted + } + encoded, err := json.Marshal(object) + if err != nil { + return nil, 0, err + } + return encoded, len(encoded), nil +} + +// applyOwnedHeaders installs the host-controlled header set. The provider +// contributes no headers at all: there is no caller-header allowlist to widen +// because the descriptor carries structured fields only. Authorization is left +// for the credential vault to inject. +func (s *Session) applyOwnedHeaders(request *http.Request, planned *providerv0.PlannedRequest, hasBody bool) { + request.Header = http.Header{} + request.Host = request.URL.Host + request.Header.Set("User-Agent", s.userAgent) + request.Header.Set("Accept", "application/json") + // The host owns decompression; it never lets a caller pick the encoding. + request.Header.Set("Accept-Encoding", "identity") + if hasBody { + request.Header.Set("Content-Type", "application/json") + } + if key := planned.GetIdempotencyKey(); key != "" && isMutating(planned.GetMethod()) { + request.Header.Set("Idempotency-Key", key) + } +} + +func isMutating(method providerv0.HTTPMethod) bool { + switch method { + case providerv0.HTTPMethod_HTTP_METHOD_GET, providerv0.HTTPMethod_HTTP_METHOD_HEAD: + return false + default: + return true + } +} + +// outboundSize estimates the full encoded request size — request line, Host, +// every host-owned header including the resolved credential, and the body — so +// the request budget bounds the whole outbound message rather than only its +// JSON body. +func outboundSize(request *http.Request, bodySize int64) int64 { + size := int64(len(request.Method) + len(" ") + len(request.URL.RequestURI()) + len(" HTTP/1.1\r\n")) + size += int64(len("Host: ") + len(request.Host) + len("\r\n")) + for name, values := range request.Header { + for _, value := range values { + size += int64(len(name) + len(": ") + len(value) + len("\r\n")) + } + } + return size + bodySize +} diff --git a/provider/broker/response.go b/provider/broker/response.go new file mode 100644 index 00000000..f85d515d --- /dev/null +++ b/provider/broker/response.go @@ -0,0 +1,100 @@ +package broker + +import ( + "fmt" + "io" + "net/http" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" +) + +// handleResponse turns a received response into a filtered safe response. The +// manifest response schema describes successful bodies only, so a non-success +// response is surfaced by status alone with its untrusted body dropped, and a +// response that carries no body by HTTP semantics is surfaced by status alone. +// A successful body is filtered through the schema and fails closed on any +// filtering error — but the delivery status is always preserved, never lost to +// a bare error. +func (s *Session) handleResponse(descriptor manifest.RequestDescriptor, resp *http.Response) (*providerv0.ExecuteRequestResponse, error) { + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + // A non-success body is untrusted and carries no schema-declared fields; + // drop it entirely so an attacker-controlled, secret-looking error body + // never reaches a provider, log, or persisted surface. + return &providerv0.ExecuteRequestResponse{ + StatusCode: uint32(resp.StatusCode), + Certainty: certaintyForStatus(resp.StatusCode), + }, nil + } + policy, err := s.responsePolicyFor(descriptor) + if err != nil { + return nil, err + } + // Read one byte past the budget: the policy rejects an over-budget body + // rather than silently truncating a successful response. HEAD and no-content + // responses read empty. + limit := s.limits.MaxCompressedBytes + raw, err := io.ReadAll(io.LimitReader(resp.Body, limit+1)) + if err != nil { + return nil, fmt.Errorf("read response body: %w", err) + } + if int64(len(raw)) > limit { + return nil, fmt.Errorf("response exceeds byte budget") + } + if len(raw) == 0 { + // A body-carrying success is expected whenever the schema requires a + // field; an empty body then is drift and fails closed. Otherwise a + // bodyless success (HEAD, 204, 205) is surfaced by status alone. + if policy.RequiresBody() { + return nil, fmt.Errorf("successful response has no body but the schema requires fields") + } + return &providerv0.ExecuteRequestResponse{ + StatusCode: uint32(resp.StatusCode), + Certainty: providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE, + }, nil + } + result, err := policy.Filter(resp.Request.Context(), raw, resp.Header.Get("Content-Encoding"), resp.Header.Get("Content-Type"), s.sink) + if err != nil { + return nil, fmt.Errorf("filter response: %w", err) + } + + response := &providerv0.ExecuteRequestResponse{ + StatusCode: uint32(resp.StatusCode), + Certainty: certaintyForStatus(resp.StatusCode), + SuppressedPresence: result.Suppressed, + } + for _, forwarded := range result.Forwarded { + response.Forwarded = append(response.Forwarded, &providerv0.FilteredField{ + Selector: forwarded.Selector, + Value: forwarded.Value, + }) + } + for _, capture := range result.Captures { + if capture.Outcome != responsepolicy.OutcomeCaptured { + continue + } + response.Captures = append(response.Captures, &providerv0.CaptureResult{ + CaptureId: capture.Selector, + Selector: capture.Selector, + SinkReference: capture.Reference, + Captured: true, + }) + } + return response, nil +} + +// certaintyForStatus maps a received status to the host's certainty about the +// remote effect. A 2xx is complete; a 4xx is a definitive client-side rejection, +// so the remote state is likewise known; a 3xx (surfaced because redirects are +// never followed) or 5xx leaves the remote effect genuinely unknown. +func certaintyForStatus(status int) providerv0.OutcomeCertainty { + switch { + case status >= 200 && status < 300: + return providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE + case status >= 400 && status < 500: + return providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE + default: + return providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_UNCERTAIN + } +} diff --git a/provider/broker/review_fixes_test.go b/provider/broker/review_fixes_test.go new file mode 100644 index 00000000..af7cfb33 --- /dev/null +++ b/provider/broker/review_fixes_test.go @@ -0,0 +1,256 @@ +package broker_test + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net" + "net/http" + "net/http/httptest" + "sync" + "testing" + "time" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/broker" + "github.com/codefly-dev/core/provider/canonical" + "github.com/codefly-dev/core/provider/cassette" + "github.com/stretchr/testify/require" + "google.golang.org/protobuf/types/known/timestamppb" +) + +// respondingServer replies to every request with a fixed status and body. +func respondingServer(t *testing.T, status int, contentType, body string) *httptest.Server { + t.Helper() + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + if contentType != "" { + w.Header().Set("Content-Type", contentType) + } + w.WriteHeader(status) + _, _ = w.Write([]byte(body)) + })) + t.Cleanup(server.Close) + return server +} + +// Finding 1: a create request that omits the declared ownership body field must +// be rejected — omission is a bypass of identity binding, not a valid request. +func TestExecute_OwnershipBodyFieldOmittedRejected(t *testing.T) { + m := mustLoad(t) + origin := admittedOrigin(t) + // A create request whose body omits the ownership field "name" entirely. + request := &providerv0.PlannedRequest{ + RequestDescriptorId: "account.create", + RequestDescriptorDigest: descriptorDigest(t, m, "account.create"), + Method: providerv0.HTTPMethod_HTTP_METHOD_POST, + AdmittedOriginDigest: origin.GetAdmissionDigest(), + Body: map[string]*providerv0.PublicValue{}, // no "name" + CredentialPurposes: []providerv0.CredentialPurpose{providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT}, + ResponsePolicyDigest: fakeDigest("account-response-policy"), + IdempotencyKey: "idem-1", + } + bound, err := canonical.BindPlannedRequestDigest(request) + require.NoError(t, err) + + h := newHarness(t) + h.action = createAction(t, bound) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + _, err = session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, bound, providerv0.HTTPMethod_HTTP_METHOD_POST), bound)) + require.ErrorContains(t, err, "ownership body field \"name\" is missing") + require.Empty(t, h.sink.stored) +} + +// Finding 2a: a non-2xx response surfaces its status with the untrusted error +// body dropped, rather than hard-erroring and losing the status. +func TestExecute_NonSuccessSurfacesStatusAndDropsBody(t *testing.T) { + h := newHarness(t) + body := fmt.Sprintf(`{"error":{"message":"bad","hint":%q}}`, poisonSecret) + server := respondingServer(t, http.StatusBadRequest, "application/json", body) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + response, err := session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.NoError(t, err) + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, response.GetDelivery()) + require.Equal(t, uint32(http.StatusBadRequest), response.GetStatusCode()) + require.Equal(t, providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE, response.GetCertainty()) + require.Empty(t, response.GetForwarded()) + require.Empty(t, response.GetCaptures()) + + encoded, err := json.Marshal(response) + require.NoError(t, err) + require.NotContains(t, string(encoded), poisonSecret) +} + +// Finding 2b: a bodyless successful response (204, no schema-required field) is +// surfaced by status alone instead of failing on an empty JSON parse. +func TestExecute_BodylessSuccessSurfacesStatus(t *testing.T) { + m := mustLoad(t) + origin := admittedOrigin(t) + del := boundDeleteRequest(t, m, origin) + + h := newHarness(t) + h.action = deleteAction(t, del) + h.checkpoints.checkpoint = checkpoint("cp1", "idem-del") + server := respondingServer(t, http.StatusNoContent, "", "") + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + response, err := session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, del, providerv0.HTTPMethod_HTTP_METHOD_DELETE), del)) + require.NoError(t, err) + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_RESPONSE_RECEIVED, response.GetDelivery()) + require.Equal(t, uint32(http.StatusNoContent), response.GetStatusCode()) + require.Empty(t, response.GetForwarded()) +} + +// Finding 2c: a bodyless response where the schema requires a field is drift and +// fails closed — a vendor cannot skip a required capture by returning 204. +func TestExecute_BodylessSuccessWithRequiredCaptureFailsClosed(t *testing.T) { + h := newHarness(t) + server := respondingServer(t, http.StatusNoContent, "", "") + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + session, err := broker.New(cfg) + require.NoError(t, err) + + _, err = session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.ErrorContains(t, err, "schema requires fields") + require.Empty(t, h.sink.stored) +} + +// Finding 3: the request budget deadline bounds the whole round trip, not just +// the pre-send checks. A response that never arrives before the deadline is +// reported as SENT_OUTCOME_UNKNOWN. +func TestExecute_BudgetDeadlineEnforcedOnRoundTrip(t *testing.T) { + release := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + <-release // never respond before the deadline + })) + // Cleanups run LIFO: unblock the handler first, then close the server. + t.Cleanup(server.Close) + t.Cleanup(func() { close(release) }) + + h := newHarness(t) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + cfg.Budget = &providerv0.RequestBudget{ + RequestCount: 4, + RequestBytes: 8192, + ResponseBytes: 65536, + Deadline: timestamppb.New(time.Now().Add(200 * time.Millisecond)), + } + session, err := broker.New(cfg) + require.NoError(t, err) + + response, err := session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.NoError(t, err) + require.Equal(t, providerv0.DeliveryState_DELIVERY_STATE_SENT_OUTCOME_UNKNOWN, response.GetDelivery()) +} + +// Finding 4: replay never resolves DNS. A cassette recorded for an unresolvable +// host replays successfully even with a resolver that fails every lookup, while +// a live session against the same host fails at resolution. +func TestExecute_ReplayDoesNotResolveDNS(t *testing.T) { + failing := &net.Resolver{ + PreferGo: true, + Dial: func(context.Context, string, string) (net.Conn, error) { + return nil, errors.New("dns disabled") + }, + } + h := newHarnessOn(t, originConfig{scheme: "https", host: "unresolvable.invalid", port: 443, class: "public"}) + + recorded := &providerv0.ExecuteRequestResponse{ + StatusCode: http.StatusCreated, + Certainty: providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE, + Forwarded: []*providerv0.FilteredField{{Selector: "$.id", Value: pubString(remoteID)}}, + } + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, recorder.Record(cassette.NewKey(h.create, h.admitted), http.StatusCreated, nil, recorded)) + data, err := recorder.Marshal() + require.NoError(t, err) + replay, err := cassette.Load(data, "1.2.3") + require.NoError(t, err) + + replayCfg := h.config(false) + replayCfg.Resolver = failing + replayCfg.Cassette = replay + replaySession, err := broker.New(replayCfg) + require.NoError(t, err) + + response, err := replaySession.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.NoError(t, err, "replay must not depend on DNS") + require.Equal(t, uint32(http.StatusCreated), response.GetStatusCode()) + + // A live session against the same unresolvable host fails at resolution, + // confirming resolution normally happens on the live path. + liveCfg := h.config(false) + liveCfg.Resolver = failing + liveSession, err := broker.New(liveCfg) + require.NoError(t, err) + _, err = liveSession.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.Error(t, err) +} + +// Finding 7: the request byte budget bounds the full outbound message, including +// host-owned headers and the injected credential, not just the JSON body. +func TestExecute_RequestByteBudgetCountsHeaders(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + // The create body alone is tiny; a 60-byte ceiling is only exceeded once the + // request line, Host, and headers are counted. + cfg.Budget = &providerv0.RequestBudget{RequestCount: 4, RequestBytes: 60, ResponseBytes: 65536} + session, err := broker.New(cfg) + require.NoError(t, err) + + _, err = session.Execute(context.Background(), h.executeRequest(h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST), h.create)) + require.ErrorContains(t, err, "exceeds byte budget") + require.Empty(t, h.sink.stored) +} + +// Finding 8: concurrent Execute calls on one session are serialized, so a +// single-callback budget admits exactly one request. Run under -race. +func TestExecute_ConcurrentCallsRespectBudget(t *testing.T) { + h := newHarness(t) + server := accountServer(t, nil) + cfg := h.config(false) + cfg.ClientFor = dialClientFor(serverAddr(t, server)) + cfg.Budget = &providerv0.RequestBudget{RequestCount: 1, RequestBytes: 8192, ResponseBytes: 65536} + session, err := broker.New(cfg) + require.NoError(t, err) + + const workers = 8 + handles := make([]*providerv0.CredentialHandle, workers) + for i := range handles { + handles[i] = h.mintHandle(t, h.create, providerv0.HTTPMethod_HTTP_METHOD_POST) + } + + var wg sync.WaitGroup + var mu sync.Mutex + successes := 0 + for i := range workers { + wg.Add(1) + go func(handle *providerv0.CredentialHandle) { + defer wg.Done() + _, err := session.Execute(context.Background(), h.executeRequest(handle, h.create)) + if err == nil { + mu.Lock() + successes++ + mu.Unlock() + } + }(handles[i]) + } + wg.Wait() + require.Equal(t, 1, successes, "a single-callback budget must admit exactly one request") +} diff --git a/provider/broker/value.go b/provider/broker/value.go new file mode 100644 index 00000000..e98f1c21 --- /dev/null +++ b/provider/broker/value.go @@ -0,0 +1,80 @@ +package broker + +import ( + "fmt" + "strconv" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/configuration" +) + +// scalarString renders a PublicValue as a path or query scalar. Only string, +// integer, decimal, and boolean values are admitted; a secret-shaped string +// fails closed so provider input cannot smuggle a credential into a URL. +func scalarString(value *providerv0.PublicValue) (string, error) { + switch kind := value.GetKind().(type) { + case *providerv0.PublicValue_StringValue: + if configuration.LooksSecret(kind.StringValue) { + return "", fmt.Errorf("value carries a secret-shaped literal") + } + return kind.StringValue, nil + case *providerv0.PublicValue_IntegerValue: + return strconv.FormatInt(kind.IntegerValue, 10), nil + case *providerv0.PublicValue_DecimalValue: + return kind.DecimalValue, nil + case *providerv0.PublicValue_BoolValue: + return strconv.FormatBool(kind.BoolValue), nil + default: + return "", fmt.Errorf("value is not a scalar") + } +} + +// nativeValue converts a PublicValue tree into a JSON-encodable Go value for +// request body serialization. Every string is checked so a secret cannot be +// placed in an outbound body by provider code. +func nativeValue(value *providerv0.PublicValue) (any, error) { + switch kind := value.GetKind().(type) { + case *providerv0.PublicValue_NullValue: + return nil, nil + case *providerv0.PublicValue_StringValue: + if configuration.LooksSecret(kind.StringValue) { + return nil, fmt.Errorf("value carries a secret-shaped literal") + } + return kind.StringValue, nil + case *providerv0.PublicValue_BoolValue: + return kind.BoolValue, nil + case *providerv0.PublicValue_IntegerValue: + return kind.IntegerValue, nil + case *providerv0.PublicValue_DecimalValue: + return jsonNumber(kind.DecimalValue), nil + case *providerv0.PublicValue_ListValue: + list := make([]any, 0, len(kind.ListValue.GetValues())) + for _, item := range kind.ListValue.GetValues() { + converted, err := nativeValue(item) + if err != nil { + return nil, err + } + list = append(list, converted) + } + return list, nil + case *providerv0.PublicValue_ObjectValue: + object := make(map[string]any, len(kind.ObjectValue.GetFields())) + for key, item := range kind.ObjectValue.GetFields() { + converted, err := nativeValue(item) + if err != nil { + return nil, err + } + object[key] = converted + } + return object, nil + default: + return nil, fmt.Errorf("value kind is unset") + } +} + +// jsonNumber marshals as a bare JSON number rather than a quoted string. +type jsonNumber string + +func (n jsonNumber) MarshalJSON() ([]byte, error) { + return []byte(n), nil +} diff --git a/provider/cassette/cassette.go b/provider/cassette/cassette.go new file mode 100644 index 00000000..b8eaf726 --- /dev/null +++ b/provider/cassette/cassette.go @@ -0,0 +1,235 @@ +// Package cassette records and replays broker responses deterministically. +// +// Recording is explicit opt-in; replay is the default and never falls back to +// the live network. Both paths run the identical broker admission, credential, +// and checkpoint checks — the cassette only substitutes for the network round +// trip, returning the already-filtered safe response that recording captured. +// A cassette never persists raw request or response bytes: it stores only the +// filtered safe response, keyed by a match digest that normalizes the +// host-injected idempotency key, request id, and timestamps without weakening +// the action, resource, origin, or response-policy binding. Recording fails +// closed if any secret-shaped value survives into what would be stored. +package cassette + +import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" + "fmt" + "sort" + "sync" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/configuration" + "google.golang.org/protobuf/encoding/protojson" + "google.golang.org/protobuf/proto" +) + +// Mode selects record or replay. The zero value is replay, so a misconfigured +// cassette never accidentally reaches the network. +type Mode int + +const ( + // ModeReplay returns recorded responses and never dials the network. + ModeReplay Mode = iota + // ModeRecord stores live filtered responses. + ModeRecord +) + +// Entry is one recorded exchange. It holds only safe, filtered material. +type Entry struct { + Sequence int `json:"sequence"` + Method string `json:"method"` + MatchDigest string `json:"match_digest"` + PolicyDigest string `json:"policy_digest"` + ProviderVersion string `json:"provider_version"` + Status uint32 `json:"status"` + SafeHeaders map[string]string `json:"safe_headers,omitempty"` + Response json.RawMessage `json:"response"` +} + +// Cassette is a deterministic record/replay store. +type Cassette struct { + mode Mode + providerVersion string + mu sync.Mutex + entries []*Entry + replayCursor map[string]int +} + +// New builds a cassette in the given mode bound to a concrete provider version. +func New(mode Mode, providerVersion string) *Cassette { + return &Cassette{mode: mode, providerVersion: providerVersion, replayCursor: map[string]int{}} +} + +// Mode reports the cassette mode. +func (c *Cassette) Mode() Mode { return c.mode } + +// Key identifies one logical request independent of host-injected volatile +// fields. Method, descriptor, path, query, body, origin, and response policy +// are bound; the idempotency key, request id, and timestamps are not. +type Key struct { + planned *providerv0.PlannedRequest + origin *providerv0.AdmittedOrigin +} + +// NewKey builds a match key from the admitted planned request and origin. +func NewKey(planned *providerv0.PlannedRequest, origin *providerv0.AdmittedOrigin) Key { + return Key{planned: planned, origin: origin} +} + +func (k Key) digests() (match string, policy string, method string, err error) { + if k.planned == nil || k.origin == nil { + return "", "", "", fmt.Errorf("cassette key requires a planned request and origin") + } + normalized := proto.Clone(k.planned).(*providerv0.PlannedRequest) + // Normalize host-injected volatile fields; keep action/resource/origin. The + // response policy is compared separately so a policy drift is a distinct, + // explicit mismatch rather than a silent non-match. + normalized.IdempotencyKey = "" + normalized.RequestDigest = "" + normalized.ResponsePolicyDigest = "" + origin := proto.Clone(k.origin).(*providerv0.AdmittedOrigin) + origin.AdmissionDigest = "" + + plannedBytes, err := proto.MarshalOptions{Deterministic: true}.Marshal(normalized) + if err != nil { + return "", "", "", err + } + originBytes, err := proto.MarshalOptions{Deterministic: true}.Marshal(origin) + if err != nil { + return "", "", "", err + } + sum := sha256.Sum256(append(plannedBytes, originBytes...)) + return "sha256:" + hex.EncodeToString(sum[:]), k.planned.GetResponsePolicyDigest(), k.planned.GetMethod().String(), nil +} + +// Record stores a live filtered response. It fails closed if any secret-shaped +// value survived filtering, so a cassette can never become an exfiltration +// surface. +func (c *Cassette) Record(key Key, status uint32, safeHeaders map[string]string, response *providerv0.ExecuteRequestResponse) error { + if c.mode != ModeRecord { + return fmt.Errorf("cassette is not in record mode") + } + match, policy, method, err := key.digests() + if err != nil { + return err + } + if err := assertSafe(response); err != nil { + return err + } + for name, value := range safeHeaders { + if configuration.LooksSecret(name) || configuration.LooksSecret(value) { + return fmt.Errorf("cassette header %q is secret-shaped", name) + } + } + encoded, err := marshalResponse(response) + if err != nil { + return err + } + if configuration.LooksSecret(string(encoded)) { + return fmt.Errorf("recorded response carries a secret-shaped literal") + } + c.mu.Lock() + defer c.mu.Unlock() + c.entries = append(c.entries, &Entry{ + Sequence: len(c.entries), + Method: method, + MatchDigest: match, + PolicyDigest: policy, + ProviderVersion: c.providerVersion, + Status: status, + SafeHeaders: safeHeaders, + Response: encoded, + }) + return nil +} + +// Replay returns the next recorded response for the key. It never falls back to +// the network: an unknown key or a response-policy mismatch is a hard error. +func (c *Cassette) Replay(key Key) (uint32, *providerv0.ExecuteRequestResponse, error) { + if c.mode != ModeReplay { + return 0, nil, fmt.Errorf("cassette is not in replay mode") + } + match, policy, _, err := key.digests() + if err != nil { + return 0, nil, err + } + c.mu.Lock() + defer c.mu.Unlock() + cursor := c.replayCursor[match] + var found *Entry + seen := 0 + for _, entry := range c.entries { + if entry.MatchDigest != match { + continue + } + if seen == cursor { + found = entry + break + } + seen++ + } + if found == nil { + return 0, nil, fmt.Errorf("no recorded response for request; replay does not fall back to live") + } + if found.PolicyDigest != policy { + return 0, nil, fmt.Errorf("recorded response policy digest does not match current policy") + } + if found.ProviderVersion != c.providerVersion { + return 0, nil, fmt.Errorf("recorded provider version does not match current provider") + } + c.replayCursor[match] = cursor + 1 + response, err := unmarshalResponse(found.Response) + if err != nil { + return 0, nil, err + } + return found.Status, response, nil +} + +// Marshal serializes the cassette deterministically for a byte-stable on-disk +// form and reviewable diffs. +func (c *Cassette) Marshal() ([]byte, error) { + c.mu.Lock() + defer c.mu.Unlock() + entries := append([]*Entry(nil), c.entries...) + sort.SliceStable(entries, func(i, j int) bool { return entries[i].Sequence < entries[j].Sequence }) + return json.MarshalIndent(entries, "", " ") +} + +// Load restores recorded entries for replay. +func Load(data []byte, providerVersion string) (*Cassette, error) { + var entries []*Entry + if err := json.Unmarshal(data, &entries); err != nil { + return nil, fmt.Errorf("load cassette: %w", err) + } + return &Cassette{mode: ModeReplay, providerVersion: providerVersion, entries: entries, replayCursor: map[string]int{}}, nil +} + +func marshalResponse(response *providerv0.ExecuteRequestResponse) (json.RawMessage, error) { + clone := proto.Clone(response).(*providerv0.ExecuteRequestResponse) + // Volatile host-injected fields are excluded from the stored form so a + // re-record is byte-identical. + clone.RequestId = "" + clone.ResponseReceivedAt = nil + encoded, err := protojson.Marshal(clone) + if err != nil { + return nil, err + } + // protojson deliberately varies whitespace and map-entry order; normalize + // through a generic decode/re-encode (which sorts object keys) so the stored + // form is both deterministic and human-reviewable in cassette diffs. + var generic any + if err := json.Unmarshal(encoded, &generic); err != nil { + return nil, err + } + return json.Marshal(generic) +} + +func unmarshalResponse(raw json.RawMessage) (*providerv0.ExecuteRequestResponse, error) { + response := &providerv0.ExecuteRequestResponse{} + if err := protojson.Unmarshal(raw, response); err != nil { + return nil, err + } + return response, nil +} diff --git a/provider/cassette/cassette_test.go b/provider/cassette/cassette_test.go new file mode 100644 index 00000000..4fe7c2c1 --- /dev/null +++ b/provider/cassette/cassette_test.go @@ -0,0 +1,160 @@ +package cassette_test + +import ( + "testing" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/cassette" + "github.com/stretchr/testify/require" +) + +func plannedRequest(policyDigest string) *providerv0.PlannedRequest { + return &providerv0.PlannedRequest{ + RequestDescriptorId: "account.create", + Method: providerv0.HTTPMethod_HTTP_METHOD_POST, + ResponsePolicyDigest: policyDigest, + IdempotencyKey: "idem-volatile", + } +} + +func admittedOrigin() *providerv0.AdmittedOrigin { + return &providerv0.AdmittedOrigin{OriginRuleId: "api", Scheme: "https", Host: "api.stripe.com", Port: 443} +} + +func safeResponse() *providerv0.ExecuteRequestResponse { + return &providerv0.ExecuteRequestResponse{ + StatusCode: 200, + Certainty: providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE, + Forwarded: []*providerv0.FilteredField{ + {Selector: "$.id", Value: &providerv0.PublicValue{Kind: &providerv0.PublicValue_StringValue{StringValue: "acct_1"}}}, + }, + } +} + +func TestRecordReplay_RoundTrips(t *testing.T) { + policy := "sha256:" + repeat64('a') + key := cassette.NewKey(plannedRequest(policy), admittedOrigin()) + + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, recorder.Record(key, 200, nil, safeResponse())) + + data, err := recorder.Marshal() + require.NoError(t, err) + + replayer, err := cassette.Load(data, "1.2.3") + require.NoError(t, err) + + // The idempotency key differs on replay (host-injected volatile), but the + // match still succeeds because it is normalized out of the key. + status, response, err := replayer.Replay(mutateIdempotency(t, policy)) + require.NoError(t, err) + require.Equal(t, uint32(200), status) + require.Len(t, response.GetForwarded(), 1) + require.Equal(t, "acct_1", response.GetForwarded()[0].GetValue().GetStringValue()) +} + +func mutateIdempotency(t *testing.T, policy string) cassette.Key { + t.Helper() + planned := plannedRequest(policy) + planned.IdempotencyKey = "idem-different" + return cassette.NewKey(planned, admittedOrigin()) +} + +func TestRecording_IsHumanReviewable(t *testing.T) { + // The stored cassette must be a readable diff, not an opaque blob: the + // selector and safe value appear as plain text. + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, recorder.Record(cassette.NewKey(plannedRequest("sha256:"+repeat64('a')), admittedOrigin()), 200, nil, safeResponse())) + data, err := recorder.Marshal() + require.NoError(t, err) + require.Contains(t, string(data), "$.id") + require.Contains(t, string(data), "acct_1") +} + +func TestRecording_IsDeterministic(t *testing.T) { + policy := "sha256:" + repeat64('b') + key := cassette.NewKey(plannedRequest(policy), admittedOrigin()) + + first := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, first.Record(key, 200, nil, safeResponse())) + firstBytes, err := first.Marshal() + require.NoError(t, err) + + second := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, second.Record(key, 200, nil, safeResponse())) + secondBytes, err := second.Marshal() + require.NoError(t, err) + + require.Equal(t, firstBytes, secondBytes, "recording must be byte-identical") +} + +func TestReplay_NoLiveFallback(t *testing.T) { + replayer := cassette.New(cassette.ModeReplay, "1.2.3") + key := cassette.NewKey(plannedRequest("sha256:"+repeat64('c')), admittedOrigin()) + _, _, err := replayer.Replay(key) + require.ErrorContains(t, err, "does not fall back to live") +} + +func TestReplay_PolicyMismatchRejected(t *testing.T) { + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, recorder.Record(cassette.NewKey(plannedRequest("sha256:"+repeat64('d')), admittedOrigin()), 200, nil, safeResponse())) + data, err := recorder.Marshal() + require.NoError(t, err) + + replayer, err := cassette.Load(data, "1.2.3") + require.NoError(t, err) + + // Same request but a drifted response policy digest. + _, _, err = replayer.Replay(cassette.NewKey(plannedRequest("sha256:"+repeat64('e')), admittedOrigin())) + require.ErrorContains(t, err, "policy digest") +} + +func TestReplay_ProviderVersionMismatchRejected(t *testing.T) { + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + require.NoError(t, recorder.Record(cassette.NewKey(plannedRequest("sha256:"+repeat64('f')), admittedOrigin()), 200, nil, safeResponse())) + data, err := recorder.Marshal() + require.NoError(t, err) + + replayer, err := cassette.Load(data, "9.9.9") // different provider version + require.NoError(t, err) + _, _, err = replayer.Replay(cassette.NewKey(plannedRequest("sha256:"+repeat64('f')), admittedOrigin())) + require.ErrorContains(t, err, "provider version") +} + +func TestRecord_PoisonSecretHardFailure(t *testing.T) { + poison := &providerv0.ExecuteRequestResponse{ + StatusCode: 200, + Certainty: providerv0.OutcomeCertainty_OUTCOME_CERTAINTY_COMPLETE, + Forwarded: []*providerv0.FilteredField{ + {Selector: "$.token", Value: &providerv0.PublicValue{Kind: &providerv0.PublicValue_StringValue{StringValue: "sk_live_1234567890abcdef"}}}, + }, + } + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + err := recorder.Record(cassette.NewKey(plannedRequest("sha256:"+repeat64('a')), admittedOrigin()), 200, nil, poison) + require.Error(t, err) +} + +func TestRecord_RejectsSecretHeader(t *testing.T) { + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + err := recorder.Record( + cassette.NewKey(plannedRequest("sha256:"+repeat64('a')), admittedOrigin()), + 200, + map[string]string{"X-Trace": "Bearer sk_live_1234567890abcdef"}, + safeResponse(), + ) + require.Error(t, err) +} + +func TestReplay_RejectsRecordModeCassette(t *testing.T) { + recorder := cassette.New(cassette.ModeRecord, "1.2.3") + _, _, err := recorder.Replay(cassette.NewKey(plannedRequest("sha256:"+repeat64('a')), admittedOrigin())) + require.ErrorContains(t, err, "not in replay mode") +} + +func repeat64(b byte) string { + out := make([]byte, 64) + for i := range out { + out[i] = b + } + return string(out) +} diff --git a/provider/cassette/safety.go b/provider/cassette/safety.go new file mode 100644 index 00000000..2614a8a0 --- /dev/null +++ b/provider/cassette/safety.go @@ -0,0 +1,83 @@ +package cassette + +import ( + "fmt" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/configuration" + "google.golang.org/protobuf/reflect/protoreflect" +) + +// assertSafe walks a filtered response and rejects any secret-shaped string. A +// forwarded value, selector, suppressed presence marker, or capture reference +// that still looks like a secret means filtering drifted, and the response must +// not be recorded. +func assertSafe(response *providerv0.ExecuteRequestResponse) error { + if response == nil { + return fmt.Errorf("cassette response is required") + } + return scanMessage(response.ProtoReflect()) +} + +func scanMessage(message protoreflect.Message) error { + var found error + message.Range(func(field protoreflect.FieldDescriptor, value protoreflect.Value) bool { + found = scanValue(field, value) + return found == nil + }) + return found +} + +func scanValue(field protoreflect.FieldDescriptor, value protoreflect.Value) error { + switch { + case field.IsMap(): + return scanMap(field, value) + case field.IsList(): + return scanList(field, value) + default: + return scanScalar(field, value) + } +} + +func scanMap(field protoreflect.FieldDescriptor, value protoreflect.Value) error { + var found error + value.Map().Range(func(key protoreflect.MapKey, entry protoreflect.Value) bool { + if configuration.LooksSecret(key.String()) { + found = fmt.Errorf("map key is secret-shaped") + return false + } + if field.MapValue().Kind() == protoreflect.MessageKind { + found = scanMessage(entry.Message()) + } else if field.MapValue().Kind() == protoreflect.StringKind && configuration.LooksSecret(entry.String()) { + found = fmt.Errorf("map value is secret-shaped") + } + return found == nil + }) + return found +} + +func scanList(field protoreflect.FieldDescriptor, value protoreflect.Value) error { + list := value.List() + for i := 0; i < list.Len(); i++ { + if field.Kind() == protoreflect.MessageKind { + if err := scanMessage(list.Get(i).Message()); err != nil { + return err + } + } else if field.Kind() == protoreflect.StringKind && configuration.LooksSecret(list.Get(i).String()) { + return fmt.Errorf("repeated value is secret-shaped") + } + } + return nil +} + +func scanScalar(field protoreflect.FieldDescriptor, value protoreflect.Value) error { + switch field.Kind() { + case protoreflect.MessageKind: + return scanMessage(value.Message()) + case protoreflect.StringKind: + if configuration.LooksSecret(value.String()) { + return fmt.Errorf("field %q is secret-shaped", field.Name()) + } + } + return nil +} diff --git a/provider/credentials/credentials.go b/provider/credentials/credentials.go new file mode 100644 index 00000000..d43f97ca --- /dev/null +++ b/provider/credentials/credentials.go @@ -0,0 +1,238 @@ +// Package credentials mints and resolves opaque credential handles as +// attenuated capabilities. A provider never receives a raw secret: it receives +// a handle bound to an exact principal, binding, action, request, origin, +// purpose, and injection location. The host resolves the handle only after +// every request check has already passed, injects the value without logging it, +// never returns it, and rejects any cross-purpose, cross-binding, or +// cross-action reuse. A handle is not permission; it is one narrow capability +// the host still re-checks at use. +package credentials + +import ( + "crypto/rand" + "encoding/hex" + "fmt" + "net/http" + "sync" + "time" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" +) + +// InjectionKind is the closed set of host-owned injection locations. +type InjectionKind string + +const ( + // InjectBearer sets Authorization: Bearer . + InjectBearer InjectionKind = "bearer" + // InjectHeader sets a named API-key header to . + InjectHeader InjectionKind = "header" + // InjectQuery sets a named query parameter to . + InjectQuery InjectionKind = "query" +) + +// Injection is where the host places a resolved secret. The provider never +// selects this; it is fixed at mint time from the plan. +type Injection struct { + Kind InjectionKind + Name string +} + +func (i Injection) validate() error { + switch i.Kind { + case InjectBearer: + if i.Name != "" { + return fmt.Errorf("bearer injection takes no name") + } + case InjectHeader, InjectQuery: + if i.Name == "" { + return fmt.Errorf("%s injection requires a name", i.Kind) + } + default: + return fmt.Errorf("unknown injection kind %q", i.Kind) + } + return nil +} + +// Scope is the exact capability a handle grants. Every field is host-derived +// and immutable after minting. +type Scope struct { + Principal string + Organization string + ArtifactDigest string + Binding *providerv0.BindingAddress + PlanID string + ActionID string + RequestDigest string + Purpose providerv0.CredentialPurpose + Origin urlguard.Origin + Method providerv0.HTTPMethod + Injection Injection + MaxUses uint32 + TTL time.Duration +} + +func (s Scope) validate() error { + switch { + case s.Principal == "" || s.Organization == "" || s.ArtifactDigest == "": + return fmt.Errorf("scope requires principal, organization, and artifact digest") + case s.Binding == nil || s.Binding.GetBindingId() == "": + return fmt.Errorf("scope requires a binding") + case s.PlanID == "" || s.ActionID == "" || s.RequestDigest == "": + return fmt.Errorf("scope requires plan, action, and request identity") + case s.Origin.Host == "": + return fmt.Errorf("scope requires an admitted origin") + case s.MaxUses == 0: + return fmt.Errorf("scope requires at least one use") + case s.TTL <= 0: + return fmt.Errorf("scope requires a positive TTL") + } + switch s.Purpose { + case providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, + providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_RUNTIME, + providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_BUILD, + providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_WEBHOOK_VERIFICATION: + default: + return fmt.Errorf("scope has an unknown credential purpose") + } + return s.Injection.validate() +} + +// Use is the exact request context presented at resolution. It must match the +// minted scope in every dimension or resolution fails. +type Use struct { + Purpose providerv0.CredentialPurpose + Binding *providerv0.BindingAddress + ActionID string + RequestDigest string + Origin urlguard.Origin + Method providerv0.HTTPMethod +} + +type entry struct { + secret string + scope Scope + expiresAt time.Time + remaining uint32 +} + +// Vault mints and resolves handles. It holds raw secrets and never exposes +// them; only opaque handle IDs and safe correlation are observable. +type Vault struct { + now func() time.Time + mu sync.Mutex + entries map[string]*entry +} + +// NewVault builds an empty vault using the wall clock. +func NewVault() *Vault { + return &Vault{now: time.Now, entries: make(map[string]*entry)} +} + +// WithClock overrides the clock for deterministic tests. +func (v *Vault) WithClock(now func() time.Time) *Vault { + v.now = now + return v +} + +// Mint stores secret under a fresh opaque handle bound to scope. The returned +// CredentialHandle carries only the opaque id and the purpose. +func (v *Vault) Mint(secret string, scope Scope) (*providerv0.CredentialHandle, error) { + if secret == "" { + return nil, fmt.Errorf("credential secret is required") + } + if err := scope.validate(); err != nil { + return nil, err + } + id, err := newHandleID() + if err != nil { + return nil, err + } + v.mu.Lock() + defer v.mu.Unlock() + v.entries[id] = &entry{ + secret: secret, + scope: scope, + expiresAt: v.now().Add(scope.TTL), + remaining: scope.MaxUses, + } + return &providerv0.CredentialHandle{Handle: id, Purpose: scope.Purpose}, nil +} + +// Inject resolves the handle against the exact use, decrements its remaining +// uses, and writes the secret into the request. It returns a scope error +// without ever echoing the secret. The request is mutated only on success. +func (v *Vault) Inject(request *http.Request, handle string, use Use) error { + v.mu.Lock() + defer v.mu.Unlock() + stored, ok := v.entries[handle] + if !ok { + return fmt.Errorf("credential handle is unknown") + } + if v.now().After(stored.expiresAt) { + delete(v.entries, handle) + return fmt.Errorf("credential handle has expired") + } + if stored.remaining == 0 { + return fmt.Errorf("credential handle has no remaining uses") + } + if err := matchScope(stored.scope, use); err != nil { + return err + } + apply(request, stored.scope.Injection, stored.secret) + stored.remaining-- + if stored.remaining == 0 { + delete(v.entries, handle) + } + return nil +} + +func matchScope(scope Scope, use Use) error { + if scope.Purpose != use.Purpose { + return fmt.Errorf("credential purpose does not match handle") + } + if !sameBinding(scope.Binding, use.Binding) { + return fmt.Errorf("credential binding does not match handle") + } + if scope.ActionID != use.ActionID { + return fmt.Errorf("credential action does not match handle") + } + if scope.RequestDigest != use.RequestDigest { + return fmt.Errorf("credential request does not match handle") + } + if scope.Origin != use.Origin { + return fmt.Errorf("credential origin does not match handle") + } + if scope.Method != use.Method { + return fmt.Errorf("credential method does not match handle") + } + return nil +} + +func sameBinding(a, b *providerv0.BindingAddress) bool { + return a.GetWorkspaceId() == b.GetWorkspaceId() && + a.GetEnvironmentId() == b.GetEnvironmentId() && + a.GetBindingId() == b.GetBindingId() +} + +func apply(request *http.Request, injection Injection, secret string) { + switch injection.Kind { + case InjectBearer: + request.Header.Set("Authorization", "Bearer "+secret) + case InjectHeader: + request.Header.Set(injection.Name, secret) + case InjectQuery: + query := request.URL.Query() + query.Set(injection.Name, secret) + request.URL.RawQuery = query.Encode() + } +} + +func newHandleID() (string, error) { + buffer := make([]byte, 16) + if _, err := rand.Read(buffer); err != nil { + return "", fmt.Errorf("mint credential handle: %w", err) + } + return "cfh_" + hex.EncodeToString(buffer), nil +} diff --git a/provider/credentials/credentials_test.go b/provider/credentials/credentials_test.go new file mode 100644 index 00000000..496c647d --- /dev/null +++ b/provider/credentials/credentials_test.go @@ -0,0 +1,165 @@ +package credentials_test + +import ( + "net/http" + "testing" + "time" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/network/urlguard" + "github.com/codefly-dev/core/provider/credentials" + "github.com/stretchr/testify/require" +) + +const secret = "sk_live_1234567890abcdef" + +func binding() *providerv0.BindingAddress { + return &providerv0.BindingAddress{WorkspaceId: "ws", EnvironmentId: "env", BindingId: "bind"} +} + +func origin() urlguard.Origin { + return urlguard.Origin{Scheme: "https", Host: "api.stripe.com", Port: 443} +} + +func baseScope() credentials.Scope { + return credentials.Scope{ + Principal: "user", + Organization: "org", + ArtifactDigest: "sha256:aa", + Binding: binding(), + PlanID: "plan-1", + ActionID: "action-1", + RequestDigest: "sha256:req", + Purpose: providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, + Origin: origin(), + Method: providerv0.HTTPMethod_HTTP_METHOD_POST, + Injection: credentials.Injection{Kind: credentials.InjectBearer}, + MaxUses: 1, + TTL: time.Minute, + } +} + +func baseUse() credentials.Use { + return credentials.Use{ + Purpose: providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, + Binding: binding(), + ActionID: "action-1", + RequestDigest: "sha256:req", + Origin: origin(), + Method: providerv0.HTTPMethod_HTTP_METHOD_POST, + } +} + +func newRequest(t *testing.T) *http.Request { + t.Helper() + request, err := http.NewRequest(http.MethodPost, "https://api.stripe.com/v1/x", nil) + require.NoError(t, err) + return request +} + +func TestMint_ReturnsOpaqueHandleNoSecret(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + require.NotContains(t, handle.GetHandle(), secret) + require.Equal(t, providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, handle.GetPurpose()) +} + +func TestInject_SucceedsThenExhaustsUses(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + + request := newRequest(t) + require.NoError(t, vault.Inject(request, handle.GetHandle(), baseUse())) + require.Equal(t, "Bearer "+secret, request.Header.Get("Authorization")) + + // Single-use handle is now spent. + err = vault.Inject(newRequest(t), handle.GetHandle(), baseUse()) + require.Error(t, err) +} + +func TestInject_CrossPurposeRejected(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + + use := baseUse() + use.Purpose = providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_RUNTIME + request := newRequest(t) + require.ErrorContains(t, vault.Inject(request, handle.GetHandle(), use), "purpose") + require.Empty(t, request.Header.Get("Authorization")) +} + +func TestInject_CrossBindingRejected(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + + use := baseUse() + use.Binding = &providerv0.BindingAddress{WorkspaceId: "ws", EnvironmentId: "env", BindingId: "other"} + require.ErrorContains(t, vault.Inject(newRequest(t), handle.GetHandle(), use), "binding") +} + +func TestInject_CrossActionRejected(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + + use := baseUse() + use.ActionID = "action-2" + require.ErrorContains(t, vault.Inject(newRequest(t), handle.GetHandle(), use), "action") +} + +func TestInject_WrongOriginOrMethodRejected(t *testing.T) { + vault := credentials.NewVault() + handle, err := vault.Mint(secret, baseScope()) + require.NoError(t, err) + + wrongOrigin := baseUse() + wrongOrigin.Origin = urlguard.Origin{Scheme: "https", Host: "evil.com", Port: 443} + require.ErrorContains(t, vault.Inject(newRequest(t), handle.GetHandle(), wrongOrigin), "origin") + + wrongMethod := baseUse() + wrongMethod.Method = providerv0.HTTPMethod_HTTP_METHOD_DELETE + require.ErrorContains(t, vault.Inject(newRequest(t), handle.GetHandle(), wrongMethod), "method") +} + +func TestInject_ExpiredRejected(t *testing.T) { + current := time.Unix(1000, 0) + vault := credentials.NewVault().WithClock(func() time.Time { return current }) + scope := baseScope() + scope.TTL = 30 * time.Second + handle, err := vault.Mint(secret, scope) + require.NoError(t, err) + + current = current.Add(time.Minute) + require.ErrorContains(t, vault.Inject(newRequest(t), handle.GetHandle(), baseUse()), "expired") +} + +func TestInject_HeaderAndQueryInjection(t *testing.T) { + vault := credentials.NewVault() + headerScope := baseScope() + headerScope.Injection = credentials.Injection{Kind: credentials.InjectHeader, Name: "X-Api-Key"} + headerHandle, err := vault.Mint(secret, headerScope) + require.NoError(t, err) + request := newRequest(t) + require.NoError(t, vault.Inject(request, headerHandle.GetHandle(), baseUse())) + require.Equal(t, secret, request.Header.Get("X-Api-Key")) + + queryScope := baseScope() + queryScope.Injection = credentials.Injection{Kind: credentials.InjectQuery, Name: "token"} + queryHandle, err := vault.Mint(secret, queryScope) + require.NoError(t, err) + request = newRequest(t) + require.NoError(t, vault.Inject(request, queryHandle.GetHandle(), baseUse())) + require.Equal(t, secret, request.URL.Query().Get("token")) +} + +func TestMint_RejectsIncompleteScope(t *testing.T) { + vault := credentials.NewVault() + scope := baseScope() + scope.ActionID = "" + _, err := vault.Mint(secret, scope) + require.Error(t, err) +} diff --git a/provider/responsepolicy/fuzz_test.go b/provider/responsepolicy/fuzz_test.go new file mode 100644 index 00000000..9dddd78c --- /dev/null +++ b/provider/responsepolicy/fuzz_test.go @@ -0,0 +1,52 @@ +package responsepolicy_test + +import ( + "context" + "testing" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" +) + +// FuzzFilter proves that response filtering never panics on arbitrary bytes, +// never forwards a secret-shaped value, and only ever stores captured secrets +// in the sink. Duplicate keys, bombs, and malformed JSON must fail closed, not +// crash. +func FuzzFilter(f *testing.F) { + for _, seed := range []string{ + `{"id":"a","secret":"whsec_x"}`, + `{"id":"a","id":"b"}`, + `{"data":[{"secret":"sk_live_1234567890abcdef"}]}`, + `not json`, + `{"n":1e999}`, + `[]`, + ``, + } { + f.Add([]byte(seed)) + } + pol := responsepolicy.Policy{ + Fields: []responsepolicy.Field{ + {Selector: manifest.Selector{Version: "v1", Path: "$.id"}, Disposition: manifest.ResponseForwardSafe}, + {Selector: manifest.Selector{Version: "v1", Path: "$.data[*].secret"}, Disposition: manifest.ResponseCaptureToSink, Purpose: 2, SinkKey: "k"}, + }, + Limits: responsepolicy.DefaultLimits(), + } + f.Fuzz(func(t *testing.T, body []byte) { + result, err := pol.Filter(context.Background(), body, "", "application/json", &fuzzSink{}) + if err != nil { + return + } + for _, forwarded := range result.Forwarded { + if forwarded.Value == nil { + t.Fatalf("forwarded field %q has nil value", forwarded.Selector) + } + } + }) +} + +type fuzzSink struct{} + +func (fuzzSink) Put(_ context.Context, target responsepolicy.SinkTarget, _ string) (*providerv0.OpaqueReference, error) { + return &providerv0.OpaqueReference{Reference: "capture://" + target.Key, Purpose: target.Purpose}, nil +} diff --git a/provider/responsepolicy/json.go b/provider/responsepolicy/json.go new file mode 100644 index 00000000..824a4cce --- /dev/null +++ b/provider/responsepolicy/json.go @@ -0,0 +1,248 @@ +package responsepolicy + +import ( + "bytes" + "compress/gzip" + "encoding/json" + "fmt" + "io" + "sort" + "strings" +) + +// Limits bounds every dimension an untrusted vendor response can grow along. +// Zero fields are rejected by Validate so a caller cannot accidentally admit an +// unbounded response. +type Limits struct { + MaxCompressedBytes int64 + MaxDecompressedBytes int64 + MaxDepth int + MaxKeys int + MaxArrayLen int + MaxStringBytes int +} + +// DefaultLimits are conservative bounds for JSON management APIs. +func DefaultLimits() Limits { + return Limits{ + MaxCompressedBytes: 1 << 20, + MaxDecompressedBytes: 8 << 20, + MaxDepth: 64, + MaxKeys: 10000, + MaxArrayLen: 10000, + MaxStringBytes: 1 << 20, + } +} + +func (l Limits) validate() error { + if l.MaxCompressedBytes <= 0 || l.MaxDecompressedBytes <= 0 || l.MaxDepth <= 0 || + l.MaxKeys <= 0 || l.MaxArrayLen <= 0 || l.MaxStringBytes <= 0 { + return fmt.Errorf("response limits must all be positive") + } + return nil +} + +// decodeBody bounds and, if needed, decompresses a response body, then parses +// it into a normalized value tree. It owns Accept-Encoding decoding entirely: +// only identity and gzip are accepted, both bounded against decompression +// bombs. +func decodeBody(raw []byte, contentEncoding string, limits Limits) (any, error) { + if err := limits.validate(); err != nil { + return nil, err + } + if int64(len(raw)) > limits.MaxCompressedBytes { + return nil, fmt.Errorf("response exceeds compressed byte budget") + } + var reader io.Reader = bytes.NewReader(raw) + switch strings.ToLower(strings.TrimSpace(contentEncoding)) { + case "", "identity": + case "gzip": + zr, err := gzip.NewReader(bytes.NewReader(raw)) + if err != nil { + return nil, fmt.Errorf("response gzip is invalid: %w", err) + } + reader = zr + default: + return nil, fmt.Errorf("response content encoding %q is not admitted", contentEncoding) + } + // Read one byte past the budget so an exact-size bomb is still detected. + bounded := io.LimitReader(reader, limits.MaxDecompressedBytes+1) + decoded, err := io.ReadAll(bounded) + if err != nil { + return nil, fmt.Errorf("read response body: %w", err) + } + if int64(len(decoded)) > limits.MaxDecompressedBytes { + return nil, fmt.Errorf("response exceeds decompressed byte budget") + } + return parseJSON(decoded, limits) +} + +// parseJSON decodes JSON while rejecting duplicate object keys and enforcing +// structural limits. It uses the streaming token API so a duplicate key is +// caught before any value is materialized — encoding/json's default decoder +// silently keeps the last of duplicate keys, which is an ambiguity an attacker +// can exploit. +func parseJSON(data []byte, limits Limits) (any, error) { + decoder := json.NewDecoder(bytes.NewReader(data)) + decoder.UseNumber() + counter := &keyCounter{max: limits.MaxKeys} + value, err := parseValue(decoder, limits, 1, counter) + if err != nil { + return nil, err + } + if decoder.More() { + return nil, fmt.Errorf("response contains trailing data") + } + return value, nil +} + +type keyCounter struct { + seen int + max int +} + +func (c *keyCounter) add() error { + c.seen++ + if c.seen > c.max { + return fmt.Errorf("response exceeds object key budget") + } + return nil +} + +func parseValue(decoder *json.Decoder, limits Limits, depth int, counter *keyCounter) (any, error) { + if depth > limits.MaxDepth { + return nil, fmt.Errorf("response exceeds nesting budget") + } + token, err := decoder.Token() + if err != nil { + return nil, fmt.Errorf("decode response: %w", err) + } + return parseFromToken(token, decoder, limits, depth, counter) +} + +func parseFromToken(token json.Token, decoder *json.Decoder, limits Limits, depth int, counter *keyCounter) (any, error) { + switch t := token.(type) { + case json.Delim: + switch t { + case '{': + return parseObject(decoder, limits, depth, counter) + case '[': + return parseArray(decoder, limits, depth, counter) + default: + return nil, fmt.Errorf("response contains unbalanced delimiter") + } + case string: + if len(t) > limits.MaxStringBytes { + return nil, fmt.Errorf("response exceeds string byte budget") + } + return t, nil + case json.Number, bool, nil: + return t, nil + default: + return nil, fmt.Errorf("response contains an unsupported token") + } +} + +func parseObject(decoder *json.Decoder, limits Limits, depth int, counter *keyCounter) (any, error) { + object := make(map[string]any) + for decoder.More() { + keyToken, err := decoder.Token() + if err != nil { + return nil, fmt.Errorf("decode response object key: %w", err) + } + key, ok := keyToken.(string) + if !ok { + return nil, fmt.Errorf("response object key is not a string") + } + if err := counter.add(); err != nil { + return nil, err + } + if _, duplicate := object[key]; duplicate { + return nil, fmt.Errorf("response contains duplicate object key %q", key) + } + value, err := parseValue(decoder, limits, depth+1, counter) + if err != nil { + return nil, err + } + object[key] = value + } + if _, err := decoder.Token(); err != nil { // consume '}' + return nil, fmt.Errorf("decode response object end: %w", err) + } + return object, nil +} + +func parseArray(decoder *json.Decoder, limits Limits, depth int, counter *keyCounter) (any, error) { + array := make([]any, 0) + for decoder.More() { + if len(array) >= limits.MaxArrayLen { + return nil, fmt.Errorf("response exceeds array length budget") + } + value, err := parseValue(decoder, limits, depth+1, counter) + if err != nil { + return nil, err + } + array = append(array, value) + } + if _, err := decoder.Token(); err != nil { // consume ']' + return nil, fmt.Errorf("decode response array end: %w", err) + } + return array, nil +} + +// canonicalJSON reserializes a value tree with sorted object keys so equal +// safe results serialize identically for cassettes and digests. +func canonicalJSON(value any) ([]byte, error) { + var buffer bytes.Buffer + if err := writeCanonical(&buffer, value); err != nil { + return nil, err + } + return buffer.Bytes(), nil +} + +func writeCanonical(buffer *bytes.Buffer, value any) error { + switch v := value.(type) { + case map[string]any: + keys := make([]string, 0, len(v)) + for key := range v { + keys = append(keys, key) + } + sort.Strings(keys) + buffer.WriteByte('{') + for i, key := range keys { + if i > 0 { + buffer.WriteByte(',') + } + encoded, err := json.Marshal(key) + if err != nil { + return err + } + buffer.Write(encoded) + buffer.WriteByte(':') + if err := writeCanonical(buffer, v[key]); err != nil { + return err + } + } + buffer.WriteByte('}') + return nil + case []any: + buffer.WriteByte('[') + for i, item := range v { + if i > 0 { + buffer.WriteByte(',') + } + if err := writeCanonical(buffer, item); err != nil { + return err + } + } + buffer.WriteByte(']') + return nil + default: + encoded, err := json.Marshal(value) + if err != nil { + return err + } + buffer.Write(encoded) + return nil + } +} diff --git a/provider/responsepolicy/policy.go b/provider/responsepolicy/policy.go new file mode 100644 index 00000000..88e0a82e --- /dev/null +++ b/provider/responsepolicy/policy.go @@ -0,0 +1,194 @@ +// Package responsepolicy filters untrusted vendor responses into safe, +// schema-declared bytes. It owns Accept-Encoding decoding, bounds every +// dimension of the response, rejects ambiguous JSON, and applies manifest +// selectors so that FORWARD_SAFE fields are the only values that leave the +// host, secret-bearing fields are captured to a pre-authorized sink or reported +// only by presence, and everything undeclared is dropped. Any drift — a missing +// required field, a value of the wrong shape, or a sink failure — fails closed +// and forwards nothing. +package responsepolicy + +import ( + "context" + "fmt" + "strings" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/manifest" +) + +// Sink durably stores a captured secret and returns only an opaque reference. +// The production backend is out of scope; the host injects a concrete Sink. +type Sink interface { + Put(ctx context.Context, target SinkTarget, secret string) (*providerv0.OpaqueReference, error) +} + +// SinkTarget is the exact pre-authorized capture destination. It is derived by +// the host from the plan, never from provider input. +type SinkTarget struct { + Purpose providerv0.CredentialPurpose + Key string +} + +// Field is one host-derived response disposition. It mirrors the manifest +// response field but adds host-decided required-ness and the pre-authorized +// sink key for captures. +type Field struct { + Selector manifest.Selector + Disposition manifest.ResponseDisposition + Purpose providerv0.CredentialPurpose + Required bool + SinkKey string +} + +// Policy is the host-derived response policy for one admitted method/path/ +// status/content type. +type Policy struct { + Fields []Field + Limits Limits +} + +// CaptureOutcome is the explicit result vocabulary for a capture selector. +type CaptureOutcome string + +const ( + OutcomeCaptured CaptureOutcome = "CAPTURED" + OutcomeAbsent CaptureOutcome = "ABSENT" + OutcomeSinkFailed CaptureOutcome = "SINK_FAILED" +) + +// Forwarded is one safe field cleared for return to the provider. +type Forwarded struct { + Selector string + Value *providerv0.PublicValue +} + +// Capture is the durable result of one capture selector. +type Capture struct { + Selector string + Reference *providerv0.OpaqueReference + Outcome CaptureOutcome +} + +// Result is the fully filtered response. SafeJSON is a canonical projection +// containing only forwarded values and presence markers — never raw bytes. +type Result struct { + Forwarded []Forwarded + Suppressed []string + Captures []Capture + SafeJSON []byte +} + +// RequiresBody reports whether any field must be present. A successful response +// with no body but a policy that requires a field is drift and must fail closed +// rather than silently report success — a vendor cannot skip a required capture +// by returning an empty body. +func (p Policy) RequiresBody() bool { + for _, field := range p.Fields { + if field.Required { + return true + } + } + return false +} + +// Filter decodes, bounds, and filters a response body against the policy. It +// returns an error — and no partial result — whenever any declared behavior +// cannot be honored, so the caller can fail closed without ever forwarding +// original bytes. +func (p Policy) Filter(ctx context.Context, raw []byte, contentEncoding, contentType string, sink Sink) (*Result, error) { + if err := validateJSONContentType(contentType); err != nil { + return nil, err + } + root, err := decodeBody(raw, contentEncoding, p.Limits) + if err != nil { + return nil, err + } + result := &Result{} + safe := map[string]any{} + for _, field := range p.Fields { + matches, err := evaluate(field.Selector, root) + if err != nil { + return nil, fmt.Errorf("selector %s: %w", field.Selector.Path, err) + } + switch field.Disposition { + case manifest.ResponseForwardSafe: + if len(matches) == 0 && field.Required { + return nil, fmt.Errorf("required forward field %s is absent", field.Selector.Path) + } + for _, m := range matches { + value, err := toPublicValue(m.value) + if err != nil { + return nil, fmt.Errorf("forward field %s: %w", m.path, err) + } + result.Forwarded = append(result.Forwarded, Forwarded{Selector: m.path, Value: value}) + safe[m.path] = m.value + } + case manifest.ResponseSuppressPresence: + if len(matches) == 0 && field.Required { + return nil, fmt.Errorf("required suppressed field %s is absent", field.Selector.Path) + } + if len(matches) > 0 { + result.Suppressed = append(result.Suppressed, field.Selector.Path) + safe[field.Selector.Path] = map[string]any{"present": true} + } + case manifest.ResponseCaptureToSink: + captures, err := p.capture(ctx, field, matches, sink) + if err != nil { + return nil, err + } + result.Captures = append(result.Captures, captures...) + for _, capture := range captures { + safe[capture.Selector] = map[string]any{"captured": capture.Outcome == OutcomeCaptured} + } + default: + return nil, fmt.Errorf("field %s has an unknown disposition", field.Selector.Path) + } + } + encoded, err := canonicalJSON(safe) + if err != nil { + return nil, err + } + result.SafeJSON = encoded + return result, nil +} + +// capture resolves and durably stores every matched secret for one capture +// field. It fails closed on a wrong-typed value, a missing required capture, or +// a sink failure so an unstored secret is never silently dropped. +func (p Policy) capture(ctx context.Context, field Field, matches []match, sink Sink) ([]Capture, error) { + if sink == nil { + return nil, fmt.Errorf("capture field %s requires a sink", field.Selector.Path) + } + if len(matches) == 0 { + if field.Required { + return nil, fmt.Errorf("required capture field %s is absent", field.Selector.Path) + } + return []Capture{{Selector: field.Selector.Path, Outcome: OutcomeAbsent}}, nil + } + captures := make([]Capture, 0, len(matches)) + for _, m := range matches { + secret, ok := m.value.(string) + if !ok { + return nil, fmt.Errorf("capture field %s resolved a non-string value", m.path) + } + reference, err := sink.Put(ctx, SinkTarget{Purpose: field.Purpose, Key: field.SinkKey}, secret) + if err != nil { + return nil, fmt.Errorf("capture field %s sink failed: %w", m.path, err) + } + if reference == nil { + return nil, fmt.Errorf("capture field %s sink returned no reference", m.path) + } + captures = append(captures, Capture{Selector: m.path, Reference: reference, Outcome: OutcomeCaptured}) + } + return captures, nil +} + +func validateJSONContentType(contentType string) error { + media, _, _ := strings.Cut(contentType, ";") + media = strings.ToLower(strings.TrimSpace(media)) + if media == "application/json" || strings.HasSuffix(media, "+json") { + return nil + } + return fmt.Errorf("response content type %q is not admitted", contentType) +} diff --git a/provider/responsepolicy/policy_test.go b/provider/responsepolicy/policy_test.go new file mode 100644 index 00000000..0813b733 --- /dev/null +++ b/provider/responsepolicy/policy_test.go @@ -0,0 +1,249 @@ +package responsepolicy_test + +import ( + "bytes" + "compress/gzip" + "context" + "fmt" + "strings" + "testing" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/configuration" + "github.com/codefly-dev/core/provider/manifest" + "github.com/codefly-dev/core/provider/responsepolicy" + "github.com/stretchr/testify/require" +) + +// memorySink is an in-memory capture sink for tests. It records every secret +// it stored so tests can prove the raw bytes went only to the sink. +type memorySink struct { + stored []string + fail bool +} + +func (s *memorySink) Put(_ context.Context, target responsepolicy.SinkTarget, secret string) (*providerv0.OpaqueReference, error) { + if s.fail { + return nil, fmt.Errorf("sink offline") + } + s.stored = append(s.stored, secret) + return &providerv0.OpaqueReference{ + Reference: fmt.Sprintf("capture://%s/%d", target.Key, len(s.stored)), + Purpose: target.Purpose, + }, nil +} + +func sel(path string) manifest.Selector { + return manifest.Selector{Version: manifest.SelectorVersionV1, Path: path} +} + +func forward(path string) responsepolicy.Field { + return responsepolicy.Field{Selector: sel(path), Disposition: manifest.ResponseForwardSafe} +} + +func capture(path string, required bool) responsepolicy.Field { + return responsepolicy.Field{ + Selector: sel(path), + Disposition: manifest.ResponseCaptureToSink, + Purpose: providerv0.CredentialPurpose_CREDENTIAL_PURPOSE_MANAGEMENT, + Required: required, + SinkKey: "binding/secret", + } +} + +func policy(fields ...responsepolicy.Field) responsepolicy.Policy { + return responsepolicy.Policy{Fields: fields, Limits: responsepolicy.DefaultLimits()} +} + +const stripeSecret = "sk_live_1234567890abcdef" + +// assertNoSecret scans arbitrary text for a planted poison value. +func assertNoSecret(t *testing.T, haystack []byte, needle string) { + t.Helper() + require.False(t, bytes.Contains(haystack, []byte(needle)), "poison secret leaked into %q", string(haystack)) +} + +func TestFilter_StripeCreateSecretCaptured(t *testing.T) { + body := fmt.Sprintf(`{"id":"pi_123","object":"payment_intent","client_secret":%q}`, stripeSecret) + sink := &memorySink{} + pol := policy(forward("$.id"), forward("$.object"), capture("$.client_secret", true)) + + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.NoError(t, err) + + require.Len(t, result.Captures, 1) + require.Equal(t, responsepolicy.OutcomeCaptured, result.Captures[0].Outcome) + require.Equal(t, []string{stripeSecret}, sink.stored) + assertNoSecret(t, result.SafeJSON, stripeSecret) + for _, f := range result.Forwarded { + assertNoSecret(t, []byte(f.Value.GetStringValue()), stripeSecret) + } +} + +func TestFilter_ResendListSecretsInArray(t *testing.T) { + body := `{"data":[{"id":"wh_1","secret":"whsec_aaa"},{"id":"wh_2","secret":"whsec_bbb"}]}` + sink := &memorySink{} + pol := policy(forward("$.data[*].id"), capture("$.data[*].secret", false)) + + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.NoError(t, err) + + require.Len(t, result.Captures, 2) + require.ElementsMatch(t, []string{"whsec_aaa", "whsec_bbb"}, sink.stored) + assertNoSecret(t, result.SafeJSON, "whsec_aaa") + assertNoSecret(t, result.SafeJSON, "whsec_bbb") + // Both forwarded ids get distinct concrete selectors. + require.Len(t, result.Forwarded, 2) + require.NotEqual(t, result.Forwarded[0].Selector, result.Forwarded[1].Selector) +} + +func TestFilter_SentryClientKeysBesidePublicDSN(t *testing.T) { + body := `{"data":[{"public":"pub_1","secret":"srt_secret_1","dsn":{"public":"https://pub@sentry.io/1","secret":"https://pub:srt_secret_2@sentry.io/1"}}]}` + sink := &memorySink{} + pol := policy( + forward("$.data[*].public"), + forward("$.data[*].dsn.public"), + capture("$.data[*].secret", true), + capture("$.data[*].dsn.secret", true), + ) + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.NoError(t, err) + require.Len(t, result.Captures, 2) + require.ElementsMatch(t, []string{"srt_secret_1", "https://pub:srt_secret_2@sentry.io/1"}, sink.stored) + assertNoSecret(t, result.SafeJSON, "srt_secret_1") + assertNoSecret(t, result.SafeJSON, "srt_secret_2") +} + +func TestFilter_RequiredCaptureMissingFailsClosed(t *testing.T) { + body := `{"id":"pi_123"}` // client_secret moved/renamed away + sink := &memorySink{} + pol := policy(forward("$.id"), capture("$.client_secret", true)) + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.Error(t, err) + require.Empty(t, sink.stored) +} + +func TestFilter_OptionalCaptureAbsent(t *testing.T) { + body := `{"id":"pi_123"}` + sink := &memorySink{} + pol := policy(forward("$.id"), capture("$.client_secret", false)) + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.NoError(t, err) + require.Len(t, result.Captures, 1) + require.Equal(t, responsepolicy.OutcomeAbsent, result.Captures[0].Outcome) +} + +func TestFilter_CaptureWrongTypeFailsClosed(t *testing.T) { + for _, body := range []string{ + `{"client_secret":null}`, + `{"client_secret":123}`, + `{"client_secret":{"nested":"x"}}`, + } { + sink := &memorySink{} + pol := policy(capture("$.client_secret", true)) + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.Error(t, err, body) + require.Empty(t, sink.stored) + } +} + +func TestFilter_DuplicateJSONKeyRejected(t *testing.T) { + body := `{"id":"a","id":"b"}` + pol := policy(forward("$.id")) + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.ErrorContains(t, err, "duplicate") +} + +func TestFilter_UnknownFieldsDropped(t *testing.T) { + body := fmt.Sprintf(`{"id":"pi_1","surprise":%q}`, stripeSecret) + pol := policy(forward("$.id")) + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.NoError(t, err) + require.Len(t, result.Forwarded, 1) + assertNoSecret(t, result.SafeJSON, stripeSecret) +} + +func TestFilter_ForwardSecretShapedFailsClosed(t *testing.T) { + // A field declared FORWARD_SAFE that drifts into carrying a secret must not + // be forwarded. + body := fmt.Sprintf(`{"token":%q}`, stripeSecret) + require.True(t, configuration.LooksSecret(stripeSecret)) + pol := policy(responsepolicy.Field{Selector: sel("$.token"), Disposition: manifest.ResponseForwardSafe, Required: true}) + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.Error(t, err) +} + +func TestFilter_SinkFailureAfterMatchFailsClosed(t *testing.T) { + body := fmt.Sprintf(`{"client_secret":%q}`, stripeSecret) + sink := &memorySink{fail: true} + pol := policy(capture("$.client_secret", true)) + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", sink) + require.Error(t, err) +} + +func TestFilter_GzipBombRejected(t *testing.T) { + var buf bytes.Buffer + zw := gzip.NewWriter(&buf) + _, _ = zw.Write(bytes.Repeat([]byte("A"), 4<<20)) + _ = zw.Close() + + pol := responsepolicy.Policy{ + Fields: []responsepolicy.Field{forward("$.id")}, + Limits: responsepolicy.Limits{ + MaxCompressedBytes: 1 << 20, MaxDecompressedBytes: 1 << 20, + MaxDepth: 16, MaxKeys: 100, MaxArrayLen: 100, MaxStringBytes: 1024, + }, + } + _, err := pol.Filter(context.Background(), buf.Bytes(), "gzip", "application/json", &memorySink{}) + require.ErrorContains(t, err, "decompressed byte budget") +} + +func TestFilter_CallerControlledEncodingRejected(t *testing.T) { + pol := policy(forward("$.id")) + _, err := pol.Filter(context.Background(), []byte(`{"id":"a"}`), "br", "application/json", &memorySink{}) + require.ErrorContains(t, err, "content encoding") +} + +func TestFilter_DeepNestingRejected(t *testing.T) { + body := strings.Repeat(`{"a":`, 100) + `1` + strings.Repeat(`}`, 100) + pol := responsepolicy.Policy{ + Fields: []responsepolicy.Field{forward("$.a")}, + Limits: responsepolicy.Limits{ + MaxCompressedBytes: 1 << 20, MaxDecompressedBytes: 1 << 20, + MaxDepth: 16, MaxKeys: 1000, MaxArrayLen: 1000, MaxStringBytes: 1024, + }, + } + _, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.ErrorContains(t, err, "nesting budget") +} + +func TestFilter_NonSuccessSecretLookingErrorNotForwarded(t *testing.T) { + // A vendor error body carries an attacker-controlled secret-looking value in + // an undeclared field. Only declared fields survive. + body := fmt.Sprintf(`{"error":{"message":"bad","hint":%q}}`, stripeSecret) + pol := policy(forward("$.error.message")) + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.NoError(t, err) + assertNoSecret(t, result.SafeJSON, stripeSecret) +} + +func TestFilter_SuppressReportsPresenceOnly(t *testing.T) { + body := fmt.Sprintf(`{"masked":%q}`, stripeSecret) + pol := policy(responsepolicy.Field{Selector: sel("$.masked"), Disposition: manifest.ResponseSuppressPresence}) + result, err := pol.Filter(context.Background(), []byte(body), "", "application/json", &memorySink{}) + require.NoError(t, err) + require.Equal(t, []string{"$.masked"}, result.Suppressed) + require.Empty(t, result.Forwarded) + assertNoSecret(t, result.SafeJSON, stripeSecret) +} + +func TestFilter_GzipRoundTrips(t *testing.T) { + var buf bytes.Buffer + zw := gzip.NewWriter(&buf) + _, _ = zw.Write([]byte(`{"id":"pi_9"}`)) + _ = zw.Close() + pol := policy(forward("$.id")) + result, err := pol.Filter(context.Background(), buf.Bytes(), "gzip", "application/json", &memorySink{}) + require.NoError(t, err) + require.Len(t, result.Forwarded, 1) +} diff --git a/provider/responsepolicy/publicvalue.go b/provider/responsepolicy/publicvalue.go new file mode 100644 index 00000000..4297286c --- /dev/null +++ b/provider/responsepolicy/publicvalue.go @@ -0,0 +1,114 @@ +package responsepolicy + +import ( + "encoding/json" + "fmt" + "sort" + "strconv" + "strings" + + providerv0 "github.com/codefly-dev/core/generated/go/codefly/services/provider/v0" + "github.com/codefly-dev/core/provider/configuration" +) + +// toPublicValue converts a parsed JSON value into a normalized PublicValue. +// Every string is checked against the secret heuristic: a FORWARD_SAFE field +// that carries a secret-shaped literal is drift and fails closed rather than +// forwarding the bytes. +func toPublicValue(value any) (*providerv0.PublicValue, error) { + switch v := value.(type) { + case nil: + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_NullValue{NullValue: true}}, nil + case bool: + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_BoolValue{BoolValue: v}}, nil + case string: + if configuration.LooksSecret(v) { + return nil, fmt.Errorf("forwarded field carries a secret-shaped literal") + } + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_StringValue{StringValue: v}}, nil + case json.Number: + return numberToPublicValue(v) + case []any: + list := &providerv0.PublicList{Values: make([]*providerv0.PublicValue, 0, len(v))} + for _, item := range v { + converted, err := toPublicValue(item) + if err != nil { + return nil, err + } + list.Values = append(list.Values, converted) + } + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_ListValue{ListValue: list}}, nil + case map[string]any: + object := &providerv0.PublicObject{Fields: make(map[string]*providerv0.PublicValue, len(v))} + keys := make([]string, 0, len(v)) + for key := range v { + keys = append(keys, key) + } + sort.Strings(keys) + for _, key := range keys { + converted, err := toPublicValue(v[key]) + if err != nil { + return nil, err + } + object.Fields[key] = converted + } + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_ObjectValue{ObjectValue: object}}, nil + default: + return nil, fmt.Errorf("unsupported response value type %T", value) + } +} + +// numberToPublicValue renders a JSON number as either a signed integer or a +// canonical base-ten decimal. Exponent forms and non-canonical spellings fail +// closed so a forwarded number has exactly one representation. +func numberToPublicValue(number json.Number) (*providerv0.PublicValue, error) { + raw := string(number) + if integer, err := strconv.ParseInt(raw, 10, 64); err == nil && strconv.FormatInt(integer, 10) == raw { + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_IntegerValue{IntegerValue: integer}}, nil + } + decimal, err := canonicalDecimal(raw) + if err != nil { + return nil, err + } + return &providerv0.PublicValue{Kind: &providerv0.PublicValue_DecimalValue{DecimalValue: decimal}}, nil +} + +// canonicalDecimal normalizes a plain decimal string to the canonical form +// consumed by the canonical package: no exponent, no leading zeros, no +// insignificant trailing fraction zeros, and no negative zero. +func canonicalDecimal(raw string) (string, error) { + if strings.ContainsAny(raw, "eE") { + return "", fmt.Errorf("decimal must not use exponent notation") + } + negative := strings.HasPrefix(raw, "-") + digits := strings.TrimPrefix(raw, "-") + integerPart, fractionPart, hasFraction := strings.Cut(digits, ".") + if integerPart == "" || !isDigits(integerPart) || (hasFraction && !isDigits(fractionPart)) { + return "", fmt.Errorf("decimal %q is not canonical", raw) + } + integerPart = strings.TrimLeft(integerPart, "0") + if integerPart == "" { + integerPart = "0" + } + fractionPart = strings.TrimRight(fractionPart, "0") + result := integerPart + if fractionPart != "" { + result += "." + fractionPart + } + if negative && result != "0" { + result = "-" + result + } + return result, nil +} + +func isDigits(value string) bool { + if value == "" { + return false + } + for _, r := range value { + if r < '0' || r > '9' { + return false + } + } + return true +} diff --git a/provider/responsepolicy/selector.go b/provider/responsepolicy/selector.go new file mode 100644 index 00000000..42eebd5f --- /dev/null +++ b/provider/responsepolicy/selector.go @@ -0,0 +1,64 @@ +package responsepolicy + +import ( + "strconv" + + "github.com/codefly-dev/core/provider/manifest" +) + +// match is one concrete value a selector resolved to. Path is the concrete +// selector with every wildcard replaced by the matched index, so two elements +// of the same array never collide on one selector string. +type match struct { + path string + value any +} + +// evaluate resolves selector against the value tree. A selector that matches +// nothing returns an empty slice rather than an error: presence is a policy +// decision the caller makes, not a parse failure. +func evaluate(selector manifest.Selector, root any) ([]match, error) { + tokens, err := manifest.ParseSelector(selector) + if err != nil { + return nil, err + } + return walk(tokens, root, "$"), nil +} + +func walk(tokens []manifest.SelectorToken, value any, path string) []match { + if len(tokens) == 0 { + return []match{{path: path, value: value}} + } + token := tokens[0] + rest := tokens[1:] + switch token.Kind { + case manifest.SelectorObjectKey: + object, ok := value.(map[string]any) + if !ok { + return nil + } + child, present := object[token.Key] + if !present { + return nil + } + return walk(rest, child, path+"."+token.Key) + case manifest.SelectorExactIndex: + array, ok := value.([]any) + if !ok || token.Index >= uint64(len(array)) { + return nil + } + return walk(rest, array[token.Index], path+"["+strconv.FormatUint(token.Index, 10)+"]") + case manifest.SelectorArrayWildcard: + array, ok := value.([]any) + if !ok { + return nil + } + var matches []match + for i, item := range array { + matches = append(matches, walk(rest, item, path+"["+strconv.Itoa(i)+"]")...) + } + return matches + default: + return nil + } +}