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 new file mode 100644 index 00000000..1f906fc8 --- /dev/null +++ b/apps/daemon/internal/gateway/gateway.go @@ -0,0 +1,275 @@ +// 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. 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/x509" + "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) + // 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. +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(context.Background(), 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(ctx, 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. 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 := relayTransport(cfg.RootCAs, sessionDial(session, (&net.Dialer{Timeout: dialTimeout, KeepAlive: 30 * time.Second}).DialContext)) + g := &gateway{transports: []*http.Transport{host}} + // The proxy forwards on its own transport; environment-origin MCP relays + // use a relay transport. + var forward, sandbox *http.Transport + if cfg.OpenNetwork != nil { + 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, forward)}) + + 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. 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, + BaseContext: func(net.Listener) context.Context { return ctx }, + } + 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 new file mode 100644 index 00000000..b2f319c9 --- /dev/null +++ b/apps/daemon/internal/gateway/gateway_test.go @@ -0,0 +1,258 @@ +package gateway + +import ( + "bufio" + "context" + "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" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxnet" + "github.com/MiniMax-AI/OpenAgentCore/internal/sandboxwire" +) + +const wait = 5 * time.Second + +// 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() + ctx, cancel := context.WithCancel(context.Background()) + t.Cleanup(cancel) + g, err := build(ctx, cfg) + if err != nil { + t.Fatal(err) + } + 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.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) + }) + } + g.serve(ctx, lns) + l.Endpoints = g.endpoints(ports) + return l +} + +// trust returns roots that trust srv. +func trust(srv *httptest.Server) *x509.CertPool { + roots := x509.NewCertPool() + roots.AddCert(srv.Certificate()) + return 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 +} + +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) + } + } +} + +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/listen_linux.go b/apps/daemon/internal/gateway/listen_linux.go new file mode 100644 index 00000000..24779783 --- /dev/null +++ b/apps/daemon/internal/gateway/listen_linux.go @@ -0,0 +1,48 @@ +//go:build linux + +package gateway + +import ( + "fmt" + "net" + "os" + "runtime" + + "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.TCPListener, error) { + type result struct { + lns []*net.TCPListener + 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.TCPListener + for _, port := range ports { + 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() + } + 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..db6a143f --- /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.TCPListener, 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..16064e10 --- /dev/null +++ b/apps/daemon/internal/gateway/mcp.go @@ -0,0 +1,62 @@ +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. 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 + 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 := sentValue(*s.BearerToken) + m.token = &token + m.transport = withhold(t, 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..196d213e --- /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, 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"]) + } + 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..718b4d43 --- /dev/null +++ b/apps/daemon/internal/gateway/model.go @@ -0,0 +1,77 @@ +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, keeps the base URL's +// path and withholds the key from response headers and trailers. +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 + } + key := sentValue(p.APIKey) + m := &modelRelay{protocol: p.Protocol, scheme: base.Scheme, host: base.Host, targets: map[string]*url.URL{}, + 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) + 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..db6faf24 --- /dev/null +++ b/apps/daemon/internal/gateway/model_test.go @@ -0,0 +1,271 @@ +package gateway + +import ( + "bufio" + "context" + "io" + "log" + "net/http" + "net/http/httptest" + "net/http/httptrace" + "net/textproto" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/internal/modelprovider" +) + +const upstreamKey = "sk-upstream-key" + +// 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: key}}}, + RootCAs: trust(srv), + }) + return eps.Models["main"], srv +} + +func TestModelInjectsTheKeyAndNeverThePlaceholder(t *testing.T) { + seen := make(chan *http.Request, 1) + 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 + 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 TestModelWithholdsTheKeyFromResponseHeaders(t *testing.T) { + // The upstream echoes the key it received in an informational response, + // 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) + 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("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("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, upstreamKey, 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, 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() + <-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, upstreamKey, 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..05eec8fb --- /dev/null +++ b/apps/daemon/internal/gateway/proxy.go @@ -0,0 +1,135 @@ +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. 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 +} + +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 l, ok := client.(interface{ SetLinger(int) error }); ok { + l.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..c7d9c85a --- /dev/null +++ b/apps/daemon/internal/gateway/proxy_test.go @@ -0,0 +1,49 @@ +package gateway + +import ( + "crypto/tls" + "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.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: &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 u, _ := url.Parse(target); resp.StatusCode != 200 || string(body) != "hello "+u.RequestURI() { + 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..d0bc5e2e --- /dev/null +++ b/apps/daemon/internal/gateway/relay.go @@ -0,0 +1,295 @@ +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" + + "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 + // 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. 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: &tls.Config{RootCAs: roots}, + ForceAttemptHTTP2: true, + DisableCompression: true, + TLSHandshakeTimeout: 10 * time.Second, + ExpectContinueTimeout: time.Second, + MaxIdleConns: 64, + IdleConnTimeout: 90 * time.Second, + } +} + +// 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. +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)) 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} + } + host, port, err := splitHostPort(addr) + if err != nil { + return nil, err + } + return connectSandbox(ctx, open, host, port) + } +} + +// 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 + } + 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, 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 +// 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 +} + +// 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. secret is the value as +// sent. +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) { + 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") + } +} diff --git a/contracts/agents-api/model-execution.md b/contracts/agents-api/model-execution.md index fb523ae5..53a35c80 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. -- It removes every value of each stripped header, matching names case-insensitively, then injects the upstream credential in the protocol's declared header. +- 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, 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.