From 28799bd705da121b891db67a54a146bb1222297e Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Thu, 1 Oct 2026 01:07:50 +0000 Subject: [PATCH 1/3] Add the Session gateway for model, MCP and generic egress The gateway serves, inside the Session's loopback-only network namespace, one listener per frozen model upstream, one per MCP HTTP binding and one generic proxy. Model listeners relay only the routes internal/modelprovider declares, strip every credential header and inject the upstream credential from the agent host. MCP listeners admit each binding's connection origin first, then relay to the server with the bearer injected: environment origin through the sandbox Network service, service origin from the agent host. The generic proxy carries CONNECT tunnels and absolute-form plain HTTP, and connects only through a new sandboxnet stream per connection. Redirects are returned unfollowed and the gateway logs nothing. Plan returns the Endpoints before the view exists, so the Harness environment can be built before sessionview starts; Start opens the listeners from the launcher's network hook at the same fixed ports. --- apps/daemon/internal/gateway/gateway.go | 263 ++++++++++++++++++ apps/daemon/internal/gateway/gateway_test.go | 136 +++++++++ apps/daemon/internal/gateway/listen_linux.go | 49 ++++ apps/daemon/internal/gateway/listen_other.go | 11 + apps/daemon/internal/gateway/mcp.go | 60 ++++ apps/daemon/internal/gateway/mcp_test.go | 63 +++++ apps/daemon/internal/gateway/model.go | 76 +++++ apps/daemon/internal/gateway/model_test.go | 159 +++++++++++ apps/daemon/internal/gateway/proxy.go | 131 +++++++++ apps/daemon/internal/gateway/proxy_test.go | 46 +++ apps/daemon/internal/gateway/relay.go | 144 ++++++++++ .../internal/gateway/view_linux_test.go | 263 ++++++++++++++++++ 12 files changed, 1401 insertions(+) create mode 100644 apps/daemon/internal/gateway/gateway.go create mode 100644 apps/daemon/internal/gateway/gateway_test.go create mode 100644 apps/daemon/internal/gateway/listen_linux.go create mode 100644 apps/daemon/internal/gateway/listen_other.go create mode 100644 apps/daemon/internal/gateway/mcp.go create mode 100644 apps/daemon/internal/gateway/mcp_test.go create mode 100644 apps/daemon/internal/gateway/model.go create mode 100644 apps/daemon/internal/gateway/model_test.go create mode 100644 apps/daemon/internal/gateway/proxy.go create mode 100644 apps/daemon/internal/gateway/proxy_test.go create mode 100644 apps/daemon/internal/gateway/relay.go create mode 100644 apps/daemon/internal/gateway/view_linux_test.go diff --git a/apps/daemon/internal/gateway/gateway.go b/apps/daemon/internal/gateway/gateway.go new file mode 100644 index 00000000..aa926f5c --- /dev/null +++ b/apps/daemon/internal/gateway/gateway.go @@ -0,0 +1,263 @@ +// Package gateway is the Session gateway on the agent host. Inside the +// Session's loopback-only network namespace it serves one listener per frozen +// model upstream, one per MCP HTTP binding and one generic proxy, so the +// Harness never holds an upstream credential and has no network route of its +// own. +// +// A listener's identity selects its upstream and credential; nothing is routed +// by hostname. A model listener relays the declared native routes of its +// protocol (internal/modelprovider) to the upstream from the agent host and +// injects the credential. An MCP listener relays to its binding's server and +// injects the bearer token: an environment-origin binding connects through the +// sandbox's Network service, a service-origin binding from the agent host, and +// the gateway does the TLS either way. The generic proxy carries HTTP CONNECT +// tunnels and plain-HTTP forward requests, and connects only through the +// sandbox's Network service. Redirects reach the Harness unchanged and are +// never followed. The gateway logs nothing. +package gateway + +import ( + "context" + "crypto/tls" + "errors" + "fmt" + "io" + "log" + "net" + "net/http" + "os" + "strconv" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink" +) + +// Config is what one Session's gateway serves. It holds credentials: keep it +// in memory and never log it. +type Config struct { + // Models are the frozen model upstreams, each under the adapter's name for + // it. Names are unique. + Models []Model + // MCP are the Session's MCP HTTP bindings. Server labels are unique. + MCP []proto.MCPHTTPServer + // Prompt is the request that declared MCP. Each binding's origin is + // admitted against its placement with ValidateConnectionOrigin before + // anything else. + Prompt proto.PromptRequestPayload + // OpenNetwork opens a new Network stream to the Session's sandbox, as + // sandboxnet.Connect takes it. Nil means the Session has no sandbox + // network: the generic proxy then refuses every request, and an + // environment-origin binding is invalid. + OpenNetwork func(context.Context) (sandboxlink.Stream, error) + // TLS configures the gateway's upstream TLS. Nil uses the system roots. + TLS *tls.Config +} + +// Model is one frozen model upstream. +type Model struct { + Name string + Provider modelprovider.Provider +} + +// Endpoints is what the Harness is given in place of upstreams and +// credentials. Every URL is plain HTTP on the Session's loopback. +type Endpoints struct { + // Placeholder is the credential a Harness sends to a model listener. It + // is not secret; the listener removes it. + Placeholder string + // Models maps each model upstream's name to its listener's base URL, + // http://127.0.0.1:, with no path. The listener adds the frozen + // base URL's path. + Models map[string]string + // MCP maps each binding's server label to the URL the Harness uses: its + // listener with the server URL's path and query. + MCP map[string]string + // Proxy is the generic proxy's URL, for HTTP and HTTPS proxy settings. + Proxy string +} + +// SessionNetwork is the Session's network namespace: the file sessionview's +// network hook receives. Start uses it only while it runs. +type SessionNetwork struct { + Namespace *os.File +} + +// ProxyPort is the generic proxy's port in the Session's namespace. The model +// listeners take the following ports in Config order, then the MCP listeners. +// The namespace is the Session's own and the gateway listens before the +// Harness starts, so the ports are free; fixing them lets the Harness's +// environment be built before the namespace exists. +const ProxyPort = 17100 + +// maxListeners bounds the listeners of one Session. +const maxListeners = 256 + +var ( + // ErrInvalidConfig is a Config that Plan and Start reject. The message + // names the item by its name or label and never includes a credential. + ErrInvalidConfig = errors.New("gateway: invalid configuration") + // ErrNetwork is a failure to listen in the Session's network namespace. + ErrNetwork = errors.New("gateway: session network") + // ErrUnsupported is returned by Start outside Linux. + ErrUnsupported = errors.New("gateway: unsupported platform") +) + +// Plan validates cfg and returns the Endpoints that Start serves for it, so +// the Harness's environment can be built before the view starts. +func Plan(cfg Config) (Endpoints, error) { + g, err := build(cfg) + if err != nil { + return Endpoints{}, err + } + return g.endpoints(fixedPorts(len(g.listeners))), nil +} + +// Start validates cfg, opens its listeners inside the Session's network +// namespace and serves them from the daemon until ctx ends. It returns the +// same Endpoints as Plan(cfg). It is meant to run in sessionview's network +// hook; nothing listens outside the namespace. The end of ctx closes the +// listeners and every connection. +func Start(ctx context.Context, n SessionNetwork, cfg Config) (Endpoints, error) { + g, err := build(cfg) + if err != nil { + return Endpoints{}, err + } + if n.Namespace == nil { + return Endpoints{}, fmt.Errorf("%w: no namespace", ErrNetwork) + } + ports := fixedPorts(len(g.listeners)) + lns, err := listen(n.Namespace, ports) + if err != nil { + return Endpoints{}, err + } + g.serve(ctx, lns) + return g.endpoints(ports), nil +} + +func fixedPorts(n int) []int { + ports := make([]int, n) + for i := range ports { + ports[i] = ProxyPort + i + } + return ports +} + +type role uint8 + +const ( + roleProxy role = iota + roleModel + roleMCP +) + +// listener is one planned listener: what it serves and how the Harness +// addresses it. +type listener struct { + role role + name string // model name or MCP server label + suffix string // MCP: the server URL's path and query + handler http.Handler +} + +type gateway struct { + listeners []listener // the proxy, then models, then MCP, in Config order + transports []*http.Transport +} + +func invalid(format string, args ...any) error { + return fmt.Errorf("%w: "+format, append([]any{ErrInvalidConfig}, args...)...) +} + +// build validates cfg and makes each listener's handler. +func build(cfg Config) (*gateway, error) { + if n := 1 + len(cfg.Models) + len(cfg.MCP); n > maxListeners { + return nil, invalid("%d listeners, at most %d", n, maxListeners) + } + host := newTransport(cfg.TLS, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext) + g := &gateway{transports: []*http.Transport{host}} + var sandbox *http.Transport + if cfg.OpenNetwork != nil { + sandbox = newTransport(cfg.TLS, sandboxDialer(cfg.OpenNetwork)) + g.transports = append(g.transports, sandbox) + } + g.listeners = append(g.listeners, listener{role: roleProxy, handler: newProxy(cfg.OpenNetwork, sandbox)}) + + names := map[string]bool{} + for _, m := range cfg.Models { + if m.Name == "" || names[m.Name] { + return nil, invalid("model upstream name %q is empty or repeated", m.Name) + } + names[m.Name] = true + h, err := newModelRelay(m.Provider, host) + if err != nil { + return nil, invalid("model upstream %q: %v", m.Name, err) + } + g.listeners = append(g.listeners, listener{role: roleModel, name: m.Name, handler: h}) + } + + labels := map[string]bool{} + for _, s := range cfg.MCP { + if err := s.ValidateConnectionOrigin(cfg.Prompt); err != nil { + return nil, invalid("MCP server %q: %v", s.ServerLabel, err) + } + if s.ServerLabel == "" || labels[s.ServerLabel] { + return nil, invalid("MCP server label %q is empty or repeated", s.ServerLabel) + } + labels[s.ServerLabel] = true + transport := host + if s.ConnectionOrigin == "environment" { + if sandbox == nil { + return nil, invalid("MCP server %q has environment origin and the Session has no sandbox network", s.ServerLabel) + } + transport = sandbox + } + h, suffix, err := newMCPRelay(s, transport) + if err != nil { + return nil, invalid("MCP server %q: %v", s.ServerLabel, err) + } + g.listeners = append(g.listeners, listener{role: roleMCP, name: s.ServerLabel, suffix: suffix, handler: h}) + } + return g, nil +} + +func (g *gateway) endpoints(ports []int) Endpoints { + e := Endpoints{Placeholder: modelprovider.Placeholder, Models: map[string]string{}, MCP: map[string]string{}} + for i, l := range g.listeners { + base := "http://" + net.JoinHostPort("127.0.0.1", strconv.Itoa(ports[i])) + switch l.role { + case roleProxy: + e.Proxy = base + case roleModel: + e.Models[l.name] = base + case roleMCP: + e.MCP[l.name] = base + l.suffix + } + } + return e +} + +// quiet discards what net/http would log: a logged request or upstream error +// could carry more than the gateway chooses to reveal. +var quiet = log.New(io.Discard, "", 0) + +// serve serves each listener with its handler until ctx ends. +func (g *gateway) serve(ctx context.Context, lns []net.Listener) { + for i, ln := range lns { + srv := &http.Server{ + Handler: g.listeners[i].handler, + ReadHeaderTimeout: 30 * time.Second, + ErrorLog: quiet, + // Requests, hijacked tunnels and upgraded connections end with ctx. + BaseContext: func(net.Listener) context.Context { return ctx }, + } + go srv.Serve(ln) + context.AfterFunc(ctx, func() { srv.Close() }) + } + context.AfterFunc(ctx, func() { + for _, t := range g.transports { + t.CloseIdleConnections() + } + }) +} diff --git a/apps/daemon/internal/gateway/gateway_test.go b/apps/daemon/internal/gateway/gateway_test.go new file mode 100644 index 00000000..c09767d3 --- /dev/null +++ b/apps/daemon/internal/gateway/gateway_test.go @@ -0,0 +1,136 @@ +package gateway + +import ( + "context" + "crypto/tls" + "crypto/x509" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "sync/atomic" + "testing" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink/relay" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink/sandboxlinktest" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxnet" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +const wait = 5 * time.Second + +// serveOnLoopback serves cfg on host loopback listeners at free ports, as +// Start serves it in the Session's namespace, until the test ends. +func serveOnLoopback(t *testing.T, cfg Config) Endpoints { + t.Helper() + g, err := build(cfg) + if err != nil { + t.Fatal(err) + } + lns := make([]net.Listener, len(g.listeners)) + ports := make([]int, len(lns)) + for i := range lns { + if lns[i], err = net.Listen("tcp4", "127.0.0.1:0"); err != nil { + t.Fatal(err) + } + ports[i] = lns[i].Addr().(*net.TCPAddr).Port + } + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + g.serve(ctx, lns) + return g.endpoints(ports) +} + +// trust returns a TLS configuration that trusts srv. +func trust(srv *httptest.Server) *tls.Config { + roots := x509.NewCertPool() + roots.AddCert(srv.Certificate()) + return &tls.Config{RootCAs: roots} +} + +// noRedirects is a Harness-side client that shows each answer as it is. +var noRedirects = &http.Client{ + Transport: &http.Transport{Proxy: nil}, + CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }, + Timeout: wait, +} + +// sandbox serves the Network protocol behind a test relay, as oac-sandbox-io +// does, and counts the connections it makes. +type sandbox struct { + open func(context.Context) (sandboxlink.Stream, error) + dials atomic.Int32 +} + +func (s *sandbox) Resolve(ctx context.Context, host string) ([]netip.Addr, error) { + return net.DefaultResolver.LookupNetIP(ctx, "ip", host) +} + +func (s *sandbox) Dial(ctx context.Context, addr netip.AddrPort) (*net.TCPConn, error) { + s.dials.Add(1) + var d net.Dialer + c, err := d.DialContext(ctx, "tcp", addr.String()) + if err != nil { + return nil, err + } + return c.(*net.TCPConn), nil +} + +func startSandbox(t *testing.T) *sandbox { + t.Helper() + auth := sandboxlinktest.NewAuthority() + srv := sandboxlinktest.StartRelay(t, relay.Config{Authority: auth}) + resource := sandboxlink.ResourceRef{TenantID: sandboxwire.NewID(), EnvironmentID: sandboxwire.NewID(), + Kind: sandboxlink.ResourceAllocation, ID: sandboxwire.NewID(), Generation: 1} + auth.AddServe([]byte("serve credential"), sandboxlink.ServePeer{PeerID: sandboxwire.NewID(), Resource: resource}) + + sb := &sandbox{} + connected := make(chan struct{}, 1) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + sandboxlink.Serve(ctx, sandboxlink.ServeConfig{URL: srv.URL, TLS: srv.TLS, Credential: []byte("serve credential"), + Resource: resource, ServerInstanceID: sandboxwire.NewID(), + Services: []sandboxlink.ServiceHandler{{Service: sandboxlink.ServiceNetwork, Version: sandboxnet.Version, + Serve: func(ctx context.Context, b sandboxlink.Bind, s sandboxlink.Stream) { + sandboxnet.Serve(ctx, s, b.Egress, sb) + }}}, + OnConnected: func(sandboxlink.HelloAccepted) { + select { + case connected <- struct{}{}: + default: + } + }}) + }() + t.Cleanup(func() { + cancel() + <-done + }) + select { + case <-connected: + case <-time.After(wait): + t.Fatal("the sandbox did not connect to the relay") + } + + runtimeID := sandboxwire.NewID() + auth.AddRuntime([]byte("runtime credential"), runtimeID) + link, err := sandboxlink.DialAttach(context.Background(), sandboxlink.AttachConfig{URL: srv.URL, TLS: srv.TLS, + RuntimeID: runtimeID, Credential: []byte("runtime credential")}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { link.Close() }) + session, assignment := sandboxwire.NewID(), sandboxwire.NewID() + sb.open = func(ctx context.Context) (sandboxlink.Stream, error) { + grant, attachment := sandboxwire.NewID(), sandboxwire.NewID() + auth.AddGrant(grant[:], sandboxlinktest.Grant{RuntimeID: runtimeID, Resource: resource, SessionID: session, AssignmentID: assignment, + AssignmentEpoch: 1, Services: []sandboxlink.Service{sandboxlink.ServiceNetwork}, Lease: time.Minute}) + s, _, err := link.OpenService(ctx, sandboxlink.Open{Service: sandboxlink.ServiceNetwork, Version: sandboxnet.Version, Resource: resource, + AttachmentID: attachment, SessionID: session, AssignmentID: assignment, AssignmentEpoch: 1, AttachGrant: grant[:]}) + return s, err + } + return sb +} diff --git a/apps/daemon/internal/gateway/listen_linux.go b/apps/daemon/internal/gateway/listen_linux.go new file mode 100644 index 00000000..1f0a453c --- /dev/null +++ b/apps/daemon/internal/gateway/listen_linux.go @@ -0,0 +1,49 @@ +//go:build linux + +package gateway + +import ( + "fmt" + "net" + "os" + "runtime" + "strconv" + + "golang.org/x/sys/unix" +) + +// listen opens a TCP listener on 127.0.0.1 at each port inside the network +// namespace ns. A socket stays in the namespace it was created in, so the +// listeners are opened on a thread that joined ns and then served from the +// daemon's own threads. +func listen(ns *os.File, ports []int) ([]net.Listener, error) { + type result struct { + lns []net.Listener + err error + } + done := make(chan result, 1) + go func() { + // Never unlocked: the thread exits with the goroutine instead of + // returning to the scheduler inside the Session's namespace. + runtime.LockOSThread() + if err := unix.Setns(int(ns.Fd()), unix.CLONE_NEWNET); err != nil { + done <- result{err: fmt.Errorf("%w: join: %w", ErrNetwork, err)} + return + } + var lns []net.Listener + for _, port := range ports { + ln, err := net.Listen("tcp4", net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + if err != nil { + for _, l := range lns { + l.Close() + } + done <- result{err: fmt.Errorf("%w: listen: %w", ErrNetwork, err)} + return + } + lns = append(lns, ln) + } + done <- result{lns: lns} + }() + r := <-done + return r.lns, r.err +} diff --git a/apps/daemon/internal/gateway/listen_other.go b/apps/daemon/internal/gateway/listen_other.go new file mode 100644 index 00000000..e326bcf6 --- /dev/null +++ b/apps/daemon/internal/gateway/listen_other.go @@ -0,0 +1,11 @@ +//go:build !linux + +package gateway + +import ( + "net" + "os" +) + +// listen needs Linux network namespaces. +func listen(*os.File, []int) ([]net.Listener, error) { return nil, ErrUnsupported } diff --git a/apps/daemon/internal/gateway/mcp.go b/apps/daemon/internal/gateway/mcp.go new file mode 100644 index 00000000..4f03207d --- /dev/null +++ b/apps/daemon/internal/gateway/mcp.go @@ -0,0 +1,60 @@ +package gateway + +import ( + "errors" + "net/http" + "net/http/httputil" + "net/url" + + "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" +) + +// mcpRelay serves one MCP HTTP binding: it relays each request to the +// server's origin with the same path and query, and injects the binding's +// bearer token when it has one. +type mcpRelay struct { + scheme string + host string + token *string + transport http.RoundTripper +} + +// newMCPRelay returns the binding's handler and the path and query the +// Harness appends to the listener's address. +func newMCPRelay(s proto.MCPHTTPServer, t http.RoundTripper) (*mcpRelay, string, error) { + u, err := url.Parse(s.ServerURL) + if err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Hostname() == "" || u.User != nil || u.Opaque != "" || u.Fragment != "" { + return nil, "", errors.New("server URL is not an absolute http or https URL") + } + m := &mcpRelay{scheme: u.Scheme, host: u.Host, transport: t} + if s.BearerToken != nil { + if u.Scheme != "https" { + return nil, "", errors.New("a bearer token needs an https server URL") + } + token := *s.BearerToken + m.token = &token + } + suffix := u.EscapedPath() + if u.RawQuery != "" || u.ForceQuery { + suffix += "?" + u.RawQuery + } + return m, suffix, nil +} + +func (m *mcpRelay) ServeHTTP(w http.ResponseWriter, r *http.Request) { + path, query, hasQuery, ok := requestTarget(r) + if !ok { + http.Error(w, "origin-form request target required", http.StatusBadRequest) + return + } + upstream := &url.URL{Scheme: m.scheme, Host: m.host, Path: r.URL.Path, RawPath: path, + RawQuery: query, ForceQuery: hasQuery && query == ""} + reverseProxy(m.transport, func(pr *httputil.ProxyRequest) { + pr.Out.URL = upstream + pr.Out.Host = "" + if m.token != nil { + stripCredentials(pr.Out.Header) + pr.Out.Header.Set("Authorization", "Bearer "+*m.token) + } + }).ServeHTTP(w, r) +} diff --git a/apps/daemon/internal/gateway/mcp_test.go b/apps/daemon/internal/gateway/mcp_test.go new file mode 100644 index 00000000..f3358c0a --- /dev/null +++ b/apps/daemon/internal/gateway/mcp_test.go @@ -0,0 +1,63 @@ +package gateway + +import ( + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" +) + +func TestMCPBrokersBothOrigins(t *testing.T) { + sb := startSandbox(t) + seen := make(chan string, 1) + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + seen <- r.Header.Get("Authorization") + " " + r.URL.RequestURI() + io.WriteString(w, "tools") + })) + defer srv.Close() + token := "vault-token" + binding := func(origin string) proto.MCPHTTPServer { + return proto.MCPHTTPServer{ConnectionOrigin: origin, ServerLabel: "tools", ServerURL: srv.URL + "/mcp", BearerToken: &token} + } + service := proto.PromptRequestPayload{DisableExecutionEnvironment: true} + environment := proto.PromptRequestPayload{LocalEnvironment: &proto.LocalEnvironment{NetworkAccess: "enabled"}} + + for _, c := range []struct { + origin string + prompt proto.PromptRequestPayload + dials int32 + }{{"service", service, 0}, {"environment", environment, 1}} { + before := sb.dials.Load() + eps := serveOnLoopback(t, Config{MCP: []proto.MCPHTTPServer{binding(c.origin)}, Prompt: c.prompt, OpenNetwork: sb.open, TLS: trust(srv)}) + if !strings.HasPrefix(eps.MCP["tools"], "http://127.0.0.1:") || !strings.HasSuffix(eps.MCP["tools"], "/mcp") { + t.Fatalf("%s: Harness URL %q", c.origin, eps.MCP["tools"]) + } + req, _ := http.NewRequest("POST", eps.MCP["tools"], strings.NewReader(`{"jsonrpc":"2.0"}`)) + req.Header.Set("Authorization", "Bearer harness-value") + resp, err := noRedirects.Do(req) + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + if resp.StatusCode != 200 || string(body) != "tools" { + t.Fatalf("%s: %d %q", c.origin, resp.StatusCode, body) + } + if got := <-seen; got != "Bearer vault-token /mcp" { + t.Errorf("%s: server saw %q", c.origin, got) + } + if n := sb.dials.Load() - before; n != c.dials { + t.Errorf("%s: the sandbox made %d connections, want %d", c.origin, n, c.dials) + } + } + + // Origin admission runs first: an environment binding needs an enabled + // workspace network. + if _, err := Plan(Config{MCP: []proto.MCPHTTPServer{binding("environment")}, Prompt: service, OpenNetwork: sb.open}); !errors.Is(err, ErrInvalidConfig) { + t.Errorf("Plan with a relocated binding: %v", err) + } +} diff --git a/apps/daemon/internal/gateway/model.go b/apps/daemon/internal/gateway/model.go new file mode 100644 index 00000000..b768a58e --- /dev/null +++ b/apps/daemon/internal/gateway/model.go @@ -0,0 +1,76 @@ +package gateway + +import ( + "errors" + "net/http" + "net/http/httputil" + "net/url" + + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" +) + +// modelRelay serves one frozen model upstream. It relays only the routes the +// protocol declares, injects the declared credential and keeps the base URL's +// path. +type modelRelay struct { + protocol modelprovider.Protocol + scheme string + host string + targets map[string]*url.URL // upstream path of each declared route, by route path + credential modelprovider.Credential + key string + transport http.RoundTripper +} + +func newModelRelay(p modelprovider.Provider, t http.RoundTripper) (*modelRelay, error) { + if err := p.Validate(); err != nil { + return nil, err + } + credential, err := modelprovider.UpstreamCredential(p.Protocol) + if err != nil { + return nil, err + } + base, err := url.Parse(p.BaseURL) + if err != nil { + return nil, err + } + m := &modelRelay{protocol: p.Protocol, scheme: base.Scheme, host: base.Host, targets: map[string]*url.URL{}, + credential: credential, key: p.APIKey, transport: t} + for _, route := range modelprovider.Routes(p.Protocol) { + raw := modelprovider.UpstreamPath(base.EscapedPath(), route.Path) + path, err := url.PathUnescape(raw) + if err != nil { + return nil, err + } + m.targets[route.Path] = &url.URL{Path: path, RawPath: raw} + } + return m, nil +} + +func (m *modelRelay) ServeHTTP(w http.ResponseWriter, r *http.Request) { + path, query, hasQuery, ok := requestTarget(r) + route, err := modelprovider.LookupRoute(m.protocol, r.Method, path) + switch { + case !ok: + http.NotFound(w, r) + return + case errors.Is(err, modelprovider.ErrMethodNotAllowed): + http.Error(w, "method not declared for this route", http.StatusMethodNotAllowed) + return + case err != nil: + http.NotFound(w, r) + return + case wantsUpgrade(r) && !route.WebSocket: + http.Error(w, "upgrade not declared for this route", http.StatusBadRequest) + return + } + target := m.targets[route.Path] + upstream := &url.URL{Scheme: m.scheme, Host: m.host, Path: target.Path, RawPath: target.RawPath, + RawQuery: query, ForceQuery: hasQuery && query == ""} + reverseProxy(m.transport, func(pr *httputil.ProxyRequest) { + pr.Out.URL = upstream + pr.Out.Host = "" + stripCredentials(pr.Out.Header) + pr.Out.Header.Set(m.credential.Header, m.credential.Value(m.key)) + }).ServeHTTP(w, r) +} diff --git a/apps/daemon/internal/gateway/model_test.go b/apps/daemon/internal/gateway/model_test.go new file mode 100644 index 00000000..2aba3d3e --- /dev/null +++ b/apps/daemon/internal/gateway/model_test.go @@ -0,0 +1,159 @@ +package gateway + +import ( + "bufio" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" +) + +const upstreamKey = "sk-upstream-key" + +// startModel serves one Anthropic upstream at the fake TLS upstream's /anthropic. +func startModel(t *testing.T, handler http.HandlerFunc) (string, *httptest.Server) { + t.Helper() + srv := httptest.NewTLSServer(handler) + t.Cleanup(srv.Close) + eps := serveOnLoopback(t, Config{ + Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: srv.URL + "/anthropic", APIKey: upstreamKey}}}, + TLS: trust(srv), + }) + return eps.Models["main"], srv +} + +func TestModelInjectsTheKeyAndNeverThePlaceholder(t *testing.T) { + seen := make(chan *http.Request, 1) + base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + r.Body = io.NopCloser(strings.NewReader(string(body))) + seen <- r + io.WriteString(w, "answer") + }) + req, _ := http.NewRequest("POST", base+"/v1/messages?beta=true", strings.NewReader(`{"model":"m"}`)) + req.Header["authorization"] = []string{"Bearer " + modelprovider.Placeholder} + req.Header["X-API-KEY"] = []string{modelprovider.Placeholder, modelprovider.Placeholder} + req.Header.Set("Cookie", "session="+modelprovider.Placeholder) + req.Header.Set("Anthropic-Version", "2023-06-01") + resp, err := noRedirects.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + if body, _ := io.ReadAll(resp.Body); resp.StatusCode != 200 || string(body) != "answer" { + t.Fatalf("answer %d %q", resp.StatusCode, body) + } + r := <-seen + if r.URL.RequestURI() != "/anthropic/v1/messages?beta=true" { + t.Errorf("upstream target %q", r.URL.RequestURI()) + } + if got := r.Header.Values("X-Api-Key"); len(got) != 1 || got[0] != upstreamKey { + t.Errorf("upstream X-Api-Key %q", got) + } + if r.Header.Get("Anthropic-Version") != "2023-06-01" { + t.Error("a native header was not relayed") + } + for name, values := range r.Header { + for _, v := range values { + if strings.Contains(v, modelprovider.Placeholder) { + t.Errorf("upstream received the placeholder in %s", name) + } + } + } + if body, _ := io.ReadAll(r.Body); string(body) != `{"model":"m"}` { + t.Errorf("upstream body %q", body) + } +} + +func TestModelRejectsUndeclaredRequests(t *testing.T) { + var hits atomic.Int32 + base, _ := startModel(t, func(http.ResponseWriter, *http.Request) { hits.Add(1) }) + cases := []struct { + method, path string + upgrade bool + want int + }{ + {"GET", "/v1/messages", false, http.StatusMethodNotAllowed}, + {"POST", "/v1/complete", false, http.StatusNotFound}, + {"POST", "/anthropic/v1/messages", false, http.StatusNotFound}, + {"POST", "/v1/%6Dessages", false, http.StatusNotFound}, + {"POST", "/v1/messages", true, http.StatusBadRequest}, + } + for _, c := range cases { + req, _ := http.NewRequest(c.method, base+c.path, nil) + if c.upgrade { + req.Header.Set("Connection", "Upgrade") + req.Header.Set("Upgrade", "websocket") + } + resp, err := noRedirects.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != c.want { + t.Errorf("%s %s upgrade=%v: %d, want %d", c.method, c.path, c.upgrade, resp.StatusCode, c.want) + } + } + if n := hits.Load(); n != 0 { + t.Errorf("the upstream was reached %d times", n) + } +} + +func TestModelStreamsEvents(t *testing.T) { + // The upstream holds the second event until the first reached the + // Harness, or until the wait ends. + release := make(chan struct{}) + var once sync.Once + var released atomic.Bool + timer := time.AfterFunc(wait, func() { + released.Store(true) + once.Do(func() { close(release) }) + }) + base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/event-stream") + io.WriteString(w, "data: one\n\n") + w.(http.Flusher).Flush() + <-release + io.WriteString(w, "data: two\n\n") + }) + resp, err := (&http.Client{Transport: &http.Transport{Proxy: nil}}).Post(base+"/v1/messages", "application/json", strings.NewReader("{}")) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + line, err := bufio.NewReader(resp.Body).ReadString('\n') + if err != nil || line != "data: one\n" { + t.Fatalf("first event %q, %v", line, err) + } + if released.Load() { + t.Fatal("the first event arrived only after the upstream finished") + } + timer.Stop() + once.Do(func() { close(release) }) +} + +func TestModelReturnsRedirectsUnfollowed(t *testing.T) { + var hits atomic.Int32 + const location = "https://elsewhere.invalid/v1/messages?x=1" + base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + http.Redirect(w, r, location, http.StatusTemporaryRedirect) + }) + resp, err := noRedirects.Post(base+"/v1/messages", "application/json", strings.NewReader("{}")) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusTemporaryRedirect || resp.Header.Get("Location") != location { + t.Errorf("answer %d Location %q", resp.StatusCode, resp.Header.Get("Location")) + } + if n := hits.Load(); n != 1 { + t.Errorf("the upstream was reached %d times", n) + } +} diff --git a/apps/daemon/internal/gateway/proxy.go b/apps/daemon/internal/gateway/proxy.go new file mode 100644 index 00000000..6cde3430 --- /dev/null +++ b/apps/daemon/internal/gateway/proxy.go @@ -0,0 +1,131 @@ +package gateway + +import ( + "context" + "io" + "net" + "net/http" + "net/http/httputil" + "strings" + "sync" + + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxnet" +) + +// proxy is the generic proxy. It serves CONNECT tunnels and absolute-form +// plain-HTTP requests, and connects only through a new Network stream to the +// sandbox for each connection. It rejects every other request form. +type proxy struct { + open func(context.Context) (sandboxlink.Stream, error) + forward *httputil.ReverseProxy +} + +func newProxy(open func(context.Context) (sandboxlink.Stream, error), sandbox *http.Transport) *proxy { + p := &proxy{open: open} + if sandbox != nil { + // The request is relayed as the Harness addressed it; the reverse + // proxy drops hop-by-hop headers, Proxy-Authorization among them. + p.forward = reverseProxy(sandbox, func(*httputil.ProxyRequest) {}) + } + return p +} + +func (p *proxy) ServeHTTP(w http.ResponseWriter, r *http.Request) { + switch { + case p.open == nil: + http.Error(w, "the Session has no sandbox network", http.StatusForbidden) + case r.Method == http.MethodConnect && r.URL.Host != "" && r.URL.Path == "": + p.tunnel(w, r) + case r.Method != http.MethodConnect && r.URL.Scheme == "http" && r.URL.Host != "" && !strings.HasPrefix(r.RequestURI, "/"): + p.forward.ServeHTTP(w, r) + default: + http.Error(w, "the proxy serves CONNECT and absolute-form http requests only", http.StatusBadRequest) + } +} + +// tunnel connects through the sandbox, answers 200 and then carries raw bytes. +func (p *proxy) tunnel(w http.ResponseWriter, r *http.Request) { + host, port, err := splitHostPort(r.URL.Host) + if err != nil { + http.Error(w, "invalid CONNECT authority", http.StatusBadRequest) + return + } + remote, err := connectSandbox(r.Context(), p.open, host, port) + if err != nil { + status := statusOf(err) + http.Error(w, http.StatusText(status), status) + return + } + client, rw, err := http.NewResponseController(w).Hijack() + if err != nil { + remote.Reset() + http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError) + return + } + // Bytes the Harness sent after the request, such as a TLS ClientHello, + // may already be buffered. + pending, _ := rw.Reader.Peek(rw.Reader.Buffered()) + if _, err := io.WriteString(client, "HTTP/1.1 200 Connection established\r\n\r\n"); err != nil { + remote.Reset() + client.Close() + return + } + splice(r.Context(), client, pending, remote) +} + +// splice carries bytes both ways between the Harness's connection and the +// sandbox's. Each direction's orderly end reaches the other side as a +// half-close after every byte before it, and starts no timeout. An error on +// either side, or the end of ctx, aborts both. +func splice(ctx context.Context, client net.Conn, pending []byte, remote *sandboxnet.Conn) { + var end sync.Once + abort := func() { + end.Do(func() { + remote.Reset() + if tcp, ok := client.(*net.TCPConn); ok { + tcp.SetLinger(0) + } + client.Close() + }) + } + stop := context.AfterFunc(ctx, abort) + defer stop() + + out := make(chan struct{}) + go func() { + defer close(out) + var err error + if len(pending) > 0 { + _, err = remote.Write(pending) + } + if err == nil { + _, err = io.Copy(struct{ io.Writer }{remote}, struct{ io.Reader }{client}) + } + if err == nil { + err = remote.CloseWrite() + } + if err != nil { + abort() + } + }() + _, err := io.Copy(struct{ io.Writer }{client}, struct{ io.Reader }{remote}) + if err == nil { + err = closeWrite(client) + } + if err != nil { + abort() + } + <-out + end.Do(func() { + remote.Close() + client.Close() + }) +} + +func closeWrite(c net.Conn) error { + if cw, ok := c.(interface{ CloseWrite() error }); ok { + return cw.CloseWrite() + } + return c.Close() +} diff --git a/apps/daemon/internal/gateway/proxy_test.go b/apps/daemon/internal/gateway/proxy_test.go new file mode 100644 index 00000000..1be49e3d --- /dev/null +++ b/apps/daemon/internal/gateway/proxy_test.go @@ -0,0 +1,46 @@ +package gateway + +import ( + "io" + "net/http" + "net/http/httptest" + "net/url" + "testing" +) + +func TestProxyConnectsOnlyThroughTheSandbox(t *testing.T) { + sb := startSandbox(t) + eps := serveOnLoopback(t, Config{OpenNetwork: sb.open}) + hello := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "hello "+r.URL.Path) }) + secure := httptest.NewTLSServer(hello) + defer secure.Close() + plain := httptest.NewServer(hello) + defer plain.Close() + + proxyURL, _ := url.Parse(eps.Proxy) + client := &http.Client{Transport: &http.Transport{Proxy: http.ProxyURL(proxyURL), TLSClientConfig: trust(secure)}, Timeout: wait} + for i, target := range []string{secure.URL + "/via-connect", plain.URL + "/via-forward"} { + resp, err := client.Get(target) + if err != nil { + t.Fatalf("%s: %v", target, err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + if want := "hello " + []string{"/via-connect", "/via-forward"}[i]; resp.StatusCode != 200 || string(body) != want { + t.Errorf("%s: %d %q", target, resp.StatusCode, body) + } + } + if n := sb.dials.Load(); n != 2 { + t.Errorf("the sandbox made %d connections, want 2", n) + } + + // An origin-form request is not a proxy request. + resp, err := noRedirects.Get(eps.Proxy + "/") + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusBadRequest { + t.Errorf("origin-form request: %d", resp.StatusCode) + } +} diff --git a/apps/daemon/internal/gateway/relay.go b/apps/daemon/internal/gateway/relay.go new file mode 100644 index 00000000..566972e9 --- /dev/null +++ b/apps/daemon/internal/gateway/relay.go @@ -0,0 +1,144 @@ +package gateway + +import ( + "context" + "crypto/tls" + "errors" + "net" + "net/http" + "net/http/httputil" + "strconv" + "strings" + "time" + + "golang.org/x/net/http/httpguts" + + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxnet" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +const ( + // dialTimeout bounds a connection from the agent host. + dialTimeout = 30 * time.Second + // connectTimeout bounds the sandbox's resolution and dial for one + // Connect. + connectTimeout = 30 * time.Second +) + +// newTransport returns an upstream transport that relays requests as they +// are: no proxy from the environment, no added compression and no redirects, +// which http.Transport never follows. +func newTransport(tlsConfig *tls.Config, dial func(context.Context, string, string) (net.Conn, error)) *http.Transport { + if tlsConfig != nil { + tlsConfig = tlsConfig.Clone() + } + return &http.Transport{ + DialContext: dial, + TLSClientConfig: tlsConfig, + ForceAttemptHTTP2: true, + DisableCompression: true, + TLSHandshakeTimeout: 10 * time.Second, + ExpectContinueTimeout: time.Second, + MaxIdleConns: 64, + IdleConnTimeout: 90 * time.Second, + } +} + +// sandboxDialer connects through a new Network stream for each connection, so +// the sandbox resolves the name and the connection has sandbox origin. +func sandboxDialer(open func(context.Context) (sandboxlink.Stream, error)) func(context.Context, string, string) (net.Conn, error) { + return func(ctx context.Context, network, addr string) (net.Conn, error) { + if network != "tcp" && network != "tcp4" && network != "tcp6" { + return nil, &sandboxnet.Error{Code: sandboxnet.CodeUnsupportedNetwork, Effect: sandboxwire.EffectNone} + } + host, port, err := splitHostPort(addr) + if err != nil { + return nil, err + } + return connectSandbox(ctx, open, host, port) + } +} + +func connectSandbox(ctx context.Context, open func(context.Context) (sandboxlink.Stream, error), host string, port uint16) (*sandboxnet.Conn, error) { + s, err := open(ctx) + if err != nil { + return nil, err + } + return sandboxnet.Connect(ctx, s, host, port, connectTimeout) +} + +// splitHostPort splits an authority into an unbracketed host and a nonzero +// port. +func splitHostPort(addr string) (string, uint16, error) { + host, p, err := net.SplitHostPort(addr) + if err != nil { + return "", 0, &sandboxnet.Error{Code: sandboxnet.CodeInvalidArgument, Effect: sandboxwire.EffectNone, Cause: err} + } + port, err := strconv.ParseUint(p, 10, 16) + if err != nil || port == 0 || host == "" { + return "", 0, &sandboxnet.Error{Code: sandboxnet.CodeInvalidArgument, Effect: sandboxwire.EffectNone} + } + return host, uint16(port), nil +} + +// reverseProxy relays one request through t after rewrite. It flushes every +// write, so event streams pass as they arrive, and relays protocol upgrades. +func reverseProxy(t http.RoundTripper, rewrite func(*httputil.ProxyRequest)) *httputil.ReverseProxy { + return &httputil.ReverseProxy{Rewrite: rewrite, Transport: t, FlushInterval: -1, ErrorLog: quiet, ErrorHandler: relayFailed} +} + +// relayFailed answers a request that reached no upstream answer. The body +// says nothing about the upstream. +func relayFailed(w http.ResponseWriter, _ *http.Request, err error) { + status := statusOf(err) + http.Error(w, http.StatusText(status), status) +} + +// statusOf maps a failure to reach a destination to an HTTP status. +func statusOf(err error) int { + var e *sandboxnet.Error + if errors.As(err, &e) { + switch e.Code { + case sandboxnet.CodeInvalidArgument: + return http.StatusBadRequest + case sandboxnet.CodeDenied: + return http.StatusForbidden + case sandboxnet.CodeTimedOut: + return http.StatusGatewayTimeout + } + } + var ne net.Error + if errors.As(err, &ne) && ne.Timeout() { + return http.StatusGatewayTimeout + } + return http.StatusBadGateway +} + +// stripCredentials removes every value of each header in +// modelprovider.StrippedHeaders, matching names case-insensitively. +func stripCredentials(h http.Header) { + for name := range h { + for _, s := range modelprovider.StrippedHeaders { + if strings.EqualFold(name, s) { + delete(h, name) + break + } + } + } +} + +// wantsUpgrade reports whether r asks for a protocol upgrade, as +// httputil.ReverseProxy recognizes one. +func wantsUpgrade(r *http.Request) bool { + return httpguts.HeaderValuesContainsToken(r.Header["Connection"], "upgrade") +} + +// requestTarget splits an origin-form request target into its escaped path +// and its query exactly as the Harness sent them. ok is false for any other +// form. +func requestTarget(r *http.Request) (path, query string, hasQuery, ok bool) { + path, query, hasQuery = strings.Cut(r.RequestURI, "?") + return path, query, hasQuery, strings.HasPrefix(path, "/") +} diff --git a/apps/daemon/internal/gateway/view_linux_test.go b/apps/daemon/internal/gateway/view_linux_test.go new file mode 100644 index 00000000..1fc830e8 --- /dev/null +++ b/apps/daemon/internal/gateway/view_linux_test.go @@ -0,0 +1,263 @@ +//go:build linux + +package gateway + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "os" + "path/filepath" + "strconv" + "testing" + "time" + + gofs "github.com/hanwen/go-fuse/v2/fs" + "github.com/hanwen/go-fuse/v2/fuse" + "golang.org/x/sys/unix" + + "github.com/MiniMax-AI/OpenAgentCore/apps/daemon/internal/sessionview" + "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" +) + +// The view test needs root with CAP_SYS_ADMIN and CAP_NET_ADMIN, /dev/fuse and no AppArmor confinement. Run it in a throwaway container: +// +// CGO_ENABLED=0 go test -c -o /tmp/gateway.test ./apps/daemon/internal/gateway +// docker run --rm --cap-add SYS_ADMIN --cap-add NET_ADMIN --device /dev/fuse --security-opt apparmor=unconfined \ +// -e OAC_TEST_SESSIONVIEW=1 -v /tmp/gateway.test:/t.test:ro debian:bookworm-slim /t.test -test.v +const ( + gateEnv = "OAC_TEST_SESSIONVIEW" + harnessEnv = "OAC_GATEWAY_HARNESS" + endpointsEnv = "OAC_GATEWAY_ENDPOINTS" + externalEnv = "OAC_GATEWAY_EXTERNAL" + viewID = 1000 +) + +// The test binary is also the Harness inside the view. +func TestMain(m *testing.M) { + sessionview.Init() + if os.Getenv(harnessEnv) != "" { + os.Exit(runHarness()) + } + os.Exit(m.Run()) +} + +func TestListenersExistOnlyInTheSession(t *testing.T) { + if os.Getenv(gateEnv) != "1" { + t.Skipf("set %s=1 and run the test binary as root in a privileged container; see the comment at the top of this file", gateEnv) + } + if err := sessionview.Probe(); err != nil { + t.Fatalf("Probe: %v", err) + } + // A server in the host namespace stands in for a service-origin MCP + // server. The Harness reaches it only through its listener. + ln, err := net.Listen("tcp4", "0.0.0.0:0") + if err != nil { + t.Fatal(err) + } + server := &http.Server{Handler: http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { io.WriteString(w, "tools") })} + go server.Serve(ln) + defer server.Close() + port := strconv.Itoa(ln.Addr().(*net.TCPAddr).Port) + + cfg := Config{ + Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: "https://127.0.0.1:1", APIKey: upstreamKey}}}, + MCP: []proto.MCPHTTPServer{{ConnectionOrigin: "service", ServerLabel: "tools", ServerURL: "http://127.0.0.1:" + port + "/mcp"}}, + Prompt: proto.PromptRequestPayload{DisableExecutionEnvironment: true}, + OpenNetwork: startSandbox(t).open, + } + // The Harness's environment is built before the view exists. + eps, err := Plan(cfg) + if err != nil { + t.Fatal(err) + } + encoded, _ := json.Marshal(eps) + + world, harness := t.TempDir(), t.TempDir() + for _, d := range []string{".oac/harness", ".oac/bin", "proc", "dev"} { + if err := os.MkdirAll(filepath.Join(world, d), 0o755); err != nil { + t.Fatal(err) + } + } + if err := os.Chmod(harness, 0o755); err != nil { + t.Fatal(err) + } + copyExecutable(t, filepath.Join(harness, "harness")) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + v, err := sessionview.Start(context.Background(), sessionview.Spec{ + World: (&loopbackWorld{dir: world}).serve, + Private: []sessionview.PrivateDir{{Name: "harness", HostDir: harness, Exec: true}}, + Process: sessionview.Process{ + Path: "/.oac/harness/harness", Args: []string{"harness"}, Dir: "/", UID: viewID, GID: viewID, Stderr: os.Stderr, + Env: []string{harnessEnv + "=1", endpointsEnv + "=" + string(encoded), externalEnv + "=" + net.JoinHostPort(hostAddress(t), port)}, + }, + Network: sessionview.Network{Setup: func(netns *os.File) error { + _, err := Start(ctx, SessionNetwork{Namespace: netns}, cfg) + return err + }}, + }) + if err != nil { + t.Fatalf("Start: %v", err) + } + defer v.Close() + out, err := io.ReadAll(v.Stdout()) + if err != nil { + t.Fatal(err) + } + if exit, err := v.Wait(); err != nil || exit != (sessionview.Exit{}) { + t.Fatalf("Wait = %+v, %v; output %s", exit, err, out) + } + var report map[string]string + if err := json.Unmarshal(out, &report); err != nil { + t.Fatalf("Harness output %q: %v", out, err) + } + for _, name := range harnessChecks { + if msg, ok := report[name]; !ok || msg != "" { + t.Errorf("%s: %q", name, msg) + } + } + + // The listeners still serve the Session's namespace; the host's has none. + for _, u := range []string{eps.Proxy, eps.Models["main"], eps.MCP["tools"]} { + parsed, _ := url.Parse(u) + if c, err := net.DialTimeout("tcp", parsed.Host, time.Second); err == nil { + c.Close() + t.Errorf("%s is reachable from the host namespace", parsed.Host) + } + } +} + +var harnessChecks = []string{"model listener", "MCP listener", "no direct route", "route through the proxy"} + +// runHarness checks the network from inside the view and prints a report. +func runHarness() int { + external := os.Getenv(externalEnv) + var eps Endpoints + if err := json.Unmarshal([]byte(os.Getenv(endpointsEnv)), &eps); err != nil { + fmt.Fprintln(os.Stderr, err) + return 2 + } + proxyURL, err := url.Parse(eps.Proxy) + if err != nil { + fmt.Fprintln(os.Stderr, err) + return 2 + } + direct := &http.Client{Transport: &http.Transport{Proxy: nil}, Timeout: wait} + proxied := &http.Client{Transport: &http.Transport{Proxy: http.ProxyURL(proxyURL)}, Timeout: wait} + answers := func(client *http.Client, u string, status int, body string) error { + resp, err := client.Get(u) + if err != nil { + return err + } + defer resp.Body.Close() + got, _ := io.ReadAll(resp.Body) + if resp.StatusCode != status || (body != "" && string(got) != body) { + return fmt.Errorf("%s answered %d %q", u, resp.StatusCode, got) + } + return nil + } + checks := map[string]func() error{ + // An undeclared route is answered by the listener itself. + "model listener": func() error { return answers(direct, eps.Models["main"]+"/", http.StatusNotFound, "") }, + "MCP listener": func() error { return answers(direct, eps.MCP["tools"], http.StatusOK, "tools") }, + "no direct route": func() error { + c, err := net.DialTimeout("tcp", external, 2*time.Second) + if err == nil { + c.Close() + return errors.New("connected without the proxy") + } + if !errors.Is(err, unix.ENETUNREACH) { + return fmt.Errorf("dial: %v, want ENETUNREACH", err) + } + return nil + }, + // The same address answers through the proxy, which connects from the sandbox. + "route through the proxy": func() error { return answers(proxied, "http://"+external+"/", http.StatusOK, "tools") }, + } + report := map[string]string{} + for _, name := range harnessChecks { + report[name] = "" + if err := checks[name](); err != nil { + report[name] = err.Error() + } + } + json.NewEncoder(os.Stdout).Encode(report) + return 0 +} + +// hostAddress returns a non-loopback IPv4 address of the host namespace. +func hostAddress(t *testing.T) string { + t.Helper() + addrs, err := net.InterfaceAddrs() + if err != nil { + t.Fatal(err) + } + for _, a := range addrs { + if ip, ok := a.(*net.IPNet); ok && ip.IP.To4() != nil && !ip.IP.IsLoopback() { + return ip.IP.String() + } + } + t.Fatal("the host namespace has no non-loopback IPv4 address; run the container with a network") + return "" +} + +func copyExecutable(t *testing.T, dst string) { + t.Helper() + self, err := os.Executable() + if err != nil { + t.Fatal(err) + } + data, err := os.ReadFile(self) + if err != nil { + t.Fatal(err) + } + if err := os.WriteFile(dst, data, 0o755); err != nil { + t.Fatal(err) + } +} + +// loopbackWorld serves a directory as the view's world. +type loopbackWorld struct { + dir string + served chan struct{} +} + +func (w *loopbackWorld) serve(dev *os.File, _ sessionview.WorldMount) (sessionview.WorldServer, error) { + fd, err := unix.Dup(int(dev.Fd())) + if err != nil { + return nil, err + } + root, err := gofs.NewLoopbackRoot(w.dir) + if err != nil { + unix.Close(fd) + return nil, err + } + srv, err := fuse.NewServer(gofs.NewNodeFS(root, &gofs.Options{}), fmt.Sprintf("/dev/fd/%d", fd), &fuse.MountOptions{}) + if err != nil { + unix.Close(fd) + return nil, err + } + w.served = make(chan struct{}) + go func() { + srv.Serve() + close(w.served) + }() + return w, nil +} + +func (w *loopbackWorld) Stop() error { + select { + case <-w.served: + return nil + case <-time.After(10 * time.Second): + return errors.New("world still serving 10s after the view ended") + } +} From 3f8061cd894a813044fbd1b9c0dd98990cf0cf63 Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Thu, 1 Oct 2026 01:35:30 +0000 Subject: [PATCH 2/3] Withhold credentials from responses and end relays with the Session - Model and MCP relays remove every response header and trailer value that contains the injected key or bearer, informational responses included. Bodies pass unchanged. model-execution.md declares the rule. - Config.TLS becomes Config.RootCAs: upstream TLS always verifies the destination's hostname. - The listeners record every accepted connection, so the end of the Session aborts hijacked tunnels and upgraded connections that http.Server.Close leaves open, including a relay blocked on a Harness that does not read. - Every upstream dial ends with the Session, and a sandbox open plus Connect has a local deadline of the connect timeout plus a margin. - The forward proxy restores the query that the reverse proxy cleans of semicolon parameters. --- apps/daemon/internal/gateway/conns.go | 66 +++++++++++ apps/daemon/internal/gateway/gateway.go | 39 ++++--- apps/daemon/internal/gateway/gateway_test.go | 110 ++++++++++++++++--- apps/daemon/internal/gateway/listen_linux.go | 9 +- apps/daemon/internal/gateway/listen_other.go | 2 +- apps/daemon/internal/gateway/mcp.go | 6 +- apps/daemon/internal/gateway/mcp_test.go | 2 +- apps/daemon/internal/gateway/model.go | 6 +- apps/daemon/internal/gateway/model_test.go | 51 ++++++++- apps/daemon/internal/gateway/proxy.go | 12 +- apps/daemon/internal/gateway/proxy_test.go | 11 +- apps/daemon/internal/gateway/relay.go | 104 ++++++++++++++++-- contracts/agents-api/model-execution.md | 3 +- 13 files changed, 363 insertions(+), 58 deletions(-) create mode 100644 apps/daemon/internal/gateway/conns.go diff --git a/apps/daemon/internal/gateway/conns.go b/apps/daemon/internal/gateway/conns.go new file mode 100644 index 00000000..ef6a555c --- /dev/null +++ b/apps/daemon/internal/gateway/conns.go @@ -0,0 +1,66 @@ +package gateway + +import ( + "net" + "sync" +) + +// sessionConns are the Harness's open connections to the gateway, hijacked +// and upgraded ones included. http.Server.Close does not close a hijacked +// connection, so the end of the Session aborts them here. +type sessionConns struct { + mu sync.Mutex + open map[*sessionConn]struct{} + ended bool +} + +// abort resets every open connection and refuses later ones. +func (s *sessionConns) abort() { + s.mu.Lock() + open := s.open + s.open, s.ended = nil, true + s.mu.Unlock() + for c := range open { + c.SetLinger(0) + c.TCPConn.Close() + } +} + +// sessionListener records each connection it accepts in conns. +type sessionListener struct { + *net.TCPListener + conns *sessionConns +} + +func (l sessionListener) Accept() (net.Conn, error) { + tc, err := l.AcceptTCP() + if err != nil { + return nil, err + } + c := &sessionConn{TCPConn: tc, conns: l.conns} + l.conns.mu.Lock() + ended := l.conns.ended + if !ended { + l.conns.open[c] = struct{}{} + } + l.conns.mu.Unlock() + if ended { + tc.SetLinger(0) + tc.Close() + return nil, net.ErrClosed + } + return c, nil +} + +// sessionConn is an accepted connection that leaves conns when it closes. +type sessionConn struct { + *net.TCPConn + conns *sessionConns +} + +func (c *sessionConn) Close() error { + c.conns.mu.Lock() + delete(c.conns.open, c) + c.conns.mu.Unlock() + return c.TCPConn.Close() +} diff --git a/apps/daemon/internal/gateway/gateway.go b/apps/daemon/internal/gateway/gateway.go index aa926f5c..13da74a8 100644 --- a/apps/daemon/internal/gateway/gateway.go +++ b/apps/daemon/internal/gateway/gateway.go @@ -13,12 +13,15 @@ // the gateway does the TLS either way. The generic proxy carries HTTP CONNECT // tunnels and plain-HTTP forward requests, and connects only through the // sandbox's Network service. Redirects reach the Harness unchanged and are -// never followed. The gateway logs nothing. +// never followed. Response headers and trailers that carry an injected +// credential are withheld; bodies pass unchanged. The end of the Session +// closes every connection, tunnels and upgraded ones included. The gateway +// logs nothing. package gateway import ( "context" - "crypto/tls" + "crypto/x509" "errors" "fmt" "io" @@ -51,8 +54,9 @@ type Config struct { // network: the generic proxy then refuses every request, and an // environment-origin binding is invalid. OpenNetwork func(context.Context) (sandboxlink.Stream, error) - // TLS configures the gateway's upstream TLS. Nil uses the system roots. - TLS *tls.Config + // RootCAs are the roots the gateway trusts for upstream TLS. Nil means + // the system roots. The server name is always the destination's hostname. + RootCAs *x509.CertPool } // Model is one frozen model upstream. @@ -107,7 +111,7 @@ var ( // Plan validates cfg and returns the Endpoints that Start serves for it, so // the Harness's environment can be built before the view starts. func Plan(cfg Config) (Endpoints, error) { - g, err := build(cfg) + g, err := build(context.Background(), cfg) if err != nil { return Endpoints{}, err } @@ -120,7 +124,7 @@ func Plan(cfg Config) (Endpoints, error) { // hook; nothing listens outside the namespace. The end of ctx closes the // listeners and every connection. func Start(ctx context.Context, n SessionNetwork, cfg Config) (Endpoints, error) { - g, err := build(cfg) + g, err := build(ctx, cfg) if err != nil { return Endpoints{}, err } @@ -170,16 +174,17 @@ func invalid(format string, args ...any) error { return fmt.Errorf("%w: "+format, append([]any{ErrInvalidConfig}, args...)...) } -// build validates cfg and makes each listener's handler. -func build(cfg Config) (*gateway, error) { +// build validates cfg and makes each listener's handler. Every upstream dial +// ends with session. +func build(session context.Context, cfg Config) (*gateway, error) { if n := 1 + len(cfg.Models) + len(cfg.MCP); n > maxListeners { return nil, invalid("%d listeners, at most %d", n, maxListeners) } - host := newTransport(cfg.TLS, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext) + host := newTransport(cfg.RootCAs, sessionDial(session, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext)) g := &gateway{transports: []*http.Transport{host}} var sandbox *http.Transport if cfg.OpenNetwork != nil { - sandbox = newTransport(cfg.TLS, sandboxDialer(cfg.OpenNetwork)) + sandbox = newTransport(cfg.RootCAs, sessionDial(session, sandboxDialer(cfg.OpenNetwork))) g.transports = append(g.transports, sandbox) } g.listeners = append(g.listeners, listener{role: roleProxy, handler: newProxy(cfg.OpenNetwork, sandbox)}) @@ -242,20 +247,24 @@ func (g *gateway) endpoints(ports []int) Endpoints { // could carry more than the gateway chooses to reveal. var quiet = log.New(io.Discard, "", 0) -// serve serves each listener with its handler until ctx ends. -func (g *gateway) serve(ctx context.Context, lns []net.Listener) { +// serve serves each listener with its handler until ctx ends. The end of ctx +// cancels every request, which closes its upstream side, and aborts every +// connection the listeners accepted, which ends a relay blocked on a Harness +// that does not read. +func (g *gateway) serve(ctx context.Context, lns []*net.TCPListener) { + conns := &sessionConns{open: map[*sessionConn]struct{}{}} for i, ln := range lns { srv := &http.Server{ Handler: g.listeners[i].handler, ReadHeaderTimeout: 30 * time.Second, ErrorLog: quiet, - // Requests, hijacked tunnels and upgraded connections end with ctx. - BaseContext: func(net.Listener) context.Context { return ctx }, + BaseContext: func(net.Listener) context.Context { return ctx }, } - go srv.Serve(ln) + go srv.Serve(sessionListener{TCPListener: ln, conns: conns}) context.AfterFunc(ctx, func() { srv.Close() }) } context.AfterFunc(ctx, func() { + conns.abort() for _, t := range g.transports { t.CloseIdleConnections() } diff --git a/apps/daemon/internal/gateway/gateway_test.go b/apps/daemon/internal/gateway/gateway_test.go index c09767d3..66fcc7d1 100644 --- a/apps/daemon/internal/gateway/gateway_test.go +++ b/apps/daemon/internal/gateway/gateway_test.go @@ -1,17 +1,22 @@ package gateway import ( + "bufio" "context" - "crypto/tls" "crypto/x509" + "fmt" + "io" "net" "net/http" "net/http/httptest" "net/netip" + "net/url" + "sync" "sync/atomic" "testing" "time" + "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink" "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink/relay" "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxlink/sandboxlinktest" @@ -21,33 +26,50 @@ import ( const wait = 5 * time.Second -// serveOnLoopback serves cfg on host loopback listeners at free ports, as -// Start serves it in the Session's namespace, until the test ends. -func serveOnLoopback(t *testing.T, cfg Config) Endpoints { +// loopback is a gateway served on host loopback listeners at free ports, as +// Start serves it in the Session's namespace. +type loopback struct { + Endpoints + // end ends the Session. + end context.CancelFunc + // handlers counts the requests being handled. + handlers sync.WaitGroup +} + +// serveOnLoopback serves cfg until the test ends. +func serveOnLoopback(t *testing.T, cfg Config) *loopback { t.Helper() - g, err := build(cfg) + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + g, err := build(ctx, cfg) if err != nil { t.Fatal(err) } - lns := make([]net.Listener, len(g.listeners)) + l := &loopback{end: cancel} + lns := make([]*net.TCPListener, len(g.listeners)) ports := make([]int, len(lns)) for i := range lns { - if lns[i], err = net.Listen("tcp4", "127.0.0.1:0"); err != nil { + if lns[i], err = net.ListenTCP("tcp4", &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1)}); err != nil { t.Fatal(err) } ports[i] = lns[i].Addr().(*net.TCPAddr).Port + h := g.listeners[i].handler + g.listeners[i].handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + l.handlers.Add(1) + defer l.handlers.Done() + h.ServeHTTP(w, r) + }) } - ctx, cancel := context.WithCancel(context.Background()) - t.Cleanup(cancel) g.serve(ctx, lns) - return g.endpoints(ports) + l.Endpoints = g.endpoints(ports) + return l } -// trust returns a TLS configuration that trusts srv. -func trust(srv *httptest.Server) *tls.Config { +// trust returns roots that trust srv. +func trust(srv *httptest.Server) *x509.CertPool { roots := x509.NewCertPool() roots.AddCert(srv.Certificate()) - return &tls.Config{RootCAs: roots} + return roots } // noRedirects is a Harness-side client that shows each answer as it is. @@ -134,3 +156,65 @@ func startSandbox(t *testing.T) *sandbox { } return sb } + +func TestSessionEndEndsBlockedRelays(t *testing.T) { + // The upstream upgrades the connection and writes until the Harness's + // side is full, then reads until the gateway closes its side. + flooded, closed := make(chan struct{}), make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + c, rw, err := http.NewResponseController(w).Hijack() + if err != nil { + return + } + defer c.Close() + rw.WriteString("HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: test\r\n\r\n") + rw.Flush() + chunk := make([]byte, 64<<10) + for { + c.SetWriteDeadline(time.Now().Add(time.Second)) + if _, err := c.Write(chunk); err != nil { + break + } + } + close(flooded) + c.SetWriteDeadline(time.Time{}) + io.Copy(io.Discard, c) + close(closed) + })) + defer srv.Close() + gw := serveOnLoopback(t, Config{ + MCP: []proto.MCPHTTPServer{{ConnectionOrigin: "service", ServerLabel: "tools", ServerURL: srv.URL + "/mcp"}}, + Prompt: proto.PromptRequestPayload{DisableExecutionEnvironment: true}, + }) + + u, _ := url.Parse(gw.MCP["tools"]) + harness, err := net.Dial("tcp", u.Host) + if err != nil { + t.Fatal(err) + } + defer harness.Close() + fmt.Fprintf(harness, "GET /mcp HTTP/1.1\r\nHost: %s\r\nConnection: Upgrade\r\nUpgrade: test\r\n\r\n", u.Host) + if resp, err := http.ReadResponse(bufio.NewReader(harness), nil); err != nil || resp.StatusCode != http.StatusSwitchingProtocols { + t.Fatalf("upgrade: %v %v", resp, err) + } + // The Harness reads nothing more, so the gateway's write to it blocks. + select { + case <-flooded: + case <-time.After(2 * wait): + t.Fatal("the upstream never filled the Harness's side") + } + + gw.end() + returned := make(chan struct{}) + go func() { + gw.handlers.Wait() + close(returned) + }() + for what, done := range map[string]chan struct{}{"the relay": returned, "the upstream connection": closed} { + select { + case <-done: + case <-time.After(wait): + t.Fatalf("%s is still open after the Session ended", what) + } + } +} diff --git a/apps/daemon/internal/gateway/listen_linux.go b/apps/daemon/internal/gateway/listen_linux.go index 1f0a453c..24779783 100644 --- a/apps/daemon/internal/gateway/listen_linux.go +++ b/apps/daemon/internal/gateway/listen_linux.go @@ -7,7 +7,6 @@ import ( "net" "os" "runtime" - "strconv" "golang.org/x/sys/unix" ) @@ -16,9 +15,9 @@ import ( // namespace ns. A socket stays in the namespace it was created in, so the // listeners are opened on a thread that joined ns and then served from the // daemon's own threads. -func listen(ns *os.File, ports []int) ([]net.Listener, error) { +func listen(ns *os.File, ports []int) ([]*net.TCPListener, error) { type result struct { - lns []net.Listener + lns []*net.TCPListener err error } done := make(chan result, 1) @@ -30,9 +29,9 @@ func listen(ns *os.File, ports []int) ([]net.Listener, error) { done <- result{err: fmt.Errorf("%w: join: %w", ErrNetwork, err)} return } - var lns []net.Listener + var lns []*net.TCPListener for _, port := range ports { - ln, err := net.Listen("tcp4", net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + ln, err := net.ListenTCP("tcp4", &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: port}) if err != nil { for _, l := range lns { l.Close() diff --git a/apps/daemon/internal/gateway/listen_other.go b/apps/daemon/internal/gateway/listen_other.go index e326bcf6..db6a143f 100644 --- a/apps/daemon/internal/gateway/listen_other.go +++ b/apps/daemon/internal/gateway/listen_other.go @@ -8,4 +8,4 @@ import ( ) // listen needs Linux network namespaces. -func listen(*os.File, []int) ([]net.Listener, error) { return nil, ErrUnsupported } +func listen(*os.File, []int) ([]*net.TCPListener, error) { return nil, ErrUnsupported } diff --git a/apps/daemon/internal/gateway/mcp.go b/apps/daemon/internal/gateway/mcp.go index 4f03207d..7453129e 100644 --- a/apps/daemon/internal/gateway/mcp.go +++ b/apps/daemon/internal/gateway/mcp.go @@ -10,8 +10,9 @@ import ( ) // mcpRelay serves one MCP HTTP binding: it relays each request to the -// server's origin with the same path and query, and injects the binding's -// bearer token when it has one. +// server's origin with the same path and query. When the binding has a bearer +// token, it injects the token and withholds it from response headers and +// trailers. type mcpRelay struct { scheme string host string @@ -33,6 +34,7 @@ func newMCPRelay(s proto.MCPHTTPServer, t http.RoundTripper) (*mcpRelay, string, } token := *s.BearerToken m.token = &token + m.transport = withhold(t, token) } suffix := u.EscapedPath() if u.RawQuery != "" || u.ForceQuery { diff --git a/apps/daemon/internal/gateway/mcp_test.go b/apps/daemon/internal/gateway/mcp_test.go index f3358c0a..196d213e 100644 --- a/apps/daemon/internal/gateway/mcp_test.go +++ b/apps/daemon/internal/gateway/mcp_test.go @@ -32,7 +32,7 @@ func TestMCPBrokersBothOrigins(t *testing.T) { dials int32 }{{"service", service, 0}, {"environment", environment, 1}} { before := sb.dials.Load() - eps := serveOnLoopback(t, Config{MCP: []proto.MCPHTTPServer{binding(c.origin)}, Prompt: c.prompt, OpenNetwork: sb.open, TLS: trust(srv)}) + eps := serveOnLoopback(t, Config{MCP: []proto.MCPHTTPServer{binding(c.origin)}, Prompt: c.prompt, OpenNetwork: sb.open, RootCAs: trust(srv)}) if !strings.HasPrefix(eps.MCP["tools"], "http://127.0.0.1:") || !strings.HasSuffix(eps.MCP["tools"], "/mcp") { t.Fatalf("%s: Harness URL %q", c.origin, eps.MCP["tools"]) } diff --git a/apps/daemon/internal/gateway/model.go b/apps/daemon/internal/gateway/model.go index b768a58e..f9a37872 100644 --- a/apps/daemon/internal/gateway/model.go +++ b/apps/daemon/internal/gateway/model.go @@ -10,8 +10,8 @@ import ( ) // modelRelay serves one frozen model upstream. It relays only the routes the -// protocol declares, injects the declared credential and keeps the base URL's -// path. +// protocol declares, injects the declared credential, keeps the base URL's +// path and withholds the key from response headers and trailers. type modelRelay struct { protocol modelprovider.Protocol scheme string @@ -35,7 +35,7 @@ func newModelRelay(p modelprovider.Provider, t http.RoundTripper) (*modelRelay, return nil, err } m := &modelRelay{protocol: p.Protocol, scheme: base.Scheme, host: base.Host, targets: map[string]*url.URL{}, - credential: credential, key: p.APIKey, transport: t} + credential: credential, key: p.APIKey, transport: withhold(t, p.APIKey)} for _, route := range modelprovider.Routes(p.Protocol) { raw := modelprovider.UpstreamPath(base.EscapedPath(), route.Path) path, err := url.PathUnescape(raw) diff --git a/apps/daemon/internal/gateway/model_test.go b/apps/daemon/internal/gateway/model_test.go index 2aba3d3e..9a06eb86 100644 --- a/apps/daemon/internal/gateway/model_test.go +++ b/apps/daemon/internal/gateway/model_test.go @@ -2,9 +2,12 @@ package gateway import ( "bufio" + "context" "io" "net/http" "net/http/httptest" + "net/http/httptrace" + "net/textproto" "strings" "sync" "sync/atomic" @@ -22,8 +25,8 @@ func startModel(t *testing.T, handler http.HandlerFunc) (string, *httptest.Serve srv := httptest.NewTLSServer(handler) t.Cleanup(srv.Close) eps := serveOnLoopback(t, Config{ - Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: srv.URL + "/anthropic", APIKey: upstreamKey}}}, - TLS: trust(srv), + Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: srv.URL + "/anthropic", APIKey: upstreamKey}}}, + RootCAs: trust(srv), }) return eps.Models["main"], srv } @@ -71,6 +74,50 @@ func TestModelInjectsTheKeyAndNeverThePlaceholder(t *testing.T) { } } +func TestModelWithholdsTheKeyFromResponseHeaders(t *testing.T) { + // The upstream echoes the key it received in an informational response, + // in headers and in a trailer. + base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + key := r.Header.Get("X-Api-Key") + w.Header().Set("Link", "; rel=preload") + w.WriteHeader(http.StatusEarlyHints) + w.Header().Del("Link") + w.Header().Set("Trailer", "X-Echo, X-Done") + w.Header().Set("X-Echo", "key="+key) + w.Header().Set("Location", "https://elsewhere.invalid/?key="+key) + w.Header().Set("X-Native", "kept") + io.WriteString(w, "answer") + w.Header().Set("X-Echo", key) + w.Header().Set("X-Done", "yes") + }) + var informational []http.Header + ctx := httptrace.WithClientTrace(context.Background(), &httptrace.ClientTrace{ + Got1xxResponse: func(_ int, h textproto.MIMEHeader) error { + informational = append(informational, http.Header(h).Clone()) + return nil + }, + }) + req, _ := http.NewRequestWithContext(ctx, "POST", base+"/v1/messages", strings.NewReader("{}")) + resp, err := noRedirects.Do(req) + if err != nil { + t.Fatal(err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + if resp.StatusCode != 200 || string(body) != "answer" || resp.Header.Get("X-Native") != "kept" || resp.Trailer.Get("X-Done") != "yes" || len(informational) != 1 { + t.Fatalf("answer %d %q, header %v, trailer %v, %d informational", resp.StatusCode, body, resp.Header, resp.Trailer, len(informational)) + } + for _, h := range append(informational, resp.Header, resp.Trailer) { + for name, values := range h { + for _, v := range values { + if strings.Contains(v, upstreamKey) { + t.Errorf("the Harness received the key in %s", name) + } + } + } + } +} + func TestModelRejectsUndeclaredRequests(t *testing.T) { var hits atomic.Int32 base, _ := startModel(t, func(http.ResponseWriter, *http.Request) { hits.Add(1) }) diff --git a/apps/daemon/internal/gateway/proxy.go b/apps/daemon/internal/gateway/proxy.go index 6cde3430..05eec8fb 100644 --- a/apps/daemon/internal/gateway/proxy.go +++ b/apps/daemon/internal/gateway/proxy.go @@ -25,8 +25,12 @@ func newProxy(open func(context.Context) (sandboxlink.Stream, error), sandbox *h p := &proxy{open: open} if sandbox != nil { // The request is relayed as the Harness addressed it; the reverse - // proxy drops hop-by-hop headers, Proxy-Authorization among them. - p.forward = reverseProxy(sandbox, func(*httputil.ProxyRequest) {}) + // proxy drops hop-by-hop headers, Proxy-Authorization among them. It + // also drops query parameters it cannot parse, so the query is + // restored as sent. + p.forward = reverseProxy(sandbox, func(pr *httputil.ProxyRequest) { + pr.Out.URL.RawQuery = pr.In.URL.RawQuery + }) } return p } @@ -83,8 +87,8 @@ func splice(ctx context.Context, client net.Conn, pending []byte, remote *sandbo abort := func() { end.Do(func() { remote.Reset() - if tcp, ok := client.(*net.TCPConn); ok { - tcp.SetLinger(0) + if l, ok := client.(interface{ SetLinger(int) error }); ok { + l.SetLinger(0) } client.Close() }) diff --git a/apps/daemon/internal/gateway/proxy_test.go b/apps/daemon/internal/gateway/proxy_test.go index 1be49e3d..c7d9c85a 100644 --- a/apps/daemon/internal/gateway/proxy_test.go +++ b/apps/daemon/internal/gateway/proxy_test.go @@ -1,6 +1,7 @@ package gateway import ( + "crypto/tls" "io" "net/http" "net/http/httptest" @@ -11,22 +12,24 @@ import ( func TestProxyConnectsOnlyThroughTheSandbox(t *testing.T) { sb := startSandbox(t) eps := serveOnLoopback(t, Config{OpenNetwork: sb.open}) - hello := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "hello "+r.URL.Path) }) + hello := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "hello "+r.RequestURI) }) secure := httptest.NewTLSServer(hello) defer secure.Close() plain := httptest.NewServer(hello) defer plain.Close() proxyURL, _ := url.Parse(eps.Proxy) - client := &http.Client{Transport: &http.Transport{Proxy: http.ProxyURL(proxyURL), TLSClientConfig: trust(secure)}, Timeout: wait} - for i, target := range []string{secure.URL + "/via-connect", plain.URL + "/via-forward"} { + client := &http.Client{Transport: &http.Transport{Proxy: http.ProxyURL(proxyURL), TLSClientConfig: &tls.Config{RootCAs: trust(secure)}}, Timeout: wait} + // A forwarded query reaches the destination as sent, parameters with + // semicolons included. + for _, target := range []string{secure.URL + "/via-connect", plain.URL + "/via-forward?filter=a;b&keep=1"} { resp, err := client.Get(target) if err != nil { t.Fatalf("%s: %v", target, err) } body, _ := io.ReadAll(resp.Body) resp.Body.Close() - if want := "hello " + []string{"/via-connect", "/via-forward"}[i]; resp.StatusCode != 200 || string(body) != want { + if u, _ := url.Parse(target); resp.StatusCode != 200 || string(body) != "hello "+u.RequestURI() { t.Errorf("%s: %d %q", target, resp.StatusCode, body) } } diff --git a/apps/daemon/internal/gateway/relay.go b/apps/daemon/internal/gateway/relay.go index 566972e9..aac59842 100644 --- a/apps/daemon/internal/gateway/relay.go +++ b/apps/daemon/internal/gateway/relay.go @@ -3,10 +3,15 @@ package gateway import ( "context" "crypto/tls" + "crypto/x509" "errors" + "io" "net" "net/http" + "net/http/httptrace" "net/http/httputil" + "net/textproto" + "slices" "strconv" "strings" "time" @@ -25,18 +30,21 @@ const ( // connectTimeout bounds the sandbox's resolution and dial for one // Connect. connectTimeout = 30 * time.Second + // connectMargin is how long the gateway waits beyond connectTimeout for + // the stream to open and the sandbox's answer to arrive. + connectMargin = 5 * time.Second ) +type dialFunc func(ctx context.Context, network, addr string) (net.Conn, error) + // newTransport returns an upstream transport that relays requests as they // are: no proxy from the environment, no added compression and no redirects, -// which http.Transport never follows. -func newTransport(tlsConfig *tls.Config, dial func(context.Context, string, string) (net.Conn, error)) *http.Transport { - if tlsConfig != nil { - tlsConfig = tlsConfig.Clone() - } +// which http.Transport never follows. TLS trusts roots, or the system roots +// when roots is nil, and always verifies the destination's hostname. +func newTransport(roots *x509.CertPool, dial dialFunc) *http.Transport { return &http.Transport{ DialContext: dial, - TLSClientConfig: tlsConfig, + TLSClientConfig: &tls.Config{RootCAs: roots}, ForceAttemptHTTP2: true, DisableCompression: true, TLSHandshakeTimeout: 10 * time.Second, @@ -46,9 +54,22 @@ func newTransport(tlsConfig *tls.Config, dial func(context.Context, string, stri } } +// sessionDial binds each dial to the Session as well as to its own context. +// http.Transport detaches a dial from the request that started it, and +// CloseIdleConnections cancels only dials that no request waits for. +func sessionDial(session context.Context, dial dialFunc) dialFunc { + return func(ctx context.Context, network, addr string) (net.Conn, error) { + ctx, cancel := context.WithCancel(ctx) + defer cancel() + stop := context.AfterFunc(session, cancel) + defer stop() + return dial(ctx, network, addr) + } +} + // sandboxDialer connects through a new Network stream for each connection, so // the sandbox resolves the name and the connection has sandbox origin. -func sandboxDialer(open func(context.Context) (sandboxlink.Stream, error)) func(context.Context, string, string) (net.Conn, error) { +func sandboxDialer(open func(context.Context) (sandboxlink.Stream, error)) dialFunc { return func(ctx context.Context, network, addr string) (net.Conn, error) { if network != "tcp" && network != "tcp4" && network != "tcp6" { return nil, &sandboxnet.Error{Code: sandboxnet.CodeUnsupportedNetwork, Effect: sandboxwire.EffectNone} @@ -61,7 +82,12 @@ func sandboxDialer(open func(context.Context) (sandboxlink.Stream, error)) func( } } +// connectSandbox opens a Network stream and connects through it. A local +// deadline bounds both, so a sandbox that never answers cannot hold the +// connection open; the returned Conn outlives it. func connectSandbox(ctx context.Context, open func(context.Context) (sandboxlink.Stream, error), host string, port uint16) (*sandboxnet.Conn, error) { + ctx, cancel := context.WithTimeout(ctx, connectTimeout+connectMargin) + defer cancel() s, err := open(ctx) if err != nil { return nil, err @@ -116,6 +142,70 @@ func statusOf(err error) int { return http.StatusBadGateway } +// withhold returns t, or, when secret is set, a transport that keeps secret +// out of the response headers the Harness receives. +func withhold(t http.RoundTripper, secret string) http.RoundTripper { + if secret == "" { + return t + } + return withholding{next: t, secret: secret} +} + +// withholding removes every header and trailer value that contains secret +// from each response, informational ones included, so an upstream that echoes +// the injected credential in a header does not disclose it. Bodies pass +// unchanged. +type withholding struct { + next http.RoundTripper + secret string +} + +func (t withholding) RoundTrip(r *http.Request) (*http.Response, error) { + // The newest trace's hooks run first, so each informational response is + // cleaned before the reverse proxy relays it. + r = r.WithContext(httptrace.WithClientTrace(r.Context(), &httptrace.ClientTrace{ + Got1xxResponse: func(_ int, h textproto.MIMEHeader) error { + t.remove(http.Header(h)) + return nil + }, + })) + resp, err := t.next.RoundTrip(r) + if err != nil { + return nil, err + } + t.remove(resp.Header) + // An upgraded response's body is the connection itself, without + // trailers, and must stay one. + if resp.StatusCode != http.StatusSwitchingProtocols { + resp.Body = &withheldTrailers{ReadCloser: resp.Body, resp: resp, t: t} + } + return resp, nil +} + +// remove deletes each value that contains the secret. A name whose values +// are all removed stays with none, so an announced trailer stays announced. +func (t withholding) remove(h http.Header) { + for name, values := range h { + h[name] = slices.DeleteFunc(values, func(v string) bool { return strings.Contains(v, t.secret) }) + } +} + +// withheldTrailers cleans the response's trailers when the body ends: they +// have arrived by then, and the reverse proxy relays them only afterwards. +type withheldTrailers struct { + io.ReadCloser + resp *http.Response + t withholding +} + +func (b *withheldTrailers) Read(p []byte) (int, error) { + n, err := b.ReadCloser.Read(p) + if err == io.EOF { + b.t.remove(b.resp.Trailer) + } + return n, err +} + // stripCredentials removes every value of each header in // modelprovider.StrippedHeaders, matching names case-insensitively. func stripCredentials(h http.Header) { diff --git a/contracts/agents-api/model-execution.md b/contracts/agents-api/model-execution.md index fb523ae5..cc526838 100644 --- a/contracts/agents-api/model-execution.md +++ b/contracts/agents-api/model-execution.md @@ -94,8 +94,9 @@ At dispatch, Core sends the snapshot as one confidential provider bundle over th The Harness reaches its frozen upstream through a Session-local credential gateway on the agent host. The Harness's native base URL points at the gateway, and the Harness receives only a non-secret placeholder credential, never the key. -- The gateway relays only the declared native routes of the provider's protocol, with the request and response unchanged. An undeclared path or method, or an undeclared WebSocket upgrade, is rejected and never reaches the upstream. +- The gateway relays only the declared native routes of the provider's protocol, with the request and response unchanged apart from the credential rules below. An undeclared path or method, or an undeclared WebSocket upgrade, is rejected and never reaches the upstream. - It removes every value of each stripped header, matching names case-insensitively, then injects the upstream credential in the protocol's declared header. +- It removes every response header and trailer value that contains the key, in informational responses too. Response bodies pass unchanged, so an upstream that echoes the key in a body discloses it to the Harness. The gateway applies the same rule to the bearer token it injects for an HTTP MCP server. - It never follows a redirect with the credential. - It never converts between protocols. From 806000fbcda26b0ee68b94120f85cceb07f91c93 Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Thu, 1 Oct 2026 02:04:19 +0000 Subject: [PATCH 3/3] Keep relay credentials out of logs and close rejected upgrades Relay transports close each HTTP/1 connection after its response, so http.Transport never logs bytes an idle connection receives; HTTP/2 connections stay shared. Relays inject and withhold each credential in the form the header writer sends. ModifyResponse rejects a mismatched 101, and every upgraded upstream connection closes with its request. --- apps/daemon/internal/gateway/gateway.go | 13 +-- apps/daemon/internal/gateway/gateway_test.go | 38 +++++++++ apps/daemon/internal/gateway/mcp.go | 2 +- apps/daemon/internal/gateway/model.go | 3 +- apps/daemon/internal/gateway/model_test.go | 87 +++++++++++++++++--- apps/daemon/internal/gateway/relay.go | 65 ++++++++++++++- contracts/agents-api/model-execution.md | 2 +- 7 files changed, 189 insertions(+), 21 deletions(-) diff --git a/apps/daemon/internal/gateway/gateway.go b/apps/daemon/internal/gateway/gateway.go index 13da74a8..1f906fc8 100644 --- a/apps/daemon/internal/gateway/gateway.go +++ b/apps/daemon/internal/gateway/gateway.go @@ -180,14 +180,17 @@ func build(session context.Context, cfg Config) (*gateway, error) { if n := 1 + len(cfg.Models) + len(cfg.MCP); n > maxListeners { return nil, invalid("%d listeners, at most %d", n, maxListeners) } - host := newTransport(cfg.RootCAs, sessionDial(session, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext)) + host := relayTransport(cfg.RootCAs, sessionDial(session, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext)) g := &gateway{transports: []*http.Transport{host}} - var sandbox *http.Transport + // The proxy forwards on its own transport; environment-origin MCP relays + // use a relay transport. + var forward, sandbox *http.Transport if cfg.OpenNetwork != nil { - sandbox = newTransport(cfg.RootCAs, sessionDial(session, sandboxDialer(cfg.OpenNetwork))) - g.transports = append(g.transports, sandbox) + dial := sessionDial(session, sandboxDialer(cfg.OpenNetwork)) + forward, sandbox = newTransport(cfg.RootCAs, dial), relayTransport(cfg.RootCAs, dial) + g.transports = append(g.transports, forward, sandbox) } - g.listeners = append(g.listeners, listener{role: roleProxy, handler: newProxy(cfg.OpenNetwork, sandbox)}) + g.listeners = append(g.listeners, listener{role: roleProxy, handler: newProxy(cfg.OpenNetwork, forward)}) names := map[string]bool{} for _, m := range cfg.Models { diff --git a/apps/daemon/internal/gateway/gateway_test.go b/apps/daemon/internal/gateway/gateway_test.go index 66fcc7d1..b2f319c9 100644 --- a/apps/daemon/internal/gateway/gateway_test.go +++ b/apps/daemon/internal/gateway/gateway_test.go @@ -218,3 +218,41 @@ func TestSessionEndEndsBlockedRelays(t *testing.T) { } } } + +func TestRejectedUpgradeClosesTheUpstream(t *testing.T) { + // The upstream switches to a protocol the request did not ask for. + closed := make(chan struct{}) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + c, rw, err := http.NewResponseController(w).Hijack() + if err != nil { + return + } + defer c.Close() + rw.WriteString("HTTP/1.1 101 Switching Protocols\r\nConnection: Upgrade\r\nUpgrade: other\r\n\r\n") + rw.Flush() + io.Copy(io.Discard, c) + close(closed) + })) + defer srv.Close() + gw := serveOnLoopback(t, Config{ + MCP: []proto.MCPHTTPServer{{ConnectionOrigin: "service", ServerLabel: "tools", ServerURL: srv.URL + "/mcp"}}, + Prompt: proto.PromptRequestPayload{DisableExecutionEnvironment: true}, + }) + + req, _ := http.NewRequest("GET", gw.MCP["tools"], nil) + req.Header.Set("Connection", "Upgrade") + req.Header.Set("Upgrade", "test") + resp, err := noRedirects.Do(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != http.StatusBadGateway { + t.Fatalf("answer %d", resp.StatusCode) + } + select { + case <-closed: + case <-time.After(wait): + t.Fatal("the upstream connection is still open") + } +} diff --git a/apps/daemon/internal/gateway/mcp.go b/apps/daemon/internal/gateway/mcp.go index 7453129e..16064e10 100644 --- a/apps/daemon/internal/gateway/mcp.go +++ b/apps/daemon/internal/gateway/mcp.go @@ -32,7 +32,7 @@ func newMCPRelay(s proto.MCPHTTPServer, t http.RoundTripper) (*mcpRelay, string, if u.Scheme != "https" { return nil, "", errors.New("a bearer token needs an https server URL") } - token := *s.BearerToken + token := sentValue(*s.BearerToken) m.token = &token m.transport = withhold(t, token) } diff --git a/apps/daemon/internal/gateway/model.go b/apps/daemon/internal/gateway/model.go index f9a37872..718b4d43 100644 --- a/apps/daemon/internal/gateway/model.go +++ b/apps/daemon/internal/gateway/model.go @@ -34,8 +34,9 @@ func newModelRelay(p modelprovider.Provider, t http.RoundTripper) (*modelRelay, if err != nil { return nil, err } + key := sentValue(p.APIKey) m := &modelRelay{protocol: p.Protocol, scheme: base.Scheme, host: base.Host, targets: map[string]*url.URL{}, - credential: credential, key: p.APIKey, transport: withhold(t, p.APIKey)} + credential: credential, key: key, transport: withhold(t, key)} for _, route := range modelprovider.Routes(p.Protocol) { raw := modelprovider.UpstreamPath(base.EscapedPath(), route.Path) path, err := url.PathUnescape(raw) diff --git a/apps/daemon/internal/gateway/model_test.go b/apps/daemon/internal/gateway/model_test.go index 9a06eb86..db6faf24 100644 --- a/apps/daemon/internal/gateway/model_test.go +++ b/apps/daemon/internal/gateway/model_test.go @@ -4,6 +4,7 @@ import ( "bufio" "context" "io" + "log" "net/http" "net/http/httptest" "net/http/httptrace" @@ -19,13 +20,14 @@ import ( const upstreamKey = "sk-upstream-key" -// startModel serves one Anthropic upstream at the fake TLS upstream's /anthropic. -func startModel(t *testing.T, handler http.HandlerFunc) (string, *httptest.Server) { +// startModel serves one Anthropic upstream with key at the fake TLS +// upstream's /anthropic. +func startModel(t *testing.T, key string, handler http.HandlerFunc) (string, *httptest.Server) { t.Helper() srv := httptest.NewTLSServer(handler) t.Cleanup(srv.Close) eps := serveOnLoopback(t, Config{ - Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: srv.URL + "/anthropic", APIKey: upstreamKey}}}, + Models: []Model{{Name: "main", Provider: modelprovider.Provider{Protocol: modelprovider.Anthropic, BaseURL: srv.URL + "/anthropic", APIKey: key}}}, RootCAs: trust(srv), }) return eps.Models["main"], srv @@ -33,7 +35,7 @@ func startModel(t *testing.T, handler http.HandlerFunc) (string, *httptest.Serve func TestModelInjectsTheKeyAndNeverThePlaceholder(t *testing.T) { seen := make(chan *http.Request, 1) - base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + base, _ := startModel(t, upstreamKey, func(w http.ResponseWriter, r *http.Request) { body, _ := io.ReadAll(r.Body) r.Body = io.NopCloser(strings.NewReader(string(body))) seen <- r @@ -76,8 +78,15 @@ func TestModelInjectsTheKeyAndNeverThePlaceholder(t *testing.T) { func TestModelWithholdsTheKeyFromResponseHeaders(t *testing.T) { // The upstream echoes the key it received in an informational response, - // in headers and in a trailer. - base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + // in headers and in a trailer. A configured key with surrounding + // whitespace reaches the upstream trimmed. + for _, key := range []string{upstreamKey, upstreamKey + " \t"} { + withholdsTheKey(t, key) + } +} + +func withholdsTheKey(t *testing.T, key string) { + base, _ := startModel(t, key, func(w http.ResponseWriter, r *http.Request) { key := r.Header.Get("X-Api-Key") w.Header().Set("Link", "; rel=preload") w.WriteHeader(http.StatusEarlyHints) @@ -105,22 +114,78 @@ func TestModelWithholdsTheKeyFromResponseHeaders(t *testing.T) { body, _ := io.ReadAll(resp.Body) resp.Body.Close() if resp.StatusCode != 200 || string(body) != "answer" || resp.Header.Get("X-Native") != "kept" || resp.Trailer.Get("X-Done") != "yes" || len(informational) != 1 { - t.Fatalf("answer %d %q, header %v, trailer %v, %d informational", resp.StatusCode, body, resp.Header, resp.Trailer, len(informational)) + t.Fatalf("key %q: answer %d %q, header %v, trailer %v, %d informational", key, resp.StatusCode, body, resp.Header, resp.Trailer, len(informational)) } for _, h := range append(informational, resp.Header, resp.Trailer) { for name, values := range h { for _, v := range values { if strings.Contains(v, upstreamKey) { - t.Errorf("the Harness received the key in %s", name) + t.Errorf("key %q: the Harness received the key in %s", key, name) } } } } } +func TestModelKeepsTheKeyOutOfTheLog(t *testing.T) { + // The upstream follows an empty answer with bytes that echo the key, + // which http.Transport logs if the connection then sits idle. + closed := make(chan struct{}) + base, _ := startModel(t, upstreamKey, func(w http.ResponseWriter, r *http.Request) { + c, rw, err := http.NewResponseController(w).Hijack() + if err != nil { + return + } + defer c.Close() + rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\nHTTP/1.1 200 OK\r\nX-Echo: " + r.Header.Get("X-Api-Key") + "\r\n\r\n") + rw.Flush() + io.Copy(io.Discard, c) + close(closed) + }) + logged := &syncBuffer{} + previous := log.Writer() + log.SetOutput(logged) + t.Cleanup(func() { log.SetOutput(previous) }) + + resp, err := noRedirects.Post(base+"/v1/messages", "application/json", strings.NewReader("{}")) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if resp.StatusCode != 200 { + t.Fatalf("answer %d", resp.StatusCode) + } + select { + case <-closed: + case <-time.After(wait): + t.Fatal("the gateway kept the upstream connection") + } + if strings.Contains(logged.String(), upstreamKey) { + t.Fatal("the key reached the log") + } +} + +// syncBuffer collects what the global logger writes. +type syncBuffer struct { + mu sync.Mutex + b strings.Builder +} + +func (s *syncBuffer) Write(p []byte) (int, error) { + s.mu.Lock() + defer s.mu.Unlock() + return s.b.Write(p) +} + +func (s *syncBuffer) String() string { + s.mu.Lock() + defer s.mu.Unlock() + return s.b.String() +} + func TestModelRejectsUndeclaredRequests(t *testing.T) { var hits atomic.Int32 - base, _ := startModel(t, func(http.ResponseWriter, *http.Request) { hits.Add(1) }) + base, _ := startModel(t, upstreamKey, func(http.ResponseWriter, *http.Request) { hits.Add(1) }) cases := []struct { method, path string upgrade bool @@ -162,7 +227,7 @@ func TestModelStreamsEvents(t *testing.T) { released.Store(true) once.Do(func() { close(release) }) }) - base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + base, _ := startModel(t, upstreamKey, func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") io.WriteString(w, "data: one\n\n") w.(http.Flusher).Flush() @@ -188,7 +253,7 @@ func TestModelStreamsEvents(t *testing.T) { func TestModelReturnsRedirectsUnfollowed(t *testing.T) { var hits atomic.Int32 const location = "https://elsewhere.invalid/v1/messages?x=1" - base, _ := startModel(t, func(w http.ResponseWriter, r *http.Request) { + base, _ := startModel(t, upstreamKey, func(w http.ResponseWriter, r *http.Request) { hits.Add(1) http.Redirect(w, r, location, http.StatusTemporaryRedirect) }) diff --git a/apps/daemon/internal/gateway/relay.go b/apps/daemon/internal/gateway/relay.go index aac59842..d0bc5e2e 100644 --- a/apps/daemon/internal/gateway/relay.go +++ b/apps/daemon/internal/gateway/relay.go @@ -54,6 +54,18 @@ func newTransport(roots *x509.CertPool, dial dialFunc) *http.Transport { } } +// relayTransport returns a transport for model and MCP relays, which inject +// credentials. It closes each HTTP/1 connection after its response, so none +// is ever idle: http.Transport logs the bytes an idle HTTP/1 connection +// receives, and those can echo the credential. A negative MaxIdleConnsPerHost +// keeps HTTP/1 connections out of the idle pool; HTTP/2 connections stay +// shared in their own pool, and the HTTP/2 transport logs no received bytes. +func relayTransport(roots *x509.CertPool, dial dialFunc) *http.Transport { + t := newTransport(roots, dial) + t.MaxIdleConnsPerHost = -1 + return t +} + // sessionDial binds each dial to the Session as well as to its own context. // http.Transport detaches a dial from the request that started it, and // CloseIdleConnections cancels only dials that no request waits for. @@ -112,7 +124,47 @@ func splitHostPort(addr string) (string, uint16, error) { // reverseProxy relays one request through t after rewrite. It flushes every // write, so event streams pass as they arrive, and relays protocol upgrades. func reverseProxy(t http.RoundTripper, rewrite func(*httputil.ProxyRequest)) *httputil.ReverseProxy { - return &httputil.ReverseProxy{Rewrite: rewrite, Transport: t, FlushInterval: -1, ErrorLog: quiet, ErrorHandler: relayFailed} + return &httputil.ReverseProxy{Rewrite: rewrite, Transport: t, ModifyResponse: checkUpgrade, FlushInterval: -1, ErrorLog: quiet, ErrorHandler: relayFailed} +} + +var errUpgrade = errors.New("the upstream switched to a protocol the request did not ask for") + +// checkUpgrade accepts a 101 response only when the request asked for an +// upgrade, the response switches to that protocol and its body is the +// connection. ReverseProxy closes the body of a response checkUpgrade +// rejects, but not after its own upgrade errors. A 101 hands the connection +// to the relay, so the connection also closes when the request ends, which +// the end of the Session ends too. +func checkUpgrade(res *http.Response) error { + if res.StatusCode != http.StatusSwitchingProtocols { + return nil + } + body := res.Body + context.AfterFunc(res.Request.Context(), func() { body.Close() }) + want, got := upgradeType(res.Request.Header), upgradeType(res.Header) + if _, ok := body.(io.ReadWriteCloser); !ok || want == "" || !printable(want) || !printable(got) || !strings.EqualFold(want, got) { + return errUpgrade + } + return nil +} + +// upgradeType returns the protocol h upgrades to, as httputil.ReverseProxy +// reads it. +func upgradeType(h http.Header) string { + if !httpguts.HeaderValuesContainsToken(h["Connection"], "Upgrade") { + return "" + } + return h.Get("Upgrade") +} + +// printable reports whether s is printable ASCII. +func printable(s string) bool { + for i := 0; i < len(s); i++ { + if s[i] < ' ' || s[i] > '~' { + return false + } + } + return true } // relayFailed answers a request that reached no upstream answer. The body @@ -142,8 +194,17 @@ func statusOf(err error) int { return http.StatusBadGateway } +// sentValue returns v as the HTTP/1 header writer sends it: each CR and LF +// becomes a space, and surrounding whitespace goes. A relay injects a +// credential in this form and withholds it in this form, so the filter +// matches what the upstream received. +func sentValue(v string) string { + return textproto.TrimString(strings.NewReplacer("\n", " ", "\r", " ").Replace(v)) +} + // withhold returns t, or, when secret is set, a transport that keeps secret -// out of the response headers the Harness receives. +// out of the response headers the Harness receives. secret is the value as +// sent. func withhold(t http.RoundTripper, secret string) http.RoundTripper { if secret == "" { return t diff --git a/contracts/agents-api/model-execution.md b/contracts/agents-api/model-execution.md index cc526838..53a35c80 100644 --- a/contracts/agents-api/model-execution.md +++ b/contracts/agents-api/model-execution.md @@ -95,7 +95,7 @@ At dispatch, Core sends the snapshot as one confidential provider bundle over th The Harness reaches its frozen upstream through a Session-local credential gateway on the agent host. The Harness's native base URL points at the gateway, and the Harness receives only a non-secret placeholder credential, never the key. - The gateway relays only the declared native routes of the provider's protocol, with the request and response unchanged apart from the credential rules below. An undeclared path or method, or an undeclared WebSocket upgrade, is rejected and never reaches the upstream. -- It removes every value of each stripped header, matching names case-insensitively, then injects the upstream credential in the protocol's declared header. +- It removes every value of each stripped header, matching names case-insensitively, then injects the upstream credential in the protocol's declared header, with the key's surrounding whitespace removed. - It removes every response header and trailer value that contains the key, in informational responses too. Response bodies pass unchanged, so an upstream that echoes the key in a body discloses it to the Harness. The gateway applies the same rule to the bearer token it injects for an HTTP MCP server. - It never follows a redirect with the credential. - It never converts between protocols.