From 9619f8aa0fd5c92f811e3fe29fb681f43857f253 Mon Sep 17 00:00:00 2001 From: adminturneddevops Date: Mon, 7 Sep 2026 15:51:29 -0400 Subject: [PATCH] security patching --- cmd/abox-guest/main.go | 67 +-- cmd/abox-vmm/start_darwin_arm64.go | 4 +- cmd/abox/main.go | 43 +- internal/agent/agent.go | 86 +++- internal/agent/agent_test.go | 7 +- internal/agentapi/types.go | 31 ++ internal/config/config.go | 5 + internal/config/config_test.go | 13 + internal/credsource/cloud_test.go | 20 +- internal/credsource/resolve.go | 35 +- internal/guest/brokerclient/client.go | 71 ++- internal/guest/brokerclient/client_test.go | 64 +++ internal/guest/egress/egress.go | 127 ----- internal/guest/egress/egress_test.go | 42 -- internal/guest/mcp/mcp.go | 276 ----------- internal/guest/mcp/mcp_test.go | 114 ----- internal/guest/mcpclient/client.go | 205 ++++++++ internal/guest/mcpclient/client_test.go | 263 ++++++++++ internal/guest/tools/builtins.go | 34 +- internal/hostbroker/broker.go | 97 ++++ internal/llmbroker/broker.go | 48 +- internal/llmbroker/broker_test.go | 103 +++- internal/mcpbroker/broker.go | 552 +++++++++++++++++++++ internal/mcpbroker/broker_test.go | 304 ++++++++++++ internal/provider/provider.go | 29 +- internal/repository/repository.go | 77 ++- internal/repository/repository_test.go | 177 ++++++- internal/runtime/runtime.go | 130 +++-- internal/runtime/runtime_push_test.go | 46 +- internal/runtime/runtime_test.go | 4 +- internal/runtime/runtime_turn_test.go | 16 +- internal/session/scrub.go | 9 +- internal/session/session.go | 12 +- internal/session/session_test.go | 10 +- internal/tui/approval_test.go | 154 ++++++ internal/tui/commands_test.go | 23 +- internal/tui/tui.go | 205 +++++++- pkg/abox/abox.go | 96 +++- protocol/protocol.go | 75 ++- 39 files changed, 2727 insertions(+), 947 deletions(-) create mode 100644 internal/agentapi/types.go delete mode 100644 internal/guest/egress/egress.go delete mode 100644 internal/guest/egress/egress_test.go delete mode 100644 internal/guest/mcp/mcp.go delete mode 100644 internal/guest/mcp/mcp_test.go create mode 100644 internal/guest/mcpclient/client.go create mode 100644 internal/guest/mcpclient/client_test.go create mode 100644 internal/hostbroker/broker.go create mode 100644 internal/mcpbroker/broker.go create mode 100644 internal/mcpbroker/broker_test.go create mode 100644 internal/tui/approval_test.go diff --git a/cmd/abox-guest/main.go b/cmd/abox-guest/main.go index 56c7fd8..8873a47 100644 --- a/cmd/abox-guest/main.go +++ b/cmd/abox-guest/main.go @@ -12,7 +12,6 @@ import ( "fmt" "io" "net" - "net/url" "os" "strings" "sync" @@ -21,8 +20,7 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/agent" "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/internal/guest/brokerclient" - "github.com/AdminTurnedDevOps/ABox/internal/guest/egress" - guestmcp "github.com/AdminTurnedDevOps/ABox/internal/guest/mcp" + "github.com/AdminTurnedDevOps/ABox/internal/guest/mcpclient" "github.com/AdminTurnedDevOps/ABox/internal/guest/tools" "github.com/AdminTurnedDevOps/ABox/protocol" "golang.org/x/sys/unix" @@ -37,9 +35,6 @@ func main() { func run() error { prepMounts() - if err := egress.ConfigureGuestResolver(); err != nil { - fmt.Fprintf(os.Stderr, "abox-guest: resolver: %v\n", err) - } cfg, err := loadConfig() if err != nil { return err @@ -48,23 +43,18 @@ func run() error { if err := os.MkdirAll(repo.Root, 0o755); err != nil { return err } - for _, s := range cfg.MCPServers { - if u, err := url.Parse(s.URL); err == nil { - egress.Allow(u.Hostname()) - } - } - mcpMgr := guestmcp.New(cfg.MCPServers, cfg.Secrets) - if err := mcpMgr.Connect(context.Background()); err != nil { - fmt.Fprintf(os.Stderr, "abox-guest: mcp: %v\n", err) + if len(cfg.Secrets) > 0 || len(cfg.MCPServers) > 0 { + return fmt.Errorf("legacy guest config contains credentials or MCP endpoints; rebuild the session") } - defer mcpMgr.Close() bclient := brokerclient.New() + mcpClient := mcpclient.New(bclient) loop := &agent.Loop{ - Model: config.ModelFromGuest(cfg.Model), - Repo: repo, - MCP: mcpMgr, - ContextFile: agent.DefaultContextFile, - Stream: bclient.Stream, + Model: config.ModelFromGuest(cfg.Model), + Repo: repo, + MCP: mcpClient, + ContextFile: agent.DefaultContextFile, + Stream: bclient.Stream, + ApproveRunCommand: bclient.RequestRunCommandApproval, } if err := loop.LoadContext(); err != nil { fmt.Fprintf(os.Stderr, "abox-guest: context: %v\n", err) @@ -93,11 +83,13 @@ func run() error { if ack.Error != nil { return ack.Error } - if len(ack.Result) > 0 { - var ackRes protocol.HelloResult - if err := json.Unmarshal(ack.Result, &ackRes); err == nil { - bclient.SetHostProtocol(ackRes.Protocol) - } + var ackRes protocol.HelloResult + if len(ack.Result) == 0 || json.Unmarshal(ack.Result, &ackRes) != nil || !ackRes.Accepted { + return fmt.Errorf("host rejected guest protocol") + } + bclient.SetHostProtocol(ackRes.Protocol) + if ackRes.Protocol < 4 { + return fmt.Errorf("host protocol %d cannot enforce brokered MCP and command approvals; run make build", ackRes.Protocol) } w := &connWriter{c: conn} @@ -156,7 +148,7 @@ func run() error { } return nil default: - resp := handle(loop, repo, mcpMgr, &archive, frame) + resp := handle(loop, repo, &archive, frame) if err := w.write(resp); err != nil { return err } @@ -312,6 +304,7 @@ func runTurn(w *connWriter, turns *turnTracker, loop *agent.Loop, bclient *broke turns.setCancel(req.ID, cancel) loop.MaxTurns = p.MaxTurns loop.Rich = p.RichEvents + loop.TurnID = req.ID if p.RichEvents { loop.Stream = bclient.StreamWithUsage } else { @@ -335,32 +328,12 @@ func runTurn(w *connWriter, turns *turnTracker, loop *agent.Loop, bclient *broke _ = w.write(protocol.Frame{ID: req.ID, Result: ok}) } -func applySecrets(secrets map[string]string) { - for k, v := range secrets { - if k != "" && v != "" { - _ = os.Setenv(k, v) - } - } -} - -func handle(loop *agent.Loop, repo tools.Repo, mcpMgr *guestmcp.Manager, archive *bytes.Buffer, req protocol.Frame) protocol.Frame { +func handle(loop *agent.Loop, repo tools.Repo, archive *bytes.Buffer, req protocol.Frame) protocol.Frame { out := protocol.Frame{V: protocol.Version, ID: req.ID} var err error switch req.Method { case "get_context": out.Result, _ = protocol.EncodeParams(protocol.GetContextResult{History: loop.History()}) - case "set_mcp_tokens": - p, e := protocol.DecodeParams[protocol.SetMCPTokensParams](req.Params) - if e != nil { - err = e - break - } - applySecrets(p.Secrets) - if mcpMgr != nil { - mcpMgr.SetSecrets(p.Secrets) - _ = mcpMgr.Connect(context.Background()) - } - out.Result, _ = protocol.EncodeParams(map[string]bool{"ok": true}) case "set_model": p, e := protocol.DecodeParams[protocol.SetModelParams](req.Params) if e != nil { diff --git a/cmd/abox-vmm/start_darwin_arm64.go b/cmd/abox-vmm/start_darwin_arm64.go index b880b23..3faf27d 100644 --- a/cmd/abox-vmm/start_darwin_arm64.go +++ b/cmd/abox-vmm/start_darwin_arm64.go @@ -35,8 +35,8 @@ func startVM(cfg vmmconfig.Config) error { if rc := C.krun_disable_implicit_vsock(id); rc < 0 { return fmt.Errorf("krun_disable_implicit_vsock: %d", int(rc)) } - if rc := C.krun_add_vsock(id, C.KRUN_TSI_HIJACK_INET); rc < 0 { - return fmt.Errorf("krun_add_vsock(TSI_INET): %d", int(rc)) + if rc := C.krun_add_vsock(id, 0); rc < 0 { + return fmt.Errorf("krun_add_vsock: %d", int(rc)) } sock := C.CString(cfg.RPCSocket) diff --git a/cmd/abox/main.go b/cmd/abox/main.go index 9d9d36f..cc88b7e 100644 --- a/cmd/abox/main.go +++ b/cmd/abox/main.go @@ -19,7 +19,7 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/internal/credentials" "github.com/AdminTurnedDevOps/ABox/internal/credsource" - "github.com/AdminTurnedDevOps/ABox/internal/llmbroker" + "github.com/AdminTurnedDevOps/ABox/internal/hostbroker" "github.com/AdminTurnedDevOps/ABox/internal/mcpauth" "github.com/AdminTurnedDevOps/ABox/internal/repository" "github.com/AdminTurnedDevOps/ABox/internal/runtime" @@ -36,6 +36,9 @@ func main() { } func run() error { + if err := scrubLegacySessions(); err != nil { + return err + } if len(os.Args) > 1 && os.Args[1] == "mcp" { return runMCP(os.Args[2:]) } @@ -65,9 +68,6 @@ func run() error { if err != nil { return err } - if err := scrubLegacySessions(); err != nil { - return err - } resolver := credsource.NewResolver() defer resolver.Close() @@ -122,16 +122,13 @@ func run() error { } var sb *runtime.Sandbox + var broker *hostbroker.Broker vmState := "not-started" image := cfg.Runtime.Image if image == "" { image = config.GuestImagePath() } - mcpServers, err := cfg.ResolvedMCPServers() - if err != nil { - return err - } - if err := runtime.Prepare(sess, image, sel, mcpServers, *resume); err != nil { + if err := runtime.Prepare(sess, image, sel, *resume); err != nil { if execMode { return err } @@ -159,10 +156,19 @@ func run() error { return fmt.Errorf("protocol-1 guest cannot use the secretless config; rebuild the guest image") } } + if started.GuestProtocol < 4 && !*probeVM { + started.Stop() + return fmt.Errorf("guest protocol %d cannot enforce brokered MCP and command approvals; rebuild the guest image and start a new session", started.GuestProtocol) + } sb = started vmState = "ready" defer sb.Stop() - sb.OnGuestCall = brokerForMode(cfg, resolver, execMode) + broker, err = brokerForMode(cfg, sel, resolver, execMode) + if err != nil { + return err + } + defer broker.Close() + sb.SetGuestCallHandler(broker) if err := pushSecrets(sb, cfg, resolver, sel); err != nil { if execMode { return err @@ -205,19 +211,25 @@ func run() error { _ = session.WriteTranscript(sess.TranscriptPath(), transcript) } } - return tui.Run(cfg, sel, sb, vmState, transcript, resolver, sess.TranscriptPath()) + return tui.Run(cfg, sel, sb, broker, vmState, transcript, resolver, sess.TranscriptPath()) } // Headless logs stream lifecycle; the TUI stays quiet so logs never paint into the UI. -func brokerForMode(cfg config.File, resolver *credsource.Resolver, execMode bool) *llmbroker.Broker { - b := llmbroker.New(cfg, resolver) +func brokerForMode(cfg config.File, sel config.Model, resolver *credsource.Resolver, execMode bool) (*hostbroker.Broker, error) { + b, err := hostbroker.New(cfg, sel, resolver) + if err != nil { + return nil, err + } if execMode { b.SetLogger(log.Printf) } - return b + return b, nil } func pushSecrets(sb *runtime.Sandbox, cfg config.File, resolver *credsource.Resolver, sel config.Model) error { + if sb.GuestProtocol >= 3 { + return nil + } ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() secrets, resolveErr := credsource.ResolveSelected(ctx, resolver, cfg, sel) @@ -280,9 +292,10 @@ func runExec(sb *runtime.Sandbox, prompt string) error { ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt) defer stop() enc := json.NewEncoder(os.Stdout) - return sb.UserTurn(ctx, prompt, func(e protocol.AgentEvent) { + _, err := sb.UserTurnCtx(ctx, prompt, runtime.TurnOptions{RichEvents: true}, func(e protocol.AgentEvent) { _ = enc.Encode(e) }) + return err } func loadResumeSession(wd, id string) (*session.Session, error) { diff --git a/internal/agent/agent.go b/internal/agent/agent.go index c005e78..59c88d8 100644 --- a/internal/agent/agent.go +++ b/internal/agent/agent.go @@ -10,15 +10,15 @@ import ( "strings" "time" + "github.com/AdminTurnedDevOps/ABox/internal/agentapi" "github.com/AdminTurnedDevOps/ABox/internal/config" - "github.com/AdminTurnedDevOps/ABox/internal/guest/mcp" "github.com/AdminTurnedDevOps/ABox/internal/guest/tools" - "github.com/AdminTurnedDevOps/ABox/internal/provider" "github.com/AdminTurnedDevOps/ABox/protocol" ) type MCPClient interface { - Tools() []mcp.Tool + Refresh(context.Context) error + Tools() []protocol.MCPTool Call(ctx context.Context, server, tool string, args json.RawMessage) (string, error) } @@ -31,20 +31,23 @@ type Loop struct { Model config.Model Repo tools.Repo MCP MCPClient - Messages []provider.Message + Messages []agentapi.Message ContextFile string OnEvent func(protocol.AgentEvent) MaxTurns int Rich bool + TurnID string - Stream func(ctx context.Context, model config.Model, messages []provider.Message, tools []provider.ToolSchema) (<-chan provider.Event, error) + ApproveRunCommand func(context.Context, protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) + + Stream func(ctx context.Context, model config.Model, messages []agentapi.Message, tools []agentapi.ToolSchema) (<-chan agentapi.Event, error) } -func BuiltinTools() []provider.ToolSchema { +func BuiltinTools() []agentapi.ToolSchema { specs := tools.BuiltinSpecs() - out := make([]provider.ToolSchema, len(specs)) + out := make([]agentapi.ToolSchema, len(specs)) for i, s := range specs { - out[i] = provider.ToolSchema{Name: s.Name, Description: s.Description, Parameters: s.Parameters} + out[i] = agentapi.ToolSchema{Name: s.Name, Description: s.Description, Parameters: s.Parameters} } return out } @@ -61,7 +64,13 @@ func (l *Loop) Turn(ctx context.Context, user string) error { l.emit(protocol.AgentEvent{Kind: "error", Err: err.Error()}) return err } - l.Messages = append(l.Messages, provider.Message{Role: "user", Content: user}) + if l.MCP != nil { + if err := l.MCP.Refresh(ctx); err != nil { + l.emit(protocol.AgentEvent{Kind: "error", Err: err.Error()}) + return err + } + } + l.Messages = append(l.Messages, agentapi.Message{Role: "user", Content: user}) limit := l.MaxTurns if limit <= 0 { limit = 16 @@ -81,7 +90,7 @@ func (l *Loop) Turn(ctx context.Context, user string) error { return err } var text string - var tool provider.Event + var tool agentapi.Event for ev := range events { switch ev.Type { case "text": @@ -109,7 +118,7 @@ func (l *Loop) Turn(ctx context.Context, user string) error { } if tool.ToolName == "" { if text != "" { - l.Messages = append(l.Messages, provider.Message{Role: "assistant", Content: text}) + l.Messages = append(l.Messages, agentapi.Message{Role: "assistant", Content: text}) } if l.Rich { ev := protocol.AgentEvent{Kind: "result", StopReason: stopReason} @@ -123,7 +132,7 @@ func (l *Loop) Turn(ctx context.Context, user string) error { _ = l.SaveContext() return nil } - l.Messages = append(l.Messages, provider.Message{ + l.Messages = append(l.Messages, agentapi.Message{ Role: "assistant", ToolID: tool.ToolID, ToolName: tool.ToolName, @@ -146,7 +155,7 @@ func (l *Loop) Turn(ctx context.Context, user string) error { } l.emit(tev) } - l.Messages = append(l.Messages, provider.Message{ + l.Messages = append(l.Messages, agentapi.Message{ Role: "tool", ToolID: tool.ToolID, ToolResult: result, @@ -183,7 +192,7 @@ func (l *Loop) History() []protocol.HistoryLine { return protocol.TrimHistory(HistoryFromMessages(l.Messages), protocol.MaxHistoryBytes) } -func HistoryFromMessages(msgs []provider.Message) []protocol.HistoryLine { +func HistoryFromMessages(msgs []agentapi.Message) []protocol.HistoryLine { var out []protocol.HistoryLine for _, m := range msgs { switch m.Role { @@ -215,7 +224,7 @@ func HistoryFromContextJSON(data []byte) ([]protocol.HistoryLine, error) { return HistoryFromMessages(msgs), nil } -func SaveMessages(path string, msgs []provider.Message) error { +func SaveMessages(path string, msgs []agentapi.Message) error { if path == "" { return fmt.Errorf("empty context path") } @@ -234,7 +243,7 @@ func SaveMessages(path string, msgs []provider.Message) error { return os.Rename(tmp, path) } -func LoadMessages(path string) ([]provider.Message, error) { +func LoadMessages(path string) ([]agentapi.Message, error) { data, err := os.ReadFile(path) if err != nil { return nil, err @@ -242,23 +251,23 @@ func LoadMessages(path string) ([]provider.Message, error) { return ParseMessages(data) } -func ParseMessages(data []byte) ([]provider.Message, error) { +func ParseMessages(data []byte) ([]agentapi.Message, error) { data = bytes.TrimSpace(data) if len(data) == 0 { return nil, nil } - var msgs []provider.Message + var msgs []agentapi.Message if err := json.Unmarshal(data, &msgs); err != nil { return nil, err } return msgs, nil } -func trimMessages(msgs []provider.Message, max int) []provider.Message { +func trimMessages(msgs []agentapi.Message, max int) []agentapi.Message { if max <= 0 { return nil } - var out []provider.Message + var out []agentapi.Message var size int for i := len(msgs) - 1; i >= 0; i-- { b, err := json.Marshal(msgs[i]) @@ -277,7 +286,7 @@ func trimMessages(msgs []provider.Message, max int) []provider.Message { return out } -func (l *Loop) allTools() []provider.ToolSchema { +func (l *Loop) allTools() []agentapi.ToolSchema { tools := BuiltinTools() if l.MCP == nil { return tools @@ -291,7 +300,7 @@ func (l *Loop) allTools() []provider.ToolSchema { if t.Server != "" { desc = "[" + t.Server + "] " + desc } - tools = append(tools, provider.ToolSchema{ + tools = append(tools, agentapi.ToolSchema{ Name: name, Description: desc, Parameters: t.Parameters, @@ -308,7 +317,7 @@ func splitMCPName(name string) (server, tool string, ok bool) { return server, tool, true } -func (l *Loop) execTool(ctx context.Context, ev provider.Event) (string, error) { +func (l *Loop) execTool(ctx context.Context, ev agentapi.Event) (string, error) { if server, tool, ok := splitMCPName(ev.ToolName); ok { if l.MCP == nil { return "", fmt.Errorf("mcp is not configured") @@ -319,7 +328,36 @@ func (l *Loop) execTool(ctx context.Context, ev provider.Event) (string, error) if deadline, ok := ctx.Deadline(); ok { timeout = time.Until(deadline) } - result, err := l.Repo.CallToolCtx(ctx, ev.ToolName, json.RawMessage(ev.ToolArgs), timeout) + approve := func(ctx context.Context, p protocol.RunCommandParams, effective time.Duration) error { + pending := protocol.AgentEvent{Kind: "approval", Tool: tools.RunCommand, Status: "pending"} + if l.Rich { + pending.ToolID = ev.ToolID + pending.ToolArgs = ev.ToolArgs + } + l.emit(pending) + if l.ApproveRunCommand == nil { + l.emit(protocol.AgentEvent{Kind: "approval", Tool: tools.RunCommand, Status: "denied", ToolID: ev.ToolID}) + return fmt.Errorf("run_command denied: no approval policy is configured") + } + seconds := int(effective / time.Second) + if seconds < 1 { + seconds = 1 + } + decision, err := l.ApproveRunCommand(ctx, protocol.RunCommandApprovalParams{ + TurnID: l.TurnID, ToolID: ev.ToolID, Command: p.Command, WorkDir: p.WorkDir, TimeoutSec: seconds, + }) + if err != nil { + l.emit(protocol.AgentEvent{Kind: "approval", Tool: tools.RunCommand, Status: "denied", ToolID: ev.ToolID}) + return fmt.Errorf("run_command denied: %w", err) + } + if decision != protocol.ApprovalAllowOnce { + l.emit(protocol.AgentEvent{Kind: "approval", Tool: tools.RunCommand, Status: "denied", ToolID: ev.ToolID}) + return fmt.Errorf("run_command denied") + } + l.emit(protocol.AgentEvent{Kind: "approval", Tool: tools.RunCommand, Status: "allowed", ToolID: ev.ToolID}) + return nil + } + result, err := l.Repo.CallApprovedToolCtx(ctx, ev.ToolName, json.RawMessage(ev.ToolArgs), timeout, approve) return tools.FormatToolResult(result, err) } diff --git a/internal/agent/agent_test.go b/internal/agent/agent_test.go index af198c8..6f06bcd 100644 --- a/internal/agent/agent_test.go +++ b/internal/agent/agent_test.go @@ -9,7 +9,6 @@ import ( "testing" "github.com/AdminTurnedDevOps/ABox/internal/config" - "github.com/AdminTurnedDevOps/ABox/internal/guest/mcp" "github.com/AdminTurnedDevOps/ABox/internal/guest/tools" "github.com/AdminTurnedDevOps/ABox/internal/provider" "github.com/AdminTurnedDevOps/ABox/protocol" @@ -155,8 +154,10 @@ type stubMCP struct { result string } -func (s *stubMCP) Tools() []mcp.Tool { - return []mcp.Tool{{ +func (s *stubMCP) Refresh(context.Context) error { return nil } + +func (s *stubMCP) Tools() []protocol.MCPTool { + return []protocol.MCPTool{{ Server: "svc", Name: "echo", Prefixed: "svc__echo", diff --git a/internal/agentapi/types.go b/internal/agentapi/types.go new file mode 100644 index 0000000..2b8dfbe --- /dev/null +++ b/internal/agentapi/types.go @@ -0,0 +1,31 @@ +// Package agentapi contains provider-neutral agent stream types shared by the +// guest loop and host provider implementation. +package agentapi + +import "github.com/AdminTurnedDevOps/ABox/protocol" + +type Event struct { + Type string + Text string + ToolName string + ToolID string + ToolArgs string + Err error + Usage *protocol.UsageInfo + StopReason string +} + +type Message struct { + Role string + Content string + ToolID string + ToolName string + ToolArgs string + ToolResult string +} + +type ToolSchema struct { + Name string + Description string + Parameters map[string]any +} diff --git a/internal/config/config.go b/internal/config/config.go index 33333ca..8bc76e9 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -282,6 +282,11 @@ func (s MCPServer) validate() error { } func (m Model) validate() error { + if m.BaseURL != "" { + if err := validateHTTPSURL("base_url", m.BaseURL); err != nil { + return err + } + } if m.CredentialEnv != "" && m.Credential != nil { return fmt.Errorf("set either credential or credential_env, not both") } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index c109699..20aa557 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -101,6 +101,19 @@ func TestValidateRejectsUnknownCredentialSource(t *testing.T) { } } +func TestModelBaseURLRequiresHTTPS(t *testing.T) { + c := Defaults() + c.Models[0].BaseURL = "http://api.example.com/v1" + if err := c.Validate(); err == nil || !strings.Contains(err.Error(), "models[0]: base_url must be https") { + t.Fatalf("http base_url: %v", err) + } + + c.Models[0].BaseURL = "https://api.example.com/v1" + if err := c.Validate(); err != nil { + t.Fatalf("https base_url: %v", err) + } +} + func TestValidateFieldVersionRules(t *testing.T) { c := Defaults() c.Models[0].CredentialEnv = "" diff --git a/internal/credsource/cloud_test.go b/internal/credsource/cloud_test.go index f460aac..2b8725f 100644 --- a/internal/credsource/cloud_test.go +++ b/internal/credsource/cloud_test.go @@ -583,8 +583,8 @@ func TestResolveSelectedOnlySelectedModel(t *testing.T) { if got["XAI_API_KEY"] != "xk" { t.Fatalf("model key missing: %#v", got) } - if got["ABOX_MCP_GH_TOKEN"] != "mtok" { - t.Fatalf("mcp token missing: %#v", got) + if _, ok := got["ABOX_MCP_GH_TOKEN"]; ok { + t.Fatalf("mcp token must remain host-brokered: %#v", got) } if _, ok := got["OPENAI_API_KEY"]; ok { t.Fatalf("unselected model key leaked: %#v", got) @@ -603,7 +603,7 @@ func TestResolveSelectedMissingModelError(t *testing.T) { } } -func TestResolveSelectedReturnsMCPTokenWithModelError(t *testing.T) { +func TestResolveSelectedDoesNotReturnMCPTokenWithModelError(t *testing.T) { r := NewResolver() source := &mapSource{vals: map[string]string{"env/MCP_TOKEN": "mcp-value"}} r.Register("env", source) @@ -615,12 +615,12 @@ func TestResolveSelectedReturnsMCPTokenWithModelError(t *testing.T) { if err == nil || !strings.Contains(err.Error(), "MODEL_TOKEN") { t.Fatalf("got error %v", err) } - if got["MCP_TOKEN"] != "mcp-value" { - t.Fatalf("partial secrets %#v", got) + if len(got) != 0 { + t.Fatalf("MCP token escaped host broker: %#v", got) } } -func TestResolveSelectedReportsMCPBackendErrorAndContinues(t *testing.T) { +func TestResolveSelectedDoesNotResolveMCPBackends(t *testing.T) { backendErr := errors.New("backend unavailable") r := NewResolver() source := &mapSource{ @@ -635,11 +635,11 @@ func TestResolveSelectedReportsMCPBackendErrorAndContinues(t *testing.T) { } model := config.Model{Name: "selected", Provider: "other", CredentialEnv: "MODEL_TOKEN"} got, err := ResolveSelected(context.Background(), r, cfg, model) - if !errors.Is(err, backendErr) || !strings.Contains(err.Error(), `mcp server "bad"`) { - t.Fatalf("got error %v", err) + if err != nil { + t.Fatalf("MCP backend affected legacy model resolution: %v", err) } - if got["MODEL_TOKEN"] != "model" || got["GOOD_TOKEN"] != "good" { - t.Fatalf("partial secrets %#v", got) + if got["MODEL_TOKEN"] != "model" || len(got) != 1 { + t.Fatalf("unexpected secrets %#v", got) } } diff --git a/internal/credsource/resolve.go b/internal/credsource/resolve.go index 364fc0e..90d8e1b 100644 --- a/internal/credsource/resolve.go +++ b/internal/credsource/resolve.go @@ -2,44 +2,23 @@ package credsource import ( "context" - "errors" "fmt" "github.com/AdminTurnedDevOps/ABox/internal/config" ) -// ResolveSelected returns the selected model credential plus enabled MCP -// tokens. A missing model credential is an error; a missing MCP token is skipped. +// ResolveSelected returns only the selected model credential for legacy +// protocol-2 guests. MCP credentials are always resolved by the host broker. func ResolveSelected(ctx context.Context, r *Resolver, cfg config.File, model config.Model) (map[string]string, error) { + _ = cfg out := map[string]string{} - var resolveErrs []error ref := model.CredentialReference() val, err := r.Resolve(ctx, FromConfig(ref)) if err != nil { val.Zero() - resolveErrs = append(resolveErrs, fmt.Errorf("credential for model %q (%s %s): %w", model.Name, ref.Source, ref.Name, err)) - } else { - out[model.EnvName()] = string(val.Bytes) - val.Zero() - } - - servers, err := cfg.ResolvedMCPServers() - if err != nil { - resolveErrs = append(resolveErrs, fmt.Errorf("resolve mcp servers: %w", err)) - return out, errors.Join(resolveErrs...) - } - for _, srv := range servers { - sref := srv.CredentialReference() - tok, err := r.Resolve(ctx, FromConfig(sref)) - if err != nil { - tok.Zero() - if !errors.Is(err, ErrNotFound) { - resolveErrs = append(resolveErrs, fmt.Errorf("credential for mcp server %q (%s %s): %w", srv.Name, sref.Source, sref.Name, err)) - } - continue - } - out[config.TokenEnv(srv)] = string(tok.Bytes) - tok.Zero() + return out, fmt.Errorf("credential for model %q (%s %s): %w", model.Name, ref.Source, ref.Name, err) } - return out, errors.Join(resolveErrs...) + out[model.EnvName()] = string(val.Bytes) + val.Zero() + return out, nil } diff --git a/internal/guest/brokerclient/client.go b/internal/guest/brokerclient/client.go index 5cdc20d..c958432 100644 --- a/internal/guest/brokerclient/client.go +++ b/internal/guest/brokerclient/client.go @@ -11,8 +11,8 @@ import ( "sync" "time" + "github.com/AdminTurnedDevOps/ABox/internal/agentapi" "github.com/AdminTurnedDevOps/ABox/internal/config" - "github.com/AdminTurnedDevOps/ABox/internal/provider" "github.com/AdminTurnedDevOps/ABox/protocol" ) @@ -58,7 +58,7 @@ func (p *pendingCall) complete(frame protocol.Frame) { } type clientStream struct { - out chan provider.Event + out chan agentapi.Event notify chan struct{} abort chan struct{} settled chan struct{} @@ -73,20 +73,20 @@ type clientStream struct { } type queuedEvent struct { - event provider.Event + event agentapi.Event size int } func newClientStream() *clientStream { s := &clientStream{ - out: make(chan provider.Event), notify: make(chan struct{}, 1), + out: make(chan agentapi.Event), notify: make(chan struct{}, 1), abort: make(chan struct{}), settled: make(chan struct{}), } go s.pump() return s } -func (s *clientStream) enqueue(ev provider.Event, terminal bool, size int) error { +func (s *clientStream) enqueue(ev agentapi.Event, terminal bool, size int) error { s.mu.Lock() if s.sealed { s.mu.Unlock() @@ -166,7 +166,7 @@ func (s *clientStream) pump() { s.terminalErr = nil s.mu.Unlock() select { - case s.out <- provider.Event{Type: "error", Err: err}: + case s.out <- agentapi.Event{Type: "error", Err: err}: case <-s.abort: return } @@ -218,6 +218,11 @@ func (c *Client) hostProtocol() int { return c.hostProto } +// HostProtocol returns the protocol version negotiated with the host. +func (c *Client) HostProtocol() int { + return c.hostProtocol() +} + func (c *Client) HandleFrame(f protocol.Frame) bool { if f.Method == "provider_event" { c.dispatchEvent(f) @@ -285,21 +290,51 @@ func (c *Client) call(ctx context.Context, method string, params any) (protocol. } } +// Call performs a guest-initiated host RPC. Semantic clients are responsible +// for checking that the negotiated protocol supports the requested method. +func (c *Client) Call(ctx context.Context, method string, params any) (protocol.Frame, error) { + return c.call(ctx, method, params) +} + +func (c *Client) RequestRunCommandApproval(ctx context.Context, params protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) { + if c.HostProtocol() < 4 { + return protocol.ApprovalDeny, fmt.Errorf("%w (host speaks protocol %d)", ErrHostTooOld, c.HostProtocol()) + } + frame, err := c.call(ctx, "request_run_command_approval", params) + if err != nil { + return protocol.ApprovalDeny, err + } + if frame.Error != nil { + return protocol.ApprovalDeny, frame.Error + } + var result protocol.RunCommandApprovalResult + if err := json.Unmarshal(frame.Result, &result); err != nil { + return protocol.ApprovalDeny, fmt.Errorf("malformed approval response") + } + if result.Decision != protocol.ApprovalAllowOnce && result.Decision != protocol.ApprovalDeny { + return protocol.ApprovalDeny, fmt.Errorf("invalid approval decision") + } + return result.Decision, nil +} + // A canceled provider_open can still succeed on the host; cancel that stream // so it does not occupy a slot until idle timeout. func (c *Client) reapCanceledCall(id string, call *pendingCall, method string) { - timer := time.NewTimer(canceledCallGrace) - defer timer.Stop() - select { - case <-call.ready: - case <-timer.C: + if method != "provider_open" { + timer := time.NewTimer(canceledCallGrace) + defer timer.Stop() + select { + case <-call.ready: + case <-timer.C: + c.forgetPending(id, call) + return + } c.forgetPending(id, call) return } + <-call.ready c.forgetPending(id, call) - if method == "provider_open" { - c.cancelOpenedFromFrame(call.frame) - } + c.cancelOpenedFromFrame(call.frame) } func (c *Client) forgetPending(id string, call *pendingCall) { @@ -327,15 +362,15 @@ func (c *Client) cancelOpenedFromFrame(frame protocol.Frame) { c.cancelHost(openRes.StreamID) } -func (c *Client) Stream(ctx context.Context, model config.Model, messages []provider.Message, tools []provider.ToolSchema) (<-chan provider.Event, error) { +func (c *Client) Stream(ctx context.Context, model config.Model, messages []agentapi.Message, tools []agentapi.ToolSchema) (<-chan agentapi.Event, error) { return c.stream(ctx, model, messages, tools, false) } -func (c *Client) StreamWithUsage(ctx context.Context, model config.Model, messages []provider.Message, tools []provider.ToolSchema) (<-chan provider.Event, error) { +func (c *Client) StreamWithUsage(ctx context.Context, model config.Model, messages []agentapi.Message, tools []agentapi.ToolSchema) (<-chan agentapi.Event, error) { return c.stream(ctx, model, messages, tools, true) } -func (c *Client) stream(ctx context.Context, model config.Model, messages []provider.Message, tools []provider.ToolSchema, rich bool) (<-chan provider.Event, error) { +func (c *Client) stream(ctx context.Context, model config.Model, messages []agentapi.Message, tools []agentapi.ToolSchema, rich bool) (<-chan agentapi.Event, error) { if c.hostProtocol() < 3 { return nil, fmt.Errorf("%w (host speaks protocol %d)", ErrHostTooOld, c.hostProtocol()) } @@ -455,7 +490,7 @@ func (c *Client) dispatchEvent(f protocol.Frame) { c.Close(fmt.Errorf("malformed provider event: %w", err)) return } - ev := provider.Event{ + ev := agentapi.Event{ Type: p.Type, Text: p.Text, ToolID: p.ToolID, ToolName: p.ToolName, ToolArgs: p.ToolArgs, Usage: p.Usage, StopReason: p.StopReason, diff --git a/internal/guest/brokerclient/client_test.go b/internal/guest/brokerclient/client_test.go index 03caeae..03a776e 100644 --- a/internal/guest/brokerclient/client_test.go +++ b/internal/guest/brokerclient/client_test.go @@ -444,6 +444,12 @@ func TestCanceledOpenCancelsLateHostStream(t *testing.T) { case <-time.After(2 * time.Second): t.Fatal("canceled open did not return") } + c.mu.Lock() + pending := len(c.pending) + c.mu.Unlock() + if pending != 1 { + t.Fatalf("canceled open tombstones = %d, want 1", pending) + } close(releaseOpen) select { case id := <-canceled: @@ -455,6 +461,64 @@ func TestCanceledOpenCancelsLateHostStream(t *testing.T) { } } +func TestCanceledCallTombstoneClearedOnConnectionClose(t *testing.T) { + c := New() + c.AttachContext(func(context.Context, protocol.Frame) error { return nil }) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + _, err := c.call(ctx, "provider_open", protocol.ProviderOpenParams{Model: "g"}) + if !errors.Is(err, context.Canceled) { + t.Fatalf("got %v", err) + } + c.mu.Lock() + pending := len(c.pending) + c.mu.Unlock() + if pending != 1 { + t.Fatalf("canceled open tombstones = %d, want 1", pending) + } + + c.Close(errors.New("disconnected")) + c.mu.Lock() + pending = len(c.pending) + c.mu.Unlock() + if pending != 0 { + t.Fatalf("connection close retained %d tombstones", pending) + } +} + +func TestLateOpenCancelIsIssuedOutsideClientLock(t *testing.T) { + c := New() + lockFree := make(chan bool, 1) + c.AttachContext(func(_ context.Context, frame protocol.Frame) error { + if frame.Method != "provider_cancel" { + return nil + } + acquired := c.mu.TryLock() + lockFree <- acquired + if !acquired { + return errors.New("provider_cancel issued while client lock held") + } + c.mu.Unlock() + result, _ := protocol.EncodeParams(map[string]bool{"ok": true}) + c.HandleFrame(protocol.Frame{ID: frame.ID, Result: result}) + return nil + }) + result, _ := protocol.EncodeParams(protocol.ProviderOpenResult{StreamID: "s-late"}) + done := make(chan struct{}) + go func() { + c.cancelOpenedFromFrame(protocol.Frame{Result: result}) + close(done) + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("provider_cancel deadlocked on client lock") + } + if !<-lockFree { + t.Fatal("provider_cancel was issued while client lock held") + } +} + func TestStreamSendErrorCancelsOpenedHostStream(t *testing.T) { c, stub := newClientPair(t, 3) canceled := make(chan string, 1) diff --git a/internal/guest/egress/egress.go b/internal/guest/egress/egress.go deleted file mode 100644 index 6c3305f..0000000 --- a/internal/guest/egress/egress.go +++ /dev/null @@ -1,127 +0,0 @@ -// Package egress allowlists guest outbound hosts. -package egress - -import ( - "context" - "fmt" - "maps" - "net" - "net/http" - "os" - "strings" - "sync" - "time" -) - -// Empty: protocol-3 guests only allow MCP origins added at boot. -var defaultAllowed = map[string]struct{}{} - -var ( - allowedMu sync.Mutex - allowedHosts = maps.Clone(defaultAllowed) -) - -var dnsServers = []string{"1.1.1.1:53", "8.8.8.8:53"} - -func normalizeHost(host string) string { - host = strings.ToLower(strings.TrimSpace(host)) - if h, _, err := net.SplitHostPort(host); err == nil { - host = h - } - return host -} - -func Allowed(host string) bool { - host = normalizeHost(host) - allowedMu.Lock() - defer allowedMu.Unlock() - _, ok := allowedHosts[host] - return ok -} - -func Allow(host string) { - host = normalizeHost(host) - if host == "" { - return - } - allowedMu.Lock() - defer allowedMu.Unlock() - allowedHosts[host] = struct{}{} -} - -func ResetForTest() { - allowedMu.Lock() - defer allowedMu.Unlock() - allowedHosts = maps.Clone(defaultAllowed) -} - -func ConfigureGuestResolver() error { - _ = os.MkdirAll("/etc", 0o755) - body := "nameserver 1.1.1.1\nnameserver 8.8.8.8\noptions ndots:1\n" - if err := os.WriteFile("/etc/resolv.conf", []byte(body), 0o644); err != nil { - return err - } - net.DefaultResolver = &net.Resolver{ - PreferGo: true, - Dial: func(ctx context.Context, _, _ string) (net.Conn, error) { - d := net.Dialer{Timeout: 5 * time.Second} - var last error - for _, srv := range dnsServers { - c, err := d.DialContext(ctx, "tcp4", srv) - if err == nil { - return c, nil - } - last = err - } - return nil, last - }, - } - return nil -} - -func Transport() *http.Transport { - dialer := &net.Dialer{Timeout: 15 * time.Second, KeepAlive: 30 * time.Second} - return &http.Transport{ - Proxy: nil, - ForceAttemptHTTP2: true, - MaxIdleConns: 8, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 15 * time.Second, - ExpectContinueTimeout: 1 * time.Second, - DialContext: func(ctx context.Context, network, address string) (net.Conn, error) { - host, port, err := net.SplitHostPort(address) - if err != nil { - return nil, err - } - if !Allowed(host) { - return nil, fmt.Errorf("egress denied: %s is not an allowed endpoint", host) - } - if port != "443" { - return nil, fmt.Errorf("egress denied: only HTTPS :443 is allowed") - } - ips, err := net.DefaultResolver.LookupIP(ctx, "ip4", host) - if err != nil { - return nil, fmt.Errorf("resolve %s: %w", host, err) - } - var last error - for _, ip := range ips { - c, err := dialer.DialContext(ctx, "tcp4", net.JoinHostPort(ip.String(), port)) - if err == nil { - return c, nil - } - last = err - } - if last == nil { - last = fmt.Errorf("no IPv4 addresses for %s", host) - } - return nil, last - }, - } -} - -func Client() *http.Client { - return &http.Client{ - Timeout: 5 * time.Minute, - Transport: Transport(), - } -} diff --git a/internal/guest/egress/egress_test.go b/internal/guest/egress/egress_test.go deleted file mode 100644 index 12be3fd..0000000 --- a/internal/guest/egress/egress_test.go +++ /dev/null @@ -1,42 +0,0 @@ -package egress - -import ( - "context" - "testing" -) - -func TestAllowedHosts(t *testing.T) { - t.Cleanup(ResetForTest) - ResetForTest() - // Protocol-3 guests broker LLM traffic through the host: no provider - // host is allowed by default, only configured MCP origins via Allow. - if Allowed("api.x.ai") || Allowed("api.openai.com") || Allowed("api.anthropic.com") { - t.Fatal("expected no LLM hosts allowed by default") - } - if Allowed("example.com") || Allowed("169.254.169.254") { - t.Fatal("expected arbitrary hosts denied") - } -} - -func TestDialDeniesOtherHosts(t *testing.T) { - tr := Transport() - _, err := tr.DialContext(context.Background(), "tcp", "example.com:443") - if err == nil { - t.Fatal("expected deny") - } -} - -func TestAllowAddsMCPHost(t *testing.T) { - t.Cleanup(ResetForTest) - ResetForTest() - if Allowed("mcp.example.com") { - t.Fatal("expected mcp host denied before Allow") - } - Allow("mcp.example.com") - if !Allowed("mcp.example.com") { - t.Fatal("expected mcp host allowed after Allow") - } - if Allowed("example.com") { - t.Fatal("expected unrelated host still denied") - } -} diff --git a/internal/guest/mcp/mcp.go b/internal/guest/mcp/mcp.go deleted file mode 100644 index f279231..0000000 --- a/internal/guest/mcp/mcp.go +++ /dev/null @@ -1,276 +0,0 @@ -// Package mcp is the guest Streamable HTTP MCP client. -package mcp - -import ( - "context" - "encoding/json" - "fmt" - "net/http" - "strings" - "sync" - "time" - - sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" - - "github.com/AdminTurnedDevOps/ABox/internal/guest/egress" - "github.com/AdminTurnedDevOps/ABox/protocol" -) - -const ( - maxSchemaBytes = 64 << 10 - maxResultBytes = 1 << 20 -) - -type Tool struct { - Server string - Name string - Prefixed string - Description string - Parameters map[string]any -} - -type Manager struct { - servers []protocol.GuestMCPServer - secrets map[string]string - client *http.Client - - mu sync.Mutex - conns map[string]*conn - tools []Tool -} - -type conn struct { - session *sdkmcp.ClientSession -} - -func New(servers []protocol.GuestMCPServer, secrets map[string]string) *Manager { - return &Manager{ - servers: servers, - secrets: secrets, - conns: map[string]*conn{}, - } -} - -func (m *Manager) WithHTTPClient(c *http.Client) *Manager { - m.client = c - return m -} - -func (m *Manager) SetSecrets(secrets map[string]string) { - m.mu.Lock() - defer m.mu.Unlock() - if m.secrets == nil { - m.secrets = map[string]string{} - } - for k, v := range secrets { - m.secrets[k] = v - } -} - -func (m *Manager) Connect(ctx context.Context) error { - for _, s := range m.servers { - if err := m.connectOne(ctx, s); err != nil { - // Best-effort: one failed server must not block the others. - continue - } - } - return nil -} - -func (m *Manager) connectOne(ctx context.Context, s protocol.GuestMCPServer) error { - if s.Name == "" || s.URL == "" { - return fmt.Errorf("mcp server missing name or url") - } - token := m.tokenFor(s) - httpClient := m.httpClientFor(token) - cli := sdkmcp.NewClient(&sdkmcp.Implementation{Name: "abox-guest", Version: "dev"}, nil) - cctx, cancel := context.WithTimeout(ctx, 30*time.Second) - defer cancel() - session, err := cli.Connect(cctx, &sdkmcp.StreamableClientTransport{ - Endpoint: s.URL, - HTTPClient: httpClient, - }, nil) - if err != nil { - return err - } - listed, err := session.ListTools(cctx, nil) - if err != nil { - _ = session.Close() - return err - } - allow := map[string]struct{}{} - for _, n := range s.Allowlist { - allow[n] = struct{}{} - } - m.mu.Lock() - defer m.mu.Unlock() - if existing := m.conns[s.Name]; existing != nil { - _ = existing.session.Close() - } - m.conns[s.Name] = &conn{session: session} - var kept []Tool - for _, t := range m.tools { - if t.Server != s.Name { - kept = append(kept, t) - } - } - m.tools = kept - for _, t := range listed.Tools { - if len(allow) > 0 { - if _, ok := allow[t.Name]; !ok { - continue - } - } - params, ok := schemaMap(t.InputSchema) - if !ok { - continue - } - m.tools = append(m.tools, Tool{ - Server: s.Name, - Name: t.Name, - Prefixed: s.Name + "__" + t.Name, - Description: t.Description, - Parameters: params, - }) - } - return nil -} - -func (m *Manager) Tools() []Tool { - m.mu.Lock() - defer m.mu.Unlock() - out := make([]Tool, len(m.tools)) - copy(out, m.tools) - return out -} - -func (m *Manager) Call(ctx context.Context, server, tool string, args json.RawMessage) (string, error) { - m.mu.Lock() - c := m.conns[server] - m.mu.Unlock() - if c == nil || c.session == nil { - return "", fmt.Errorf("mcp server %q is not connected", server) - } - var arguments any - if len(args) > 0 { - if err := json.Unmarshal(args, &arguments); err != nil { - return "", fmt.Errorf("mcp args: %w", err) - } - } - if _, ok := ctx.Deadline(); !ok { - var cancel context.CancelFunc - ctx, cancel = context.WithTimeout(ctx, 60*time.Second) - defer cancel() - } - res, err := c.session.CallTool(ctx, &sdkmcp.CallToolParams{Name: tool, Arguments: arguments}) - if err != nil { - return "", err - } - return resultText(res) -} - -func (m *Manager) Close() error { - m.mu.Lock() - defer m.mu.Unlock() - var last error - for name, c := range m.conns { - if c != nil && c.session != nil { - if err := c.session.Close(); err != nil { - last = err - } - } - delete(m.conns, name) - } - m.tools = nil - return last -} - -func (m *Manager) tokenFor(s protocol.GuestMCPServer) string { - m.mu.Lock() - defer m.mu.Unlock() - if s.TokenEnv == "" || m.secrets == nil { - return "" - } - return strings.TrimSpace(m.secrets[s.TokenEnv]) -} - -func (m *Manager) httpClientFor(token string) *http.Client { - base := m.client - if base == nil { - base = &http.Client{Timeout: 5 * time.Minute, Transport: egress.Transport()} - } - rt := base.Transport - if token == "" { - return base - } - return &http.Client{ - Timeout: base.Timeout, - Transport: bearerRT{base: rt, token: token}, - } -} - -type bearerRT struct { - base http.RoundTripper - token string -} - -func (b bearerRT) RoundTrip(req *http.Request) (*http.Response, error) { - r := req.Clone(req.Context()) - if b.token != "" { - r.Header.Set("Authorization", "Bearer "+b.token) - } - base := b.base - if base == nil { - base = http.DefaultTransport - } - return base.RoundTrip(r) -} - -func schemaMap(v any) (map[string]any, bool) { - if v == nil { - return map[string]any{"type": "object", "properties": map[string]any{}}, true - } - if m, ok := v.(map[string]any); ok { - b, err := json.Marshal(m) - if err != nil || len(b) > maxSchemaBytes { - return nil, false - } - return m, true - } - b, err := json.Marshal(v) - if err != nil || len(b) > maxSchemaBytes { - return nil, false - } - var m map[string]any - if err := json.Unmarshal(b, &m); err != nil { - return map[string]any{"type": "object", "properties": map[string]any{}}, true - } - return m, true -} - -func resultText(res *sdkmcp.CallToolResult) (string, error) { - if res == nil { - return "", fmt.Errorf("empty mcp result") - } - var b strings.Builder - for _, c := range res.Content { - switch t := c.(type) { - case *sdkmcp.TextContent: - b.WriteString(t.Text) - default: - raw, err := json.Marshal(c) - if err != nil { - continue - } - b.Write(raw) - } - } - s := b.String() - if len(s) > maxResultBytes { - s = s[:maxResultBytes] + "…" - } - if res.IsError { - return "", fmt.Errorf("%s", s) - } - return s, nil -} diff --git a/internal/guest/mcp/mcp_test.go b/internal/guest/mcp/mcp_test.go deleted file mode 100644 index 7bf119e..0000000 --- a/internal/guest/mcp/mcp_test.go +++ /dev/null @@ -1,114 +0,0 @@ -package mcp - -import ( - "context" - "encoding/json" - "io" - "net/http" - "net/http/httptest" - "strings" - "testing" - - sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" - - "github.com/AdminTurnedDevOps/ABox/protocol" -) - -type echoArgs struct { - Text string `json:"text"` -} - -func echoTool(_ context.Context, _ *sdkmcp.CallToolRequest, args echoArgs) (*sdkmcp.CallToolResult, any, error) { - return &sdkmcp.CallToolResult{ - Content: []sdkmcp.Content{&sdkmcp.TextContent{Text: "echo:" + args.Text}}, - }, nil, nil -} - -func serveMCP(t *testing.T, stateless bool) *httptest.Server { - t.Helper() - server := sdkmcp.NewServer(&sdkmcp.Implementation{Name: "fixture", Version: "1"}, nil) - sdkmcp.AddTool(server, &sdkmcp.Tool{Name: "echo", Description: "echo text"}, echoTool) - sdkmcp.AddTool(server, &sdkmcp.Tool{Name: "hidden", Description: "should be filtered"}, echoTool) - h := sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server { return server }, &sdkmcp.StreamableHTTPOptions{ - Stateless: stateless, - }) - return httptest.NewServer(h) -} - -func TestConnectStatefulListsPrefixedTools(t *testing.T) { - ts := serveMCP(t, false) - defer ts.Close() - m := New([]protocol.GuestMCPServer{{ - Name: "svc", - URL: ts.URL, - Allowlist: []string{"echo"}, - }}, nil).WithHTTPClient(ts.Client()) - if err := m.Connect(context.Background()); err != nil { - t.Fatal(err) - } - defer m.Close() - tools := m.Tools() - if len(tools) != 1 || tools[0].Prefixed != "svc__echo" { - t.Fatalf("tools=%#v", tools) - } -} - -func TestConnectStatelessCall(t *testing.T) { - ts := serveMCP(t, true) - defer ts.Close() - m := New([]protocol.GuestMCPServer{{Name: "svc", URL: ts.URL}}, nil).WithHTTPClient(ts.Client()) - if err := m.Connect(context.Background()); err != nil { - t.Fatal(err) - } - defer m.Close() - out, err := m.Call(context.Background(), "svc", "echo", json.RawMessage(`{"text":"hi"}`)) - if err != nil { - t.Fatal(err) - } - if !strings.Contains(out, "echo:hi") { - t.Fatalf("out=%q", out) - } -} - -func TestBearerHeaderSent(t *testing.T) { - inner := serveMCP(t, true) - defer inner.Close() - ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if got := r.Header.Get("Authorization"); got != "Bearer test-token" { - http.Error(w, "missing bearer", http.StatusUnauthorized) - return - } - inner.Config.Handler.ServeHTTP(w, r) - })) - defer ts.Close() - m := New([]protocol.GuestMCPServer{{ - Name: "svc", - URL: ts.URL, - TokenEnv: "SVC_TOKEN", - }}, map[string]string{"SVC_TOKEN": "test-token"}).WithHTTPClient(ts.Client()) - if err := m.Connect(context.Background()); err != nil { - t.Fatal(err) - } - defer m.Close() -} - -func TestConnect401DoesNotAbortOthers(t *testing.T) { - ok := serveMCP(t, true) - defer ok.Close() - deny := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - http.Error(w, "nope", http.StatusUnauthorized) - _, _ = io.Copy(io.Discard, r.Body) - })) - defer deny.Close() - m := New([]protocol.GuestMCPServer{ - {Name: "bad", URL: deny.URL}, - {Name: "ok", URL: ok.URL}, - }, nil).WithHTTPClient(http.DefaultClient) - if err := m.Connect(context.Background()); err != nil { - t.Fatal(err) - } - defer m.Close() - if len(m.Tools()) == 0 { - t.Fatal("expected ok server tools") - } -} diff --git a/internal/guest/mcpclient/client.go b/internal/guest/mcpclient/client.go new file mode 100644 index 0000000..32f8774 --- /dev/null +++ b/internal/guest/mcpclient/client.go @@ -0,0 +1,205 @@ +// Package mcpclient proxies semantic MCP operations through the host broker. +package mcpclient + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "sync" + "sync/atomic" + "time" + + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +const cancelTimeout = 2 * time.Second + +var ( + ErrHostTooOld = errors.New("host does not support MCP RPC; run make build") + nextCallID atomic.Uint64 +) + +// RPC is the guest broker functionality needed by the MCP proxy. +type RPC interface { + HostProtocol() int + Call(context.Context, string, any) (protocol.Frame, error) +} + +// Tool is an MCP tool advertised by the host. +type Tool = protocol.MCPTool + +type Client struct { + rpc RPC + + mu sync.RWMutex + tools []Tool + known map[toolRef]struct{} +} + +func New(rpc RPC) *Client { + return &Client{rpc: rpc, known: make(map[toolRef]struct{})} +} + +// Refresh fetches and atomically replaces the cached host MCP tool list. +func (c *Client) Refresh(ctx context.Context) error { + if err := c.requireV4(); err != nil { + return err + } + frame, err := c.rpc.Call(ctx, "mcp_list", protocol.MCPListParams{}) + if err != nil { + return fmt.Errorf("mcp_list: %w", err) + } + if frame.Error != nil { + return fmt.Errorf("mcp_list: %w", frame.Error) + } + + var result *protocol.MCPListResult + if len(frame.Result) == 0 || json.Unmarshal(frame.Result, &result) != nil || result == nil { + return errors.New("mcp_list: malformed result") + } + if len(result.Tools) > protocol.MaxMCPTools { + return fmt.Errorf("mcp_list: too many tools (maximum %d)", protocol.MaxMCPTools) + } + + tools := make([]Tool, len(result.Tools)) + known := make(map[toolRef]struct{}, len(result.Tools)) + prefixed := make(map[string]struct{}, len(result.Tools)) + for i, tool := range result.Tools { + if tool.Server == "" || tool.Name == "" || tool.Prefixed != tool.Server+"__"+tool.Name || tool.Parameters == nil { + return fmt.Errorf("mcp_list: invalid tool at index %d", i) + } + key := toolRef{server: tool.Server, tool: tool.Name} + if _, exists := known[key]; exists { + return fmt.Errorf("mcp_list: duplicate tool %q", tool.Prefixed) + } + if _, exists := prefixed[tool.Prefixed]; exists { + return fmt.Errorf("mcp_list: duplicate prefixed tool %q", tool.Prefixed) + } + schema, err := json.Marshal(tool.Parameters) + if err != nil || len(schema) > protocol.MaxMCPSchemaBytes { + return fmt.Errorf("mcp_list: invalid schema for tool %q", tool.Prefixed) + } + var parameters map[string]any + if json.Unmarshal(schema, ¶meters) != nil || parameters == nil { + return fmt.Errorf("mcp_list: invalid schema for tool %q", tool.Prefixed) + } + tool.Parameters = parameters + tools[i] = tool + known[key] = struct{}{} + prefixed[tool.Prefixed] = struct{}{} + } + + c.mu.Lock() + c.tools = tools + c.known = known + c.mu.Unlock() + return nil +} + +// Tools returns a deep copy of the currently cached tools. +func (c *Client) Tools() []Tool { + c.mu.RLock() + defer c.mu.RUnlock() + out := make([]Tool, len(c.tools)) + for i, tool := range c.tools { + out[i] = tool + out[i].Parameters = cloneMap(tool.Parameters) + } + return out +} + +// Call invokes a tool from the most recently refreshed cache. +func (c *Client) Call(ctx context.Context, server, tool string, args json.RawMessage) (string, error) { + if err := c.requireV4(); err != nil { + return "", err + } + c.mu.RLock() + _, ok := c.known[toolRef{server: server, tool: tool}] + c.mu.RUnlock() + if !ok { + return "", fmt.Errorf("mcp tool %q on server %q is not available", tool, server) + } + if len(args) > protocol.MaxMCPArgsBytes { + return "", fmt.Errorf("mcp arguments too large (maximum %d bytes)", protocol.MaxMCPArgsBytes) + } + if len(args) > 0 && !json.Valid(args) { + return "", errors.New("mcp arguments are not valid JSON") + } + + callID := fmt.Sprintf("mcp-%d", nextCallID.Add(1)) + frame, err := c.rpc.Call(ctx, "mcp_call", protocol.MCPCallParams{ + CallID: callID, Server: server, Tool: tool, Arguments: args, + }) + if err != nil && ctx.Err() != nil { + go c.cancel(callID) + } + if err != nil { + return "", fmt.Errorf("mcp_call: %w", err) + } + if frame.Error != nil { + return "", fmt.Errorf("mcp_call: %w", frame.Error) + } + + var result *protocol.MCPCallResult + if len(frame.Result) == 0 || json.Unmarshal(frame.Result, &result) != nil || result == nil { + return "", errors.New("mcp_call: malformed result") + } + if len(result.Text) > protocol.MaxMCPResultBytes { + return "", fmt.Errorf("mcp_call: result too large (maximum %d bytes)", protocol.MaxMCPResultBytes) + } + if result.IsError { + if result.Text == "" { + return "", errors.New("mcp tool returned an error") + } + return "", errors.New(result.Text) + } + return result.Text, nil +} + +func (c *Client) requireV4() error { + if c.rpc == nil { + return errors.New("mcp broker is not configured") + } + if got := c.rpc.HostProtocol(); got < 4 { + return fmt.Errorf("%w (host speaks protocol %d)", ErrHostTooOld, got) + } + return nil +} + +func (c *Client) cancel(callID string) { + ctx, cancel := context.WithTimeout(context.Background(), cancelTimeout) + defer cancel() + _, _ = c.rpc.Call(ctx, "mcp_cancel", protocol.MCPCancelParams{CallID: callID}) +} + +type toolRef struct { + server string + tool string +} + +func cloneMap(src map[string]any) map[string]any { + if src == nil { + return nil + } + out := make(map[string]any, len(src)) + for key, value := range src { + out[key] = cloneValue(value) + } + return out +} + +func cloneValue(value any) any { + switch value := value.(type) { + case map[string]any: + return cloneMap(value) + case []any: + out := make([]any, len(value)) + for i, item := range value { + out[i] = cloneValue(item) + } + return out + default: + return value + } +} diff --git a/internal/guest/mcpclient/client_test.go b/internal/guest/mcpclient/client_test.go new file mode 100644 index 0000000..3fa2f2f --- /dev/null +++ b/internal/guest/mcpclient/client_test.go @@ -0,0 +1,263 @@ +package mcpclient + +import ( + "context" + "encoding/json" + "errors" + "strings" + "sync" + "testing" + "time" + + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +type rpcCall struct { + ctx context.Context + method string + params any +} + +type stubRPC struct { + protocol int + + mu sync.Mutex + calls []rpcCall + call func(context.Context, string, any) (protocol.Frame, error) +} + +func (s *stubRPC) HostProtocol() int { return s.protocol } + +func (s *stubRPC) Call(ctx context.Context, method string, params any) (protocol.Frame, error) { + s.mu.Lock() + s.calls = append(s.calls, rpcCall{ctx: ctx, method: method, params: params}) + fn := s.call + s.mu.Unlock() + if fn == nil { + return protocol.Frame{}, nil + } + return fn(ctx, method, params) +} + +func resultFrame(t *testing.T, value any) protocol.Frame { + t.Helper() + raw, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + return protocol.Frame{Result: raw} +} + +func listedTool() protocol.MCPTool { + return protocol.MCPTool{ + Server: "svc", Name: "echo", Prefixed: "svc__echo", Description: "echo text", + Parameters: map[string]any{ + "type": "object", + "properties": map[string]any{"text": map[string]any{"type": "string"}}, + }, + } +} + +func TestRefreshCachesValidatedCopies(t *testing.T) { + rpc := &stubRPC{protocol: 4} + rpc.call = func(_ context.Context, method string, _ any) (protocol.Frame, error) { + if method != "mcp_list" { + t.Fatalf("method = %q", method) + } + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}}), nil + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + + got := c.Tools() + if len(got) != 1 || got[0].Prefixed != "svc__echo" { + t.Fatalf("tools = %#v", got) + } + got[0].Name = "changed" + got[0].Parameters["type"] = "changed" + got[0].Parameters["properties"].(map[string]any)["text"] = nil + again := c.Tools() + if again[0].Name != "echo" || again[0].Parameters["type"] != "object" { + t.Fatalf("cached tool mutated: %#v", again[0]) + } + properties := again[0].Parameters["properties"].(map[string]any) + if properties["text"] == nil { + t.Fatal("nested cached schema mutated") + } +} + +func TestRefreshFailureKeepsPreviousCache(t *testing.T) { + rpc := &stubRPC{protocol: 4} + result := protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}} + rpc.call = func(context.Context, string, any) (protocol.Frame, error) { + return resultFrame(t, result), nil + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + + bad := listedTool() + bad.Prefixed = "wrong" + result = protocol.MCPListResult{Tools: []protocol.MCPTool{bad}} + if err := c.Refresh(context.Background()); err == nil { + t.Fatal("expected invalid refresh error") + } + if got := c.Tools(); len(got) != 1 || got[0].Prefixed != "svc__echo" { + t.Fatalf("cache replaced after failed refresh: %#v", got) + } +} + +func TestRefreshRejectsTooManyTools(t *testing.T) { + rpc := &stubRPC{protocol: 4} + rpc.call = func(context.Context, string, any) (protocol.Frame, error) { + return resultFrame(t, protocol.MCPListResult{Tools: make([]protocol.MCPTool, protocol.MaxMCPTools+1)}), nil + } + if err := New(rpc).Refresh(context.Background()); err == nil || !strings.Contains(err.Error(), "too many tools") { + t.Fatalf("error = %v", err) + } +} + +func TestRefreshRejectsOversizedSchema(t *testing.T) { + rpc := &stubRPC{protocol: 4} + rpc.call = func(context.Context, string, any) (protocol.Frame, error) { + tool := listedTool() + tool.Parameters = map[string]any{"description": strings.Repeat("x", protocol.MaxMCPSchemaBytes)} + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{tool}}), nil + } + if err := New(rpc).Refresh(context.Background()); err == nil || !strings.Contains(err.Error(), "invalid schema") { + t.Fatalf("error = %v", err) + } +} + +func TestRequiresProtocolFour(t *testing.T) { + rpc := &stubRPC{protocol: 3} + c := New(rpc) + if err := c.Refresh(context.Background()); !errors.Is(err, ErrHostTooOld) { + t.Fatalf("refresh error = %v", err) + } + if _, err := c.Call(context.Background(), "svc", "echo", nil); !errors.Is(err, ErrHostTooOld) { + t.Fatalf("call error = %v", err) + } + if len(rpc.calls) != 0 { + t.Fatalf("unexpected RPCs: %#v", rpc.calls) + } +} + +func TestCallValidatesCacheAndArguments(t *testing.T) { + rpc := &stubRPC{protocol: 4} + rpc.call = func(_ context.Context, method string, _ any) (protocol.Frame, error) { + if method == "mcp_list" { + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}}), nil + } + return resultFrame(t, protocol.MCPCallResult{Text: "ok"}), nil + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + if _, err := c.Call(context.Background(), "svc", "missing", nil); err == nil { + t.Fatal("expected uncached tool error") + } + if _, err := c.Call(context.Background(), "svc", "echo", json.RawMessage(`{`)); err == nil { + t.Fatal("expected invalid JSON error") + } + if _, err := c.Call(context.Background(), "svc", "echo", make(json.RawMessage, protocol.MaxMCPArgsBytes+1)); err == nil { + t.Fatal("expected oversized arguments error") + } + if len(rpc.calls) != 1 { + t.Fatalf("invalid calls reached host: %d RPCs", len(rpc.calls)) + } +} + +func TestCallUsesUniqueIDsAndReturnsResult(t *testing.T) { + rpc := &stubRPC{protocol: 4} + var ids []string + rpc.call = func(_ context.Context, method string, params any) (protocol.Frame, error) { + if method == "mcp_list" { + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}}), nil + } + p, ok := params.(protocol.MCPCallParams) + if !ok { + t.Fatalf("params type = %T", params) + } + ids = append(ids, p.CallID) + if p.Server != "svc" || p.Tool != "echo" || string(p.Arguments) != `{"text":"hi"}` { + t.Fatalf("params = %#v", p) + } + return resultFrame(t, protocol.MCPCallResult{Text: "echo:hi"}), nil + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + for range 2 { + got, err := c.Call(context.Background(), "svc", "echo", json.RawMessage(`{"text":"hi"}`)) + if err != nil || got != "echo:hi" { + t.Fatalf("got %q, error %v", got, err) + } + } + if len(ids) != 2 || ids[0] == "" || ids[0] == ids[1] { + t.Fatalf("call IDs = %v", ids) + } +} + +func TestCallTurnsToolErrorIntoError(t *testing.T) { + rpc := &stubRPC{protocol: 4} + rpc.call = func(_ context.Context, method string, _ any) (protocol.Frame, error) { + if method == "mcp_list" { + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}}), nil + } + return resultFrame(t, protocol.MCPCallResult{Text: "tool failed", IsError: true}), nil + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + if _, err := c.Call(context.Background(), "svc", "echo", nil); err == nil || err.Error() != "tool failed" { + t.Fatalf("error = %v", err) + } +} + +func TestCanceledCallSendsCancelWithIndependentContext(t *testing.T) { + canceled := make(chan protocol.MCPCancelParams, 1) + rpc := &stubRPC{protocol: 4} + rpc.call = func(ctx context.Context, method string, params any) (protocol.Frame, error) { + switch method { + case "mcp_list": + return resultFrame(t, protocol.MCPListResult{Tools: []protocol.MCPTool{listedTool()}}), nil + case "mcp_call": + <-ctx.Done() + return protocol.Frame{}, ctx.Err() + case "mcp_cancel": + if ctx.Err() != nil { + t.Fatal("cancel RPC reused canceled context") + } + p := params.(protocol.MCPCancelParams) + canceled <- p + return resultFrame(t, map[string]bool{"ok": true}), nil + default: + t.Fatalf("method = %q", method) + return protocol.Frame{}, nil + } + } + c := New(rpc) + if err := c.Refresh(context.Background()); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := c.Call(ctx, "svc", "echo", nil); !errors.Is(err, context.Canceled) { + t.Fatalf("error = %v", err) + } + select { + case p := <-canceled: + if p.CallID == "" { + t.Fatal("empty canceled call ID") + } + case <-time.After(time.Second): + t.Fatal("mcp_cancel was not sent") + } +} diff --git a/internal/guest/tools/builtins.go b/internal/guest/tools/builtins.go index 730159c..80eafba 100644 --- a/internal/guest/tools/builtins.go +++ b/internal/guest/tools/builtins.go @@ -30,6 +30,8 @@ type Spec struct { Parameters map[string]any } +type RunCommandApprover func(context.Context, protocol.RunCommandParams, time.Duration) error + func BuiltinSpecs() []Spec { obj := func(props map[string]any) map[string]any { return map[string]any{"type": "object", "properties": props} @@ -64,7 +66,7 @@ func IsBuiltin(name string) bool { // CallBuiltin runs a host-guest RPC builtin. Empty params are an error. func (r Repo) CallBuiltin(name string, raw json.RawMessage, timeout time.Duration) (any, error) { - return r.call(context.Background(), name, raw, timeout, false) + return r.call(context.Background(), name, raw, timeout, false, nil) } // CallTool runs a model-facing builtin. Empty or invalid JSON uses zero params plus agent defaults. @@ -73,10 +75,14 @@ func (r Repo) CallTool(name string, raw json.RawMessage, timeout time.Duration) } func (r Repo) CallToolCtx(ctx context.Context, name string, raw json.RawMessage, timeout time.Duration) (any, error) { - return r.call(ctx, name, raw, timeout, true) + return r.call(ctx, name, raw, timeout, true, nil) +} + +func (r Repo) CallApprovedToolCtx(ctx context.Context, name string, raw json.RawMessage, timeout time.Duration, approve RunCommandApprover) (any, error) { + return r.call(ctx, name, raw, timeout, true, approve) } -func (r Repo) call(ctx context.Context, name string, raw json.RawMessage, timeout time.Duration, agent bool) (any, error) { +func (r Repo) call(ctx context.Context, name string, raw json.RawMessage, timeout time.Duration, agent bool, approve RunCommandApprover) (any, error) { switch name { case ListFiles: p, err := unmarshalParams[protocol.ListFilesParams](raw, agent) @@ -139,10 +145,32 @@ func (r Repo) call(ctx context.Context, name string, raw json.RawMessage, timeou if err != nil { return nil, err } + if len(p.Command) > protocol.MaxModelCommandBytes { + return nil, fmt.Errorf("command exceeds %d bytes", protocol.MaxModelCommandBytes) + } to := timeout if to <= 0 && p.Timeout > 0 { to = time.Duration(p.Timeout) * time.Second } + if to <= 0 { + to = 60 * time.Second + } + if agent { + if approve == nil { + return nil, fmt.Errorf("run_command approval is required") + } + if p.WorkDir != "" { + if _, err := r.Resolve(p.WorkDir); err != nil { + return nil, err + } + } + if err := approve(ctx, p, to); err != nil { + return nil, err + } + if err := ctx.Err(); err != nil { + return nil, err + } + } exit, stdout, stderr, dur, trunc, err := r.RunContext(ctx, p.Command, p.WorkDir, to, DefaultMaxOutput) return protocol.RunCommandResult{ ExitCode: exit, Stdout: stdout, Stderr: stderr, Duration: dur.String(), Trunc: trunc, diff --git a/internal/hostbroker/broker.go b/internal/hostbroker/broker.go new file mode 100644 index 0000000..3fde050 --- /dev/null +++ b/internal/hostbroker/broker.go @@ -0,0 +1,97 @@ +// Package hostbroker routes the guest's bounded semantic network operations to +// host-owned provider and MCP clients. +package hostbroker + +import ( + "context" + "encoding/json" + "fmt" + "sync" + + "github.com/AdminTurnedDevOps/ABox/internal/config" + "github.com/AdminTurnedDevOps/ABox/internal/credsource" + "github.com/AdminTurnedDevOps/ABox/internal/llmbroker" + "github.com/AdminTurnedDevOps/ABox/internal/mcpbroker" + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +type Broker struct { + mu sync.RWMutex + llm *llmbroker.Broker + mcp *mcpbroker.Broker + resolver *credsource.Resolver +} + +func New(cfg config.File, model config.Model, resolver *credsource.Resolver) (*Broker, error) { + servers, err := cfg.ResolvedMCPServers() + if err != nil { + return nil, err + } + return &Broker{ + llm: llmbroker.New(cfg, model, resolver), + mcp: mcpbroker.New(servers, resolver), + resolver: resolver, + }, nil +} + +func (b *Broker) Handle(ctx context.Context, method string, params json.RawMessage, notify func(string, any) error) (any, *protocol.Error) { + b.mu.RLock() + defer b.mu.RUnlock() + switch method { + case "provider_open", "provider_send", "provider_cancel": + return b.llm.Handle(ctx, method, params, notify) + case "mcp_list", "mcp_call", "mcp_cancel": + return b.mcp.Handle(ctx, method, params) + default: + return nil, &protocol.Error{Code: "host", Message: "unknown broker method " + method} + } +} + +func (b *Broker) UpdateModel(cfg config.File, model config.Model) { + b.mu.RLock() + llm := b.llm + b.mu.RUnlock() + llm.UpdateModel(cfg, model) +} + +func (b *Broker) UpdateMCP(cfg config.File) error { + servers, err := cfg.ResolvedMCPServers() + if err != nil { + return err + } + next := mcpbroker.New(servers, b.resolver) + b.mu.Lock() + old := b.mcp + b.mcp = next + b.mu.Unlock() + if old != nil { + return old.Close() + } + return nil +} + +func (b *Broker) SetMCPTokens(secrets map[string]string) error { + b.mu.RLock() + defer b.mu.RUnlock() + if b.mcp == nil { + return fmt.Errorf("MCP broker is not configured") + } + return b.mcp.SetTokens(secrets) +} + +func (b *Broker) SetLogger(logf func(string, ...any)) { + b.mu.RLock() + defer b.mu.RUnlock() + b.llm.SetLogger(logf) +} + +func (b *Broker) Close() error { + b.mu.Lock() + mcp := b.mcp + b.mcp = nil + b.mu.Unlock() + if mcp != nil { + return mcp.Close() + } + return nil +} diff --git a/internal/llmbroker/broker.go b/internal/llmbroker/broker.go index 9098292..ff65075 100644 --- a/internal/llmbroker/broker.go +++ b/internal/llmbroker/broker.go @@ -22,14 +22,15 @@ import ( const idleTimeout = 5 * time.Minute type Broker struct { - cfg config.File resolver *credsource.Resolver client *http.Client logf func(format string, args ...any) - mu sync.Mutex - next int - streams map[string]*stream + mu sync.Mutex + connectivityMode string + model config.Model + next int + streams map[string]*stream } type streamState uint8 @@ -53,15 +54,25 @@ type stream struct { bytesIn int } -func New(cfg config.File, resolver *credsource.Resolver) *Broker { +func New(cfg config.File, model config.Model, resolver *credsource.Resolver) *Broker { return &Broker{ - cfg: cfg, - resolver: resolver, - client: &http.Client{Timeout: 5 * time.Minute}, - streams: map[string]*stream{}, + connectivityMode: cfg.Connectivity.Mode, + model: model, + resolver: resolver, + client: &http.Client{Timeout: 5 * time.Minute}, + streams: map[string]*stream{}, } } +// UpdateModel changes the model available to future opens. Existing streams +// retain the model they opened with and continue to count toward the limit. +func (b *Broker) UpdateModel(cfg config.File, model config.Model) { + b.mu.Lock() + b.connectivityMode = cfg.Connectivity.Mode + b.model = model + b.mu.Unlock() +} + func (b *Broker) WithHTTPClient(c *http.Client) *Broker { if c != nil { b.client = c @@ -97,19 +108,21 @@ func (b *Broker) open(parent context.Context, raw json.RawMessage) (any, *protoc if err != nil { return nil, &protocol.Error{Code: "host", Message: err.Error()} } - if b.cfg.Connectivity.Mode == "offline" { - return nil, &protocol.Error{Code: "host", Message: "offline mode: provider access is disabled"} - } if strings.TrimSpace(p.Model) == "" { return nil, &protocol.Error{Code: "host", Message: "model alias required"} } - model, ok := b.cfg.ModelNamed(p.Model) - if !ok { - return nil, &protocol.Error{Code: "host", Message: fmt.Sprintf("unknown model profile %q", p.Model)} - } b.mu.Lock() - defer b.mu.Unlock() + if b.connectivityMode == "offline" { + b.mu.Unlock() + return nil, &protocol.Error{Code: "host", Message: "offline mode: provider access is disabled"} + } + model := b.model + if p.Model != model.Name { + b.mu.Unlock() + return nil, &protocol.Error{Code: "host", Message: fmt.Sprintf("model profile %q is not selected for this session", p.Model)} + } if len(b.streams) >= protocol.MaxProviderStreams { + b.mu.Unlock() return nil, &protocol.Error{Code: "host", Message: "too many open provider streams"} } b.next++ @@ -118,6 +131,7 @@ func (b *Broker) open(parent context.Context, raw json.RawMessage) (any, *protoc st := &stream{id: id, model: model, rich: p.Rich, state: streamReceiving, ctx: ctx, cancel: cancel} st.idleTimer = time.AfterFunc(idleTimeout, cancel) b.streams[id] = st + b.mu.Unlock() b.log("provider stream %s opened (model=%s rich=%v)", id, p.Model, p.Rich) go func() { <-ctx.Done() diff --git a/internal/llmbroker/broker_test.go b/internal/llmbroker/broker_test.go index 06322cb..88983f0 100644 --- a/internal/llmbroker/broker_test.go +++ b/internal/llmbroker/broker_test.go @@ -151,6 +151,15 @@ func defaultCfg() config.File { return cfg } +func newBroker(t *testing.T, cfg config.File, alias string, resolver *credsource.Resolver) *Broker { + t.Helper() + model, ok := cfg.ModelNamed(alias) + if !ok { + t.Fatalf("model %q not found in test config", alias) + } + return New(cfg, model, resolver) +} + func TestBrokerStreamsOpenAITextHostAuth(t *testing.T) { srv, authSeen := sseProvider(t, "openai", func(w http.ResponseWriter, auth, body string) { if !strings.HasPrefix(auth, "Bearer k1") { @@ -169,7 +178,7 @@ func TestBrokerStreamsOpenAITextHostAuth(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"k1"}}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) rec := ¬ifyRecorder{} id := openStream(t, b, "openai-default") @@ -206,7 +215,7 @@ func TestBrokerAnthropicKeyHeader(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"k-ant"}}) cfg.Models[2].CredentialEnv = "" cfg.Models[2].Credential = &config.CredentialRef{Source: "rot", Name: "anthropic"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "claude-default", r).WithHTTPClient(srv.Client()) rec := ¬ifyRecorder{} id := openStream(t, b, "claude-default") @@ -223,13 +232,63 @@ func TestBrokerAnthropicKeyHeader(t *testing.T) { } func TestBrokerUnknownAliasRejected(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("nope")), nil) - if perr == nil || !strings.Contains(perr.Message, "unknown model profile") { + if perr == nil || !strings.Contains(perr.Message, "not selected") { t.Fatalf("got %+v", perr) } } +func TestBrokerOpenScopedToSelectedModel(t *testing.T) { + cfg := defaultCfg() + b := newBroker(t, cfg, "openai-default", credsource.NewResolver()) + + _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("grok-default")), nil) + if perr == nil || !strings.Contains(perr.Message, "not selected") { + t.Fatalf("global but unselected alias was accepted: %+v", perr) + } + id := openStream(t, b, "openai-default") + b.mu.Lock() + st := b.streams[id] + b.mu.Unlock() + b.finish(st) +} + +func TestBrokerUpdateModelPreservesStreamsAndLimit(t *testing.T) { + cfg := defaultCfg() + b := newBroker(t, cfg, "openai-default", credsource.NewResolver()) + firstID := openStream(t, b, "openai-default") + + grok, ok := cfg.ModelNamed("grok-default") + if !ok { + t.Fatal("grok-default missing from test config") + } + b.UpdateModel(cfg, grok) + + b.mu.Lock() + first := b.streams[firstID] + b.mu.Unlock() + if first == nil || first.model.Name != "openai-default" { + t.Fatalf("existing stream model changed: %+v", first) + } + if _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("openai-default")), nil); perr == nil { + t.Fatal("previously selected model remained available") + } + secondID := openStream(t, b, "grok-default") + if _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("grok-default")), nil); perr == nil || !strings.Contains(perr.Message, "too many") { + t.Fatalf("stream limit did not survive update: %+v", perr) + } + + b.finish(first) + thirdID := openStream(t, b, "grok-default") + b.mu.Lock() + second := b.streams[secondID] + third := b.streams[thirdID] + b.mu.Unlock() + b.finish(second) + b.finish(third) +} + func providerOpenJSON(name string) protocol.ProviderOpenParams { return protocol.ProviderOpenParams{Model: name} } @@ -237,7 +296,7 @@ func providerOpenJSON(name string) protocol.ProviderOpenParams { func TestBrokerOfflineRejected(t *testing.T) { cfg := defaultCfg() cfg.Connectivity.Mode = "offline" - b := New(cfg, credsource.NewResolver()) + b := newBroker(t, cfg, "openai-default", credsource.NewResolver()) _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("openai-default")), nil) if perr == nil || !strings.Contains(perr.Message, "offline") { t.Fatalf("got %+v", perr) @@ -245,7 +304,7 @@ func TestBrokerOfflineRejected(t *testing.T) { } func TestBrokerUnknownStreamRejected(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) _, perr := b.Handle(context.Background(), "provider_send", mustJSON(t, protocol.ProviderSendParams{StreamID: "s99", Last: true}), nil) if perr == nil || !strings.Contains(perr.Message, "unknown provider stream") { t.Fatalf("got %+v", perr) @@ -255,7 +314,7 @@ func TestBrokerUnknownStreamRejected(t *testing.T) { func TestBrokerChunkBudgetRejected(t *testing.T) { cfg := defaultCfg() r := credsource.NewResolver() - b := New(cfg, r) + b := newBroker(t, cfg, "openai-default", r) id := openStream(t, b, "openai-default") remaining := protocol.MaxProviderRequest - 4 for remaining > 0 { @@ -276,7 +335,7 @@ func TestBrokerChunkBudgetRejected(t *testing.T) { } func TestBrokerRejectsOversizedChunkAndCleansStream(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) id := openStream(t, b, "openai-default") _, perr := sendChunk(t, b, ¬ifyRecorder{}, id, make([]byte, protocol.MaxProviderChunk+1), false) if perr == nil || !strings.Contains(perr.Message, "chunk too large") { @@ -299,7 +358,7 @@ func TestBrokerLastStartsProviderExactlyOnce(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"key"}}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) rec := ¬ifyRecorder{} id := openStream(t, b, "openai-default") body, _ := json.Marshal(protocol.ProviderRequest{Messages: []protocol.ProviderMessage{{Role: "user", Content: "q"}}}) @@ -326,7 +385,7 @@ func TestBrokerRejectsDataAfterStart(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"key"}}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) id := openStream(t, b, "openai-default") body, _ := json.Marshal(protocol.ProviderRequest{}) if _, perr := sendChunk(t, b, ¬ifyRecorder{}, id, body, true); perr != nil { @@ -338,7 +397,7 @@ func TestBrokerRejectsDataAfterStart(t *testing.T) { } func TestBrokerRequestCountBounds(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) id := openStream(t, b, "openai-default") req := protocol.ProviderRequest{Messages: make([]protocol.ProviderMessage, protocol.MaxProviderMessages+1)} body, _ := json.Marshal(req) @@ -349,7 +408,7 @@ func TestBrokerRequestCountBounds(t *testing.T) { } func TestBrokerToolCountBounds(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) id := openStream(t, b, "openai-default") req := protocol.ProviderRequest{Tools: make([]protocol.ProviderToolSchema, protocol.MaxProviderTools+1)} body, _ := json.Marshal(req) @@ -370,7 +429,7 @@ func TestBrokerRejectsOversizedProviderEvent(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"key"}}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) rec := ¬ifyRecorder{} id := openStream(t, b, "openai-default") body, _ := json.Marshal(protocol.ProviderRequest{}) @@ -414,7 +473,7 @@ func TestBrokerDrainsProviderAfterNotifyFailure(t *testing.T) { } func TestBrokerOpenContextClosesReceivingStream(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) ctx, cancel := context.WithCancel(context.Background()) res, perr := b.Handle(ctx, "provider_open", mustJSON(t, protocol.ProviderOpenParams{Model: "openai-default"}), nil) if perr != nil { @@ -451,7 +510,7 @@ func TestBrokerCredentialResolvedPerCall(t *testing.T) { r.Register("rot", rot) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) for i := 0; i < 2; i++ { rec := ¬ifyRecorder{} @@ -480,7 +539,7 @@ func TestBrokerMissingCredentialTypedError(t *testing.T) { r.Register("rot", &rotatingSource{vals: nil}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r) + b := newBroker(t, cfg, "openai-default", r) id := openStream(t, b, "openai-default") raw, _ := json.Marshal(protocol.ProviderRequest{}) _, perr := sendChunk(t, b, ¬ifyRecorder{}, id, raw, true) @@ -508,7 +567,7 @@ func TestBrokerCancelAbortsHTTP(t *testing.T) { r.Register("rot", &rotatingSource{vals: []string{"k"}}) cfg.Models[1].CredentialEnv = "" cfg.Models[1].Credential = &config.CredentialRef{Source: "rot", Name: "openai"} - b := New(cfg, r).WithHTTPClient(srv.Client()) + b := newBroker(t, cfg, "openai-default", r).WithHTTPClient(srv.Client()) rec := ¬ifyRecorder{} id := openStream(t, b, "openai-default") @@ -533,7 +592,7 @@ func TestBrokerCancelAbortsHTTP(t *testing.T) { func TestBrokerToolArgsBound(t *testing.T) { cfg := defaultCfg() r := credsource.NewResolver() - b := New(cfg, r) + b := newBroker(t, cfg, "openai-default", r) id := openStream(t, b, "openai-default") req := protocol.ProviderRequest{Messages: []protocol.ProviderMessage{{ Role: "assistant", ToolName: "x", ToolArgs: strings.Repeat("a", protocol.MaxProviderToolArgs+1), @@ -553,7 +612,7 @@ func TestBrokerToolArgsBound(t *testing.T) { } func TestBrokerUnknownMethodTypedError(t *testing.T) { - b := New(defaultCfg(), credsource.NewResolver()) + b := newBroker(t, defaultCfg(), "openai-default", credsource.NewResolver()) _, perr := b.Handle(context.Background(), "fetch_url", []byte(`{}`), nil) if perr == nil || !strings.Contains(perr.Message, "unknown broker method") { t.Fatalf("got %+v", perr) @@ -562,10 +621,10 @@ func TestBrokerUnknownMethodTypedError(t *testing.T) { func TestBrokerMaxConcurrentStreams(t *testing.T) { cfg := defaultCfg() - b := New(cfg, credsource.NewResolver()) + b := newBroker(t, cfg, "openai-default", credsource.NewResolver()) first := openStream(t, b, "openai-default") - second := openStream(t, b, "grok-default") - _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("claude-default")), nil) + second := openStream(t, b, "openai-default") + _, perr := b.Handle(context.Background(), "provider_open", mustJSON(t, providerOpenJSON("openai-default")), nil) if perr == nil || !strings.Contains(perr.Message, "too many open provider streams") { t.Fatalf("got %+v", perr) } diff --git a/internal/mcpbroker/broker.go b/internal/mcpbroker/broker.go new file mode 100644 index 0000000..a910948 --- /dev/null +++ b/internal/mcpbroker/broker.go @@ -0,0 +1,552 @@ +// Package mcpbroker is the host-side semantic MCP broker. The guest can name +// configured servers and discovered tools, but cannot supply URLs or headers. +package mcpbroker + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/url" + "strings" + "sync" + "time" + "unicode/utf8" + + sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/AdminTurnedDevOps/ABox/internal/config" + "github.com/AdminTurnedDevOps/ABox/internal/credsource" + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +const ( + connectTimeout = 30 * time.Second + maxConcurrentCalls = protocol.MaxGuestCalls - 1 // The RPC layer reserves one slot for cancellation. +) + +type Broker struct { + servers []config.MCPServer + resolver *credsource.Resolver + client *http.Client + + discoverMu sync.Mutex + mu sync.Mutex + discovered bool + closed bool + sessions map[string]*serverSession + tools []protocol.MCPTool + calls map[string]context.CancelFunc + overrides map[string]string +} + +type serverSession struct { + configured map[string]struct{} + discovered map[string]struct{} + session *sdkmcp.ClientSession +} + +// New accepts the already policy-resolved MCP server list. An empty list is +// the offline configuration and performs no network or credential access. +func New(servers []config.MCPServer, resolver *credsource.Resolver) *Broker { + configured := make([]config.MCPServer, len(servers)) + for i, server := range servers { + configured[i] = server + configured[i].Scopes = append([]string(nil), server.Scopes...) + configured[i].ToolAllowlist = append([]string(nil), server.ToolAllowlist...) + if server.Credential != nil { + credential := *server.Credential + configured[i].Credential = &credential + } + } + return &Broker{ + servers: configured, + resolver: resolver, + client: &http.Client{Timeout: 5 * time.Minute}, + sessions: map[string]*serverSession{}, + calls: map[string]context.CancelFunc{}, + overrides: map[string]string{}, + } +} + +// WithHTTPClient installs the base client used by MCP transports. Call it +// before Handle; each server gets an origin-bound clone of this client. +func (b *Broker) WithHTTPClient(client *http.Client) *Broker { + if client != nil { + b.client = client + } + return b +} + +// Handle serves the protocol v4 semantic MCP methods. +func (b *Broker) Handle(ctx context.Context, method string, raw json.RawMessage) (any, *protocol.Error) { + switch method { + case "mcp_list": + return b.list(ctx, raw) + case "mcp_call": + return b.call(ctx, raw) + case "mcp_cancel": + return b.cancel(raw) + default: + return nil, hostError("unknown broker method " + method) + } +} + +func (b *Broker) list(ctx context.Context, raw json.RawMessage) (any, *protocol.Error) { + if _, err := protocol.DecodeParams[protocol.MCPListParams](raw); err != nil { + return nil, hostError(err.Error()) + } + if err := b.discover(ctx); err != nil { + return nil, brokerError(err) + } + b.mu.Lock() + defer b.mu.Unlock() + tools := append([]protocol.MCPTool(nil), b.tools...) + return protocol.MCPListResult{Tools: tools}, nil +} + +func (b *Broker) discover(ctx context.Context) error { + b.discoverMu.Lock() + defer b.discoverMu.Unlock() + + b.mu.Lock() + if b.closed { + b.mu.Unlock() + return errors.New("MCP broker is closed") + } + if b.discovered { + b.mu.Unlock() + return nil + } + b.mu.Unlock() + + sessions := make(map[string]*serverSession, len(b.servers)) + var tools []protocol.MCPTool + for _, server := range b.servers { + if server.Name == "" || server.URL == "" { + closeSessions(sessions) + return errors.New("configured MCP server is missing an ID or URL") + } + if _, exists := sessions[server.Name]; exists { + closeSessions(sessions) + return fmt.Errorf("duplicate configured MCP server ID %q", server.Name) + } + if len(tools) >= protocol.MaxMCPTools { + break + } + session, serverTools, err := b.connect(ctx, server, protocol.MaxMCPTools-len(tools)) + if err != nil { + closeSessions(sessions) + return fmt.Errorf("MCP server %q: %w", server.Name, err) + } + sessions[server.Name] = session + tools = append(tools, serverTools...) + } + + b.mu.Lock() + if b.closed { + b.mu.Unlock() + closeSessions(sessions) + return errors.New("MCP broker is closed") + } + b.sessions = sessions + b.tools = tools + b.discovered = true + b.mu.Unlock() + return nil +} + +func (b *Broker) connect(parent context.Context, server config.MCPServer, limit int) (*serverSession, []protocol.MCPTool, error) { + endpoint, err := url.Parse(server.URL) + if err != nil || endpoint.Scheme == "" || endpoint.Host == "" || endpoint.User != nil { + return nil, nil, errors.New("configured endpoint is invalid") + } + token, err := b.resolveToken(parent, server) + if err != nil { + return nil, nil, err + } + client := originClient(b.client, endpoint, token) + token = "" + + ctx, cancel := context.WithTimeout(parent, connectTimeout) + defer cancel() + cli := sdkmcp.NewClient(&sdkmcp.Implementation{Name: "abox-host", Version: "dev"}, nil) + session, err := cli.Connect(ctx, &sdkmcp.StreamableClientTransport{ + Endpoint: server.URL, + HTTPClient: client, + }, nil) + if err != nil { + return nil, nil, err + } + + configured := make(map[string]struct{}, len(server.ToolAllowlist)) + for _, name := range server.ToolAllowlist { + configured[name] = struct{}{} + } + ss := &serverSession{ + configured: configured, + discovered: map[string]struct{}{}, + session: session, + } + tools, err := listTools(ctx, server.Name, ss, limit) + if err != nil { + _ = session.Close() + return nil, nil, err + } + return ss, tools, nil +} + +func (b *Broker) resolveToken(ctx context.Context, server config.MCPServer) (string, error) { + b.mu.Lock() + override, overridden := b.overrides[config.TokenEnv(server)] + b.mu.Unlock() + if overridden { + return strings.TrimSpace(override), nil + } + if b.resolver == nil { + return "", nil + } + value, err := b.resolver.Resolve(ctx, credsource.FromConfig(server.CredentialReference())) + if err != nil { + value.Zero() + if errors.Is(err, credsource.ErrNotFound) { + return "", nil + } + return "", errors.New("credential resolution failed") + } + token := strings.TrimSpace(string(value.Bytes)) + value.Zero() + return token, nil +} + +// SetTokens updates host-memory credential overrides and forces MCP sessions +// to reconnect. Unknown destinations are rejected transactionally. +func (b *Broker) SetTokens(secrets map[string]string) error { + allowed := make(map[string]struct{}, len(b.servers)) + for _, server := range b.servers { + allowed[config.TokenEnv(server)] = struct{}{} + } + for name := range secrets { + if _, ok := allowed[name]; !ok { + return fmt.Errorf("unknown MCP credential destination %q", name) + } + } + b.mu.Lock() + if b.closed { + b.mu.Unlock() + return errors.New("MCP broker is closed") + } + for name, value := range secrets { + b.overrides[name] = value + } + sessions := b.sessions + calls := b.calls + b.sessions = map[string]*serverSession{} + b.calls = map[string]context.CancelFunc{} + b.tools = nil + b.discovered = false + b.mu.Unlock() + for _, cancel := range calls { + cancel() + } + return closeSessions(sessions) +} + +func listTools(ctx context.Context, serverID string, ss *serverSession, limit int) ([]protocol.MCPTool, error) { + tools := make([]protocol.MCPTool, 0, limit) + seenCursors := map[string]struct{}{} + cursor := "" + for len(tools) < limit { + listed, err := ss.session.ListTools(ctx, &sdkmcp.ListToolsParams{Cursor: cursor}) + if err != nil { + return nil, err + } + for _, tool := range listed.Tools { + if tool == nil || tool.Name == "" { + continue + } + if len(ss.configured) > 0 { + if _, ok := ss.configured[tool.Name]; !ok { + continue + } + } + if _, duplicate := ss.discovered[tool.Name]; duplicate { + continue + } + parameters, ok := boundedSchema(tool.InputSchema) + if !ok { + continue + } + ss.discovered[tool.Name] = struct{}{} + tools = append(tools, protocol.MCPTool{ + Server: serverID, + Name: tool.Name, + Prefixed: serverID + "__" + tool.Name, + Description: tool.Description, + Parameters: parameters, + }) + if len(tools) == limit { + break + } + } + if listed.NextCursor == "" || len(tools) == limit { + break + } + if _, duplicate := seenCursors[listed.NextCursor]; duplicate { + return nil, errors.New("tool pagination repeated a cursor") + } + seenCursors[listed.NextCursor] = struct{}{} + cursor = listed.NextCursor + } + return tools, nil +} + +func boundedSchema(schema any) (map[string]any, bool) { + if schema == nil { + return map[string]any{"type": "object", "properties": map[string]any{}}, true + } + raw, err := json.Marshal(schema) + if err != nil || len(raw) > protocol.MaxMCPSchemaBytes { + return nil, false + } + var parameters map[string]any + if err := json.Unmarshal(raw, ¶meters); err != nil || parameters == nil { + return nil, false + } + return parameters, true +} + +func (b *Broker) call(parent context.Context, raw json.RawMessage) (any, *protocol.Error) { + p, err := protocol.DecodeParams[protocol.MCPCallParams](raw) + if err != nil { + return nil, hostError(err.Error()) + } + if p.CallID == "" { + return nil, hostError("MCP call ID required") + } + if len(p.Arguments) > protocol.MaxMCPArgsBytes { + return nil, hostError("MCP arguments too large") + } + var arguments any = map[string]any{} + if len(p.Arguments) > 0 { + if err := json.Unmarshal(p.Arguments, &arguments); err != nil { + return nil, hostError("malformed MCP arguments") + } + } + + b.mu.Lock() + if b.closed { + b.mu.Unlock() + return nil, hostError("MCP broker is closed") + } + ss := b.sessions[p.Server] + if ss == nil { + b.mu.Unlock() + return nil, hostError(fmt.Sprintf("MCP server %q is not configured and discovered", p.Server)) + } + if len(ss.configured) > 0 { + if _, ok := ss.configured[p.Tool]; !ok { + b.mu.Unlock() + return nil, hostError("MCP tool is not allowed") + } + } + if _, ok := ss.discovered[p.Tool]; !ok { + b.mu.Unlock() + return nil, hostError("MCP tool was not discovered") + } + if _, duplicate := b.calls[p.CallID]; duplicate { + b.mu.Unlock() + return nil, hostError("MCP call ID is already active") + } + if len(b.calls) >= maxConcurrentCalls { + b.mu.Unlock() + return nil, hostError("too many concurrent MCP calls") + } + ctx, cancel := context.WithCancel(parent) + b.calls[p.CallID] = cancel + b.mu.Unlock() + defer func() { + cancel() + b.mu.Lock() + delete(b.calls, p.CallID) + b.mu.Unlock() + }() + + if _, ok := ctx.Deadline(); !ok { + var timeoutCancel context.CancelFunc + ctx, timeoutCancel = context.WithTimeout(ctx, protocol.DefaultRPCTimeout) + defer timeoutCancel() + } + result, err := ss.session.CallTool(ctx, &sdkmcp.CallToolParams{Name: p.Tool, Arguments: arguments}) + if err != nil { + if errors.Is(ctx.Err(), context.Canceled) { + return nil, &protocol.Error{Code: "canceled", Message: "MCP call canceled"} + } + return nil, hostError(err.Error()) + } + return boundedResult(result) +} + +func boundedResult(result *sdkmcp.CallToolResult) (any, *protocol.Error) { + if result == nil { + return nil, hostError("empty MCP result") + } + var text strings.Builder + truncated := false + for _, content := range result.Content { + var part string + switch value := content.(type) { + case *sdkmcp.TextContent: + part = value.Text + default: + raw, err := json.Marshal(content) + if err != nil { + continue + } + part = string(raw) + } + remaining := protocol.MaxMCPResultBytes - text.Len() + if len(part) > remaining { + text.WriteString(validPrefix(part, remaining)) + truncated = true + break + } + text.WriteString(part) + } + return protocol.MCPCallResult{Text: text.String(), IsError: result.IsError, Truncated: truncated}, nil +} + +func validPrefix(value string, limit int) string { + if limit <= 0 { + return "" + } + if len(value) <= limit { + return value + } + value = value[:limit] + for !utf8.ValidString(value) { + value = value[:len(value)-1] + } + return value +} + +func (b *Broker) cancel(raw json.RawMessage) (any, *protocol.Error) { + p, err := protocol.DecodeParams[protocol.MCPCancelParams](raw) + if err != nil { + return nil, hostError(err.Error()) + } + if p.CallID == "" { + return nil, hostError("MCP call ID required") + } + b.mu.Lock() + cancel := b.calls[p.CallID] + b.mu.Unlock() + if cancel != nil { + cancel() + } + return map[string]bool{"ok": true}, nil +} + +// Close cancels active calls and closes all MCP sessions. It is idempotent. +func (b *Broker) Close() error { + b.mu.Lock() + if b.closed { + b.mu.Unlock() + return nil + } + b.closed = true + calls := b.calls + b.calls = map[string]context.CancelFunc{} + sessions := b.sessions + b.sessions = map[string]*serverSession{} + b.tools = nil + b.mu.Unlock() + + for _, cancel := range calls { + cancel() + } + return closeSessions(sessions) +} + +func closeSessions(sessions map[string]*serverSession) error { + var first error + for _, ss := range sessions { + if ss != nil && ss.session != nil { + if err := ss.session.Close(); err != nil && first == nil { + first = err + } + } + } + return first +} + +func originClient(base *http.Client, endpoint *url.URL, token string) *http.Client { + if base == nil { + base = http.DefaultClient + } + client := *base + origin := normalizedOrigin(endpoint) + transport := client.Transport + if transport == nil { + transport = http.DefaultTransport + } + client.Transport = originRoundTripper{base: transport, origin: origin, token: token} + previous := client.CheckRedirect + client.CheckRedirect = func(req *http.Request, via []*http.Request) error { + if normalizedOrigin(req.URL) != origin { + return errors.New("MCP redirect crossed the configured origin") + } + if previous != nil { + return previous(req, via) + } + if len(via) >= 10 { + return errors.New("stopped after 10 redirects") + } + return nil + } + return &client +} + +type originRoundTripper struct { + base http.RoundTripper + origin string + token string +} + +func (rt originRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { + if normalizedOrigin(req.URL) != rt.origin { + return nil, errors.New("MCP request left the configured origin") + } + clone := req.Clone(req.Context()) + if rt.token != "" { + clone.Header.Set("Authorization", "Bearer "+rt.token) + } + return rt.base.RoundTrip(clone) +} + +func normalizedOrigin(u *url.URL) string { + scheme := strings.ToLower(u.Scheme) + host := strings.ToLower(u.Hostname()) + port := u.Port() + if port == "" { + switch scheme { + case "http": + port = "80" + case "https": + port = "443" + } + } + return scheme + "\x00" + host + "\x00" + port +} + +func hostError(message string) *protocol.Error { + return &protocol.Error{Code: "host", Message: message} +} + +func brokerError(err error) *protocol.Error { + if errors.Is(err, context.Canceled) { + return &protocol.Error{Code: "canceled", Message: "MCP operation canceled"} + } + return hostError(err.Error()) +} diff --git a/internal/mcpbroker/broker_test.go b/internal/mcpbroker/broker_test.go new file mode 100644 index 0000000..849fcda --- /dev/null +++ b/internal/mcpbroker/broker_test.go @@ -0,0 +1,304 @@ +package mcpbroker + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" + + "github.com/AdminTurnedDevOps/ABox/internal/config" + "github.com/AdminTurnedDevOps/ABox/internal/credsource" + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +type staticSource struct { + value string + err error +} + +func (s staticSource) Resolve(context.Context, credsource.Reference) (credsource.Value, error) { + return credsource.Value{Bytes: []byte(s.value)}, s.err +} + +func (staticSource) Close() error { return nil } + +func raw(t *testing.T, value any) json.RawMessage { + t.Helper() + b, err := json.Marshal(value) + if err != nil { + t.Fatal(err) + } + return b +} + +func testServer(t *testing.T, configure func(*sdkmcp.Server)) *httptest.Server { + t.Helper() + server := sdkmcp.NewServer(&sdkmcp.Implementation{Name: "fixture", Version: "1"}, nil) + configure(server) + handler := sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server { return server }, &sdkmcp.StreamableHTTPOptions{Stateless: true}) + httpServer := httptest.NewServer(handler) + t.Cleanup(httpServer.Close) + return httpServer +} + +func addTextTool(server *sdkmcp.Server, name string, handler sdkmcp.ToolHandler) { + server.AddTool(&sdkmcp.Tool{ + Name: name, + Description: name + " description", + InputSchema: map[string]any{"type": "object", "properties": map[string]any{"text": map[string]any{"type": "string"}}}, + }, handler) +} + +func echoHandler(_ context.Context, req *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) { + var args struct { + Text string `json:"text"` + } + if err := json.Unmarshal(req.Params.Arguments, &args); err != nil { + return nil, err + } + return &sdkmcp.CallToolResult{Content: []sdkmcp.Content{&sdkmcp.TextContent{Text: "echo:" + args.Text}}}, nil +} + +func list(t *testing.T, broker *Broker) protocol.MCPListResult { + t.Helper() + result, perr := broker.Handle(context.Background(), "mcp_list", raw(t, protocol.MCPListParams{})) + if perr != nil { + t.Fatalf("mcp_list: %v", perr) + } + return result.(protocol.MCPListResult) +} + +func TestDiscoveryAndCallUseExactConfiguredID(t *testing.T) { + server := testServer(t, func(server *sdkmcp.Server) { + addTextTool(server, "echo", echoHandler) + }) + broker := New([]config.MCPServer{{Name: "svc", URL: server.URL}}, nil).WithHTTPClient(server.Client()) + t.Cleanup(func() { _ = broker.Close() }) + + listed := list(t, broker) + if len(listed.Tools) != 1 || listed.Tools[0].Server != "svc" || listed.Tools[0].Prefixed != "svc__echo" { + t.Fatalf("tools = %#v", listed.Tools) + } + result, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{ + CallID: "c1", Server: "svc", Tool: "echo", Arguments: json.RawMessage(`{"text":"hi"}`), + })) + if perr != nil { + t.Fatalf("mcp_call: %v", perr) + } + if got := result.(protocol.MCPCallResult); got.Text != "echo:hi" || got.IsError || got.Truncated { + t.Fatalf("result = %#v", got) + } + _, perr = broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "c2", Server: "SVC", Tool: "echo"})) + if perr == nil { + t.Fatal("case-changed server ID was accepted") + } +} + +func TestCallRechecksConfiguredAndDiscoveredAllowlist(t *testing.T) { + var hiddenCalls atomic.Int32 + server := testServer(t, func(server *sdkmcp.Server) { + addTextTool(server, "echo", echoHandler) + addTextTool(server, "hidden", func(context.Context, *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) { + hiddenCalls.Add(1) + return &sdkmcp.CallToolResult{}, nil + }) + }) + broker := New([]config.MCPServer{{Name: "svc", URL: server.URL, ToolAllowlist: []string{"echo"}}}, nil).WithHTTPClient(server.Client()) + t.Cleanup(func() { _ = broker.Close() }) + if got := list(t, broker).Tools; len(got) != 1 || got[0].Name != "echo" { + t.Fatalf("tools = %#v", got) + } + + _, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "bypass", Server: "svc", Tool: "hidden"})) + if perr == nil || !strings.Contains(perr.Message, "not allowed") { + t.Fatalf("allowlist bypass error = %#v", perr) + } + if hiddenCalls.Load() != 0 { + t.Fatal("hidden tool reached the MCP server") + } +} + +func TestBearerResolvedOnHostAndMissingCredentialAllowed(t *testing.T) { + var auth atomic.Value + inner := testServer(t, func(server *sdkmcp.Server) { addTextTool(server, "echo", echoHandler) }) + front := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + auth.Store(req.Header.Get("Authorization")) + inner.Config.Handler.ServeHTTP(w, req) + })) + t.Cleanup(front.Close) + + resolver := credsource.NewResolver() + resolver.Register("test", staticSource{value: "secret-token"}) + broker := New([]config.MCPServer{{ + Name: "svc", URL: front.URL, Credential: &config.CredentialRef{Source: "test", Name: "token"}, + }}, resolver).WithHTTPClient(front.Client()) + t.Cleanup(func() { _ = broker.Close() }) + list(t, broker) + if got, _ := auth.Load().(string); got != "Bearer secret-token" { + t.Fatalf("authorization = %q", got) + } + + missing := credsource.NewResolver() + missing.Register("test", staticSource{err: credsource.ErrNotFound}) + anonymous := New([]config.MCPServer{{ + Name: "anon", URL: inner.URL, Credential: &config.CredentialRef{Source: "test", Name: "missing"}, + }}, missing).WithHTTPClient(inner.Client()) + t.Cleanup(func() { _ = anonymous.Close() }) + if got := list(t, anonymous).Tools; len(got) != 1 { + t.Fatalf("anonymous tools = %#v", got) + } +} + +func TestCrossOriginRedirectRejectedWithoutLeakingBearer(t *testing.T) { + var targetRequests atomic.Int32 + var targetAuth atomic.Value + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + targetRequests.Add(1) + targetAuth.Store(req.Header.Get("Authorization")) + http.Error(w, "unexpected", http.StatusInternalServerError) + })) + t.Cleanup(target.Close) + redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) { + http.Redirect(w, req, target.URL, http.StatusTemporaryRedirect) + })) + t.Cleanup(redirect.Close) + + resolver := credsource.NewResolver() + resolver.Register("test", staticSource{value: "do-not-leak"}) + broker := New([]config.MCPServer{{ + Name: "svc", URL: redirect.URL, Credential: &config.CredentialRef{Source: "test", Name: "token"}, + }}, resolver).WithHTTPClient(redirect.Client()) + t.Cleanup(func() { _ = broker.Close() }) + _, perr := broker.Handle(context.Background(), "mcp_list", raw(t, protocol.MCPListParams{})) + if perr == nil || !strings.Contains(perr.Message, "crossed the configured origin") { + t.Fatalf("redirect error = %#v", perr) + } + if targetRequests.Load() != 0 { + got, _ := targetAuth.Load().(string) + t.Fatalf("redirect target received %d requests, auth %q", targetRequests.Load(), got) + } +} + +func TestProtocolBounds(t *testing.T) { + server := testServer(t, func(server *sdkmcp.Server) { + addTextTool(server, "large-result", func(context.Context, *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) { + return &sdkmcp.CallToolResult{Content: []sdkmcp.Content{&sdkmcp.TextContent{Text: strings.Repeat("x", protocol.MaxMCPResultBytes+100)}}}, nil + }) + server.AddTool(&sdkmcp.Tool{ + Name: "large-schema", + InputSchema: map[string]any{ + "type": "object", "description": strings.Repeat("s", protocol.MaxMCPSchemaBytes), + }, + }, echoHandler) + for i := 0; i < protocol.MaxMCPTools+5; i++ { + addTextTool(server, fmt.Sprintf("tool-%03d", i), echoHandler) + } + }) + broker := New([]config.MCPServer{{Name: "svc", URL: server.URL}}, nil).WithHTTPClient(server.Client()) + t.Cleanup(func() { _ = broker.Close() }) + listed := list(t, broker) + if len(listed.Tools) != protocol.MaxMCPTools { + t.Fatalf("tool count = %d, want %d", len(listed.Tools), protocol.MaxMCPTools) + } + for _, tool := range listed.Tools { + if tool.Name == "large-schema" { + t.Fatal("oversized schema was exposed") + } + } + + _, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{ + CallID: "large-args", Server: "svc", Tool: "large-result", Arguments: json.RawMessage(`{"x":"` + strings.Repeat("a", protocol.MaxMCPArgsBytes) + `"}`), + })) + if perr == nil || !strings.Contains(perr.Message, "arguments too large") { + t.Fatalf("argument bound error = %#v", perr) + } + result, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "large-result", Server: "svc", Tool: "large-result"})) + if perr != nil { + t.Fatalf("large result: %v", perr) + } + got := result.(protocol.MCPCallResult) + if len(got.Text) != protocol.MaxMCPResultBytes || !got.Truncated { + t.Fatalf("result bytes = %d, truncated = %v", len(got.Text), got.Truncated) + } +} + +func TestCancellationByCallID(t *testing.T) { + started := make(chan struct{}) + release := make(chan struct{}) + server := testServer(t, func(server *sdkmcp.Server) { + addTextTool(server, "wait", func(ctx context.Context, _ *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) { + close(started) + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-release: + return nil, errors.New("released") + } + }) + }) + broker := New([]config.MCPServer{{Name: "svc", URL: server.URL}}, nil).WithHTTPClient(server.Client()) + t.Cleanup(func() { _ = broker.Close() }) + list(t, broker) + + done := make(chan *protocol.Error, 1) + go func() { + _, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "cancel-me", Server: "svc", Tool: "wait"})) + done <- perr + }() + select { + case <-started: + case <-time.After(2 * time.Second): + t.Fatal("tool call did not start") + } + if _, perr := broker.Handle(context.Background(), "mcp_cancel", raw(t, protocol.MCPCancelParams{CallID: "cancel-me"})); perr != nil { + t.Fatalf("cancel: %v", perr) + } + select { + case perr := <-done: + if perr == nil || perr.Code != "canceled" { + t.Fatalf("call error = %#v", perr) + } + case <-time.After(2 * time.Second): + t.Fatal("canceled call did not return") + } + close(release) +} + +func TestConcurrentCallBound(t *testing.T) { + server := testServer(t, func(server *sdkmcp.Server) { addTextTool(server, "echo", echoHandler) }) + broker := New([]config.MCPServer{{Name: "svc", URL: server.URL}}, nil).WithHTTPClient(server.Client()) + t.Cleanup(func() { _ = broker.Close() }) + list(t, broker) + + broker.mu.Lock() + for i := 0; i < maxConcurrentCalls; i++ { + broker.calls[fmt.Sprintf("active-%d", i)] = func() {} + } + broker.mu.Unlock() + _, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "overflow", Server: "svc", Tool: "echo"})) + if perr == nil || !strings.Contains(perr.Message, "too many concurrent") { + t.Fatalf("concurrency bound error = %#v", perr) + } +} + +func TestOfflineEmptyListDoesNotResolveCredentials(t *testing.T) { + resolver := credsource.NewResolver() + resolver.Register("test", staticSource{err: errors.New("must not resolve")}) + broker := New(nil, resolver) + t.Cleanup(func() { _ = broker.Close() }) + if got := list(t, broker).Tools; len(got) != 0 { + t.Fatalf("offline tools = %#v", got) + } + _, perr := broker.Handle(context.Background(), "mcp_call", raw(t, protocol.MCPCallParams{CallID: "offline", Server: "svc", Tool: "echo"})) + if perr == nil { + t.Fatal("offline call was accepted") + } +} diff --git a/internal/provider/provider.go b/internal/provider/provider.go index 53731aa..b1d1a66 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -11,35 +11,14 @@ import ( "strings" "time" + "github.com/AdminTurnedDevOps/ABox/internal/agentapi" "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/protocol" ) -type Event struct { - Type string - Text string - ToolName string - ToolID string - ToolArgs string - Err error - Usage *protocol.UsageInfo - StopReason string -} - -type Message struct { - Role string - Content string - ToolID string - ToolName string - ToolArgs string - ToolResult string -} - -type ToolSchema struct { - Name string - Description string - Parameters map[string]any -} +type Event = agentapi.Event +type Message = agentapi.Message +type ToolSchema = agentapi.ToolSchema // Stream never reads the environment or builds its own client. func Stream(ctx context.Context, model config.Model, key string, client *http.Client, messages []Message, tools []ToolSchema) (<-chan Event, error) { diff --git a/internal/repository/repository.go b/internal/repository/repository.go index b659965..6d63e2a 100644 --- a/internal/repository/repository.go +++ b/internal/repository/repository.go @@ -39,17 +39,20 @@ func ValidateClean(start string) (Snapshot, error) { } // OpenForSession uses a clean committed worktree when one exists. -// Otherwise it copies the current directory into scratchDir, makes a +// Otherwise it copies the repository worktree into scratchDir, makes a // private commit there, and returns that. The host Git repo is not changed. func OpenForSession(start, scratchDir string) (Snapshot, error) { - if snap, err := ValidateClean(start); err == nil { + root, err := TopLevel(start) + if err != nil { + return Snapshot{}, fmt.Errorf("not a git worktree: %w", err) + } + if snap, err := ValidateClean(root); err == nil { return snap, nil } - abs, err := filepath.Abs(start) - if err != nil { - return Snapshot{}, err + if hasUnsupportedSubmodules(root) { + return Snapshot{}, fmt.Errorf("submodules are not supported in milestone one") } - if err := copyWorktree(abs, scratchDir); err != nil { + if err := copyWorktree(root, scratchDir); err != nil { return Snapshot{}, fmt.Errorf("ephemeral snapshot: %w", err) } if err := initScratchRepo(scratchDir); err != nil { @@ -60,7 +63,7 @@ func OpenForSession(start, scratchDir string) (Snapshot, error) { return Snapshot{}, fmt.Errorf("ephemeral snapshot: %w", err) } snap.Ephemeral = true - snap.HostSource = abs + snap.HostSource = root return snap, nil } @@ -116,28 +119,46 @@ func copyWorktree(src, dst string) error { if err := os.MkdirAll(dst, 0o700); err != nil { return err } - return filepath.Walk(src, func(path string, info os.FileInfo, err error) error { - if err != nil { - return err - } - rel, err := filepath.Rel(src, path) - if err != nil { - return err + cmd := exec.Command("git", "ls-files", "--cached", "--others", "--exclude-standard", "-z") + cmd.Dir = src + var stderr bytes.Buffer + cmd.Stderr = &stderr + out, err := cmd.Output() + if err != nil { + return fmt.Errorf("git ls-files: %w: %s", err, strings.TrimSpace(stderr.String())) + } + for _, name := range bytes.Split(out, []byte{0}) { + if len(name) == 0 { + continue } - if rel == "." { - return nil + rel := filepath.FromSlash(string(name)) + clean := filepath.Clean(rel) + if filepath.IsAbs(clean) || clean == "." || clean == ".." || strings.HasPrefix(clean, ".."+string(filepath.Separator)) { + return fmt.Errorf("unsafe repository path %q", rel) } - base := filepath.Base(path) - if info.IsDir() && (base == ".git" || base == "bin") { - return filepath.SkipDir + path := src + var info os.FileInfo + parts := strings.Split(clean, string(filepath.Separator)) + for i, part := range parts { + path = filepath.Join(path, part) + info, err = os.Lstat(path) + if err != nil { + break + } + if info.Mode()&os.ModeSymlink != 0 || i < len(parts)-1 && !info.IsDir() { + return fmt.Errorf("unsupported file type %q", rel) + } } - target := filepath.Join(dst, rel) - if info.IsDir() { - return os.MkdirAll(target, 0o755) + if err != nil { + if os.IsNotExist(err) { + continue // A tracked file deleted in the dirty worktree stays deleted. + } + return err } if !info.Mode().IsRegular() { - return nil + return fmt.Errorf("unsupported file type %q", rel) } + target := filepath.Join(dst, clean) if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil { return err } @@ -145,8 +166,14 @@ func copyWorktree(src, dst string) error { if err != nil { return err } - return os.WriteFile(target, data, info.Mode().Perm()) - }) + if err := os.WriteFile(target, data, info.Mode().Perm()); err != nil { + return err + } + if err := os.Chmod(target, info.Mode().Perm()); err != nil { + return err + } + } + return nil } func initScratchRepo(dir string) error { diff --git a/internal/repository/repository_test.go b/internal/repository/repository_test.go index 95d882b..434b4cb 100644 --- a/internal/repository/repository_test.go +++ b/internal/repository/repository_test.go @@ -1,6 +1,9 @@ package repository import ( + "archive/tar" + "bytes" + "io" "os" "os/exec" "path/filepath" @@ -40,20 +43,188 @@ func TestValidateClean(t *testing.T) { func TestOpenForSessionEphemeral(t *testing.T) { dir := t.TempDir() - if err := os.WriteFile(filepath.Join(dir, "hello.txt"), []byte("hi"), 0o644); err != nil { + run := func(args ...string) { + t.Helper() + cmd := exec.Command(args[0], args[1:]...) + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_AUTHOR_NAME=t", "GIT_AUTHOR_EMAIL=t@t", "GIT_COMMITTER_NAME=t", "GIT_COMMITTER_EMAIL=t@t") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("%v: %s", err, out) + } + } + run("git", "init", "-b", "main") + if err := os.WriteFile(filepath.Join(dir, "script.sh"), []byte("#!/bin/sh\necho clean\n"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "tracked.log"), []byte("clean"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "deleted.txt"), []byte("remove me"), 0o644); err != nil { + t.Fatal(err) + } + run("git", "add", "script.sh", "tracked.log", "deleted.txt") + run("git", "commit", "-m", "init") + + if err := os.WriteFile(filepath.Join(dir, ".gitignore"), []byte("*.env\n*.log\n"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, ".git", "info", "exclude"), []byte("info-secret\n"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "script.sh"), []byte("#!/bin/sh\necho dirty\n"), 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "tracked.log"), []byte("dirty"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.Remove(filepath.Join(dir, "deleted.txt")); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "ignored.env"), []byte("secret"), 0o600); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(dir, "info-secret"), []byte("secret"), 0o600); err != nil { + t.Fatal(err) + } + subdir := filepath.Join(dir, "subdir") + if err := os.Mkdir(subdir, 0o755); err != nil { t.Fatal(err) } + oddName := "untracked\nfile.txt" + if err := os.WriteFile(filepath.Join(subdir, oddName), []byte("included"), 0o644); err != nil { + t.Fatal(err) + } + scratch := t.TempDir() - snap, err := OpenForSession(dir, scratch) + snap, err := OpenForSession(subdir, scratch) if err != nil { t.Fatal(err) } if !snap.Ephemeral { t.Fatal("expected ephemeral snapshot") } - if _, err := os.Stat(filepath.Join(snap.Root, "hello.txt")); err != nil { + hostInfo, err := os.Stat(snap.HostSource) + if err != nil { t.Fatal(err) } + dirInfo, err := os.Stat(dir) + if err != nil { + t.Fatal(err) + } + if !os.SameFile(hostInfo, dirInfo) { + t.Fatalf("HostSource=%q is not repository root %q", snap.HostSource, dir) + } + for name, want := range map[string]string{ + "script.sh": "#!/bin/sh\necho dirty\n", + "tracked.log": "dirty", + ".gitignore": "*.env\n*.log\n", + filepath.Join("subdir", oddName): "included", + } { + got, err := os.ReadFile(filepath.Join(snap.Root, name)) + if err != nil { + t.Fatalf("read %q: %v", name, err) + } + if string(got) != want { + t.Fatalf("%q=%q want %q", name, got, want) + } + } + for _, name := range []string{"deleted.txt", "ignored.env", "info-secret"} { + if _, err := os.Stat(filepath.Join(snap.Root, name)); !os.IsNotExist(err) { + t.Fatalf("excluded file %q entered host-tree: %v", name, err) + } + } + info, err := os.Stat(filepath.Join(snap.Root, "script.sh")) + if err != nil { + t.Fatal(err) + } + if info.Mode().Perm()&0o111 == 0 { + t.Fatalf("script mode=%o, executable bit lost", info.Mode().Perm()) + } + + archive, err := ArchiveHEAD(snap.Root) + if err != nil { + t.Fatal(err) + } + archived := make(map[string]int64) + tr := tar.NewReader(bytes.NewReader(archive)) + for { + hdr, err := tr.Next() + if err == io.EOF { + break + } + if err != nil { + t.Fatal(err) + } + archived[hdr.Name] = hdr.Mode + } + for _, name := range []string{"deleted.txt", "ignored.env", "info-secret"} { + if _, ok := archived[name]; ok { + t.Fatalf("excluded file %q entered archive", name) + } + } + if archived["script.sh"]&0o111 == 0 { + t.Fatalf("archived script mode=%o, executable bit lost", archived["script.sh"]) + } + if _, ok := archived[filepath.ToSlash(filepath.Join("subdir", oddName))]; !ok { + t.Fatalf("NUL-delimited untracked name missing from archive: %#v", archived) + } +} + +func TestOpenForSessionRejectsSymlink(t *testing.T) { + dir := t.TempDir() + cmd := exec.Command("git", "init", "-b", "main") + cmd.Dir = dir + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("git init: %v: %s", err, out) + } + if err := os.WriteFile(filepath.Join(dir, "target"), []byte("data"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.Symlink("target", filepath.Join(dir, "link")); err != nil { + t.Skipf("symlinks unavailable: %v", err) + } + _, err := OpenForSession(dir, t.TempDir()) + if err == nil || !strings.Contains(err.Error(), "unsupported file type") { + t.Fatalf("got %v", err) + } +} + +func TestOpenForSessionRejectsSymlinkedDirectory(t *testing.T) { + dir := t.TempDir() + run := func(args ...string) { + t.Helper() + cmd := exec.Command(args[0], args[1:]...) + cmd.Dir = dir + cmd.Env = append(os.Environ(), "GIT_AUTHOR_NAME=t", "GIT_AUTHOR_EMAIL=t@t", "GIT_COMMITTER_NAME=t", "GIT_COMMITTER_EMAIL=t@t") + if out, err := cmd.CombinedOutput(); err != nil { + t.Fatalf("%v: %s", err, out) + } + } + run("git", "init", "-b", "main") + nested := filepath.Join(dir, "nested") + if err := os.Mkdir(nested, 0o755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(nested, "file.txt"), []byte("inside"), 0o644); err != nil { + t.Fatal(err) + } + run("git", "add", "nested/file.txt") + run("git", "commit", "-m", "init") + + outside := t.TempDir() + if err := os.WriteFile(filepath.Join(outside, "file.txt"), []byte("outside"), 0o644); err != nil { + t.Fatal(err) + } + if err := os.RemoveAll(nested); err != nil { + t.Fatal(err) + } + if err := os.Symlink(outside, nested); err != nil { + t.Skipf("symlinks unavailable: %v", err) + } + _, err := OpenForSession(dir, t.TempDir()) + if err == nil || !strings.Contains(err.Error(), "unsupported file type") { + t.Fatalf("got %v", err) + } } func TestValidateCleanEmptyRepo(t *testing.T) { diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index 19951a1..c14612a 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -58,6 +58,16 @@ type GuestCallHandler interface { notify func(method string, params any) error) (any, *protocol.Error) } +type RunCommandApprover interface { + ApproveRunCommand(context.Context, protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) +} + +type RunCommandApproverFunc func(context.Context, protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) + +func (f RunCommandApproverFunc) ApproveRunCommand(ctx context.Context, params protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) { + return f(ctx, params) +} + type Sandbox struct { Sess *session.Session History []protocol.HistoryLine @@ -74,7 +84,10 @@ type Sandbox struct { nextID int calls map[string]*frameQueue activeTurn string + turnCtx context.Context + turnCancel context.CancelFunc turnQ *frameQueue + approver RunCommandApprover lifeCtx context.Context lifeCancel context.CancelFunc guestSlots chan struct{} @@ -208,7 +221,7 @@ func (q *frameQueue) pop(ctx context.Context) (protocol.Frame, bool, error) { } } -func Prepare(sess *session.Session, imagePath string, model config.Model, mcpServers []config.MCPServer, resume bool) error { +func Prepare(sess *session.Session, imagePath string, model config.Model, resume bool) error { if imagePath == "" { imagePath = config.GuestImagePath() } @@ -224,7 +237,7 @@ func Prepare(sess *session.Session, imagePath string, model config.Model, mcpSer return fmt.Errorf("clone session disk: %w", err) } } - if err := sess.WriteGuestConfig(model, mcpServers); err != nil { + if err := sess.WriteGuestConfig(model); err != nil { return err } data, err := os.ReadFile(sess.GuestConfigJSON()) @@ -348,6 +361,8 @@ func (s *Sandbox) waitHello(ctx context.Context) error { } if hello.Protocol == 0 { s.GuestProtocol = 1 + } else if hello.Protocol > protocol.Version { + s.GuestProtocol = protocol.Version } else { s.GuestProtocol = hello.Protocol } @@ -379,7 +394,7 @@ func (s *Sandbox) readLoop() { switch { case strings.HasPrefix(frame.ID, "g-"): slots := s.guestSlots - if frame.Method == "provider_cancel" { + if protocol.GuestMethodIsCancellation(frame.Method) { slots = s.cancelSlot } select { @@ -457,9 +472,15 @@ func (s *Sandbox) failAll(err error) { pending := s.calls s.calls = map[string]*frameQueue{} turnQ := s.turnQ + turnCancel := s.turnCancel s.turnQ = nil s.activeTurn = "" + s.turnCtx = nil + s.turnCancel = nil s.mu.Unlock() + if turnCancel != nil { + turnCancel() + } for _, q := range pending { q.close(err) } @@ -472,14 +493,32 @@ func (s *Sandbox) failAll(err error) { func (s *Sandbox) dispatchGuestCall(frame protocol.Frame) { out := protocol.Frame{V: protocol.Version, ID: frame.ID} - if s.GuestProtocol < 3 { - out.Error = &protocol.Error{Code: "host", Message: fmt.Sprintf("provider broker requires protocol 3, guest speaks %d", s.GuestProtocol)} + minimum, known := protocol.GuestMethodMinVersion(frame.Method) + if !known { + out.Error = &protocol.Error{Code: "host", Message: "unknown guest method " + frame.Method} + if err := s.writeFrame(out); err != nil { + s.failConnection(err) + } + return + } + if s.GuestProtocol < minimum { + out.Error = &protocol.Error{Code: "host", Message: fmt.Sprintf("%s requires protocol %d, guest speaks %d", frame.Method, minimum, s.GuestProtocol)} + if err := s.writeFrame(out); err != nil { + s.failConnection(err) + } + return + } + if frame.Method == "request_run_command_approval" { + s.dispatchRunCommandApproval(frame, &out) if err := s.writeFrame(out); err != nil { s.failConnection(err) } return } - if s.OnGuestCall == nil { + s.mu.Lock() + handler := s.OnGuestCall + s.mu.Unlock() + if handler == nil { out.Error = &protocol.Error{Code: "host", Message: "guest calls not supported by this host"} if err := s.writeFrame(out); err != nil { s.failConnection(err) @@ -501,7 +540,7 @@ func (s *Sandbox) dispatchGuestCall(frame protocol.Frame) { } return nil } - res, perr := s.OnGuestCall.Handle(s.lifeCtx, frame.Method, frame.Params, notify) + res, perr := handler.Handle(s.lifeCtx, frame.Method, frame.Params, notify) switch { case perr != nil: out.Error = perr @@ -520,6 +559,43 @@ func (s *Sandbox) dispatchGuestCall(frame protocol.Frame) { } } +func (s *Sandbox) dispatchRunCommandApproval(frame protocol.Frame, out *protocol.Frame) { + params, err := protocol.DecodeParams[protocol.RunCommandApprovalParams](frame.Params) + if err != nil || params.TurnID == "" || len(params.Command) > protocol.MaxModelCommandBytes { + out.Error = &protocol.Error{Code: "host", Message: "invalid run_command approval request"} + return + } + s.mu.Lock() + active := s.activeTurn + turnCtx := s.turnCtx + approver := s.approver + s.mu.Unlock() + decision := protocol.ApprovalDeny + if active == params.TurnID && turnCtx != nil && approver != nil && turnCtx.Err() == nil { + got, approveErr := approver.ApproveRunCommand(turnCtx, params) + if approveErr == nil && got == protocol.ApprovalAllowOnce && turnCtx.Err() == nil { + decision = protocol.ApprovalAllowOnce + } + } + out.Result, _ = protocol.EncodeParams(protocol.RunCommandApprovalResult{Decision: decision}) +} + +func (s *Sandbox) SetGuestCallHandler(handler GuestCallHandler) { + s.mu.Lock() + s.OnGuestCall = handler + s.mu.Unlock() +} + +func (s *Sandbox) SetRunCommandApprover(approver RunCommandApprover) error { + s.mu.Lock() + defer s.mu.Unlock() + if s.activeTurn != "" { + return fmt.Errorf("cannot change command approver during a turn") + } + s.approver = approver + return nil +} + func (s *Sandbox) writeBusyReplies() { for { select { @@ -639,6 +715,9 @@ func (s *Sandbox) UserTurnCtx(ctx context.Context, text string, opts TurnOptions } func (s *Sandbox) userTurnLocked(ctx context.Context, text string, opts TurnOptions, onEvent func(protocol.AgentEvent), v2API bool) (*TurnOutcome, error) { + if s.GuestProtocol < 4 { + return nil, fmt.Errorf("%w: command approval requires protocol 4, guest speaks %d", ErrGuestTooOld, s.GuestProtocol) + } if v2API && opts.needsV2() && s.GuestProtocol < 2 { return nil, fmt.Errorf("%w: need protocol 2, guest speaks %d", ErrGuestTooOld, s.GuestProtocol) } @@ -665,7 +744,10 @@ func (s *Sandbox) userTurnLocked(ctx context.Context, text string, opts TurnOpti return nil, err } turnQ := newFrameQueue(turnQueueFrames, turnQueueBytes) + turnCtx, turnCancel := context.WithCancel(ctx) s.activeTurn = id + s.turnCtx = turnCtx + s.turnCancel = turnCancel s.turnQ = turnQ s.mu.Unlock() if err := s.writeFrame(protocol.Frame{V: protocol.Version, ID: id, Method: "user_turn", Params: raw}); err != nil { @@ -676,9 +758,12 @@ func (s *Sandbox) userTurnLocked(ctx context.Context, text string, opts TurnOpti s.mu.Lock() if s.activeTurn == id { s.activeTurn = "" + s.turnCtx = nil + s.turnCancel = nil s.turnQ = nil } s.mu.Unlock() + turnCancel() turnQ.close(context.Canceled) }() @@ -765,36 +850,23 @@ func (s *Sandbox) PushSecrets(ctx context.Context, model config.Model, secrets m if len(secrets) == 0 { return nil } + if s.GuestProtocol >= 3 { + return nil + } if s.GuestProtocol < 2 { return fmt.Errorf("guest image predates secret push; run make image") } - if s.GuestProtocol == 2 { - fmt.Fprintf(os.Stderr, "abox: guest speaks protocol 2; pushing legacy secrets (run make image to upgrade)\n") - } + fmt.Fprintf(os.Stderr, "abox: guest speaks protocol 2; pushing one legacy model credential (run make image to upgrade)\n") modelKey := model.EnvName() - rest := map[string]string{} - for k, v := range secrets { - if k != modelKey { - rest[k] = v - } + modelSecrets := map[string]string{} + if v, ok := secrets[modelKey]; ok { + modelSecrets[modelKey] = v } - var modelSecrets map[string]string - if s.GuestProtocol < 3 { - if v, ok := secrets[modelKey]; ok { - modelSecrets = map[string]string{modelKey: v} - } - } - if err := s.SetModel(ctx, model, modelSecrets); err != nil { - return err - } - if len(rest) > 0 { - return s.SetMCPTokens(ctx, rest) - } - return nil + return s.SetModel(ctx, model, modelSecrets) } func (s *Sandbox) SetMCPTokens(ctx context.Context, secrets map[string]string) error { - return s.Call(ctx, "set_mcp_tokens", protocol.SetMCPTokensParams{Secrets: secrets}, nil) + return fmt.Errorf("MCP credentials are host-brokered and cannot be sent to the guest") } func (s *Sandbox) SetModel(ctx context.Context, model config.Model, secrets map[string]string) error { diff --git a/internal/runtime/runtime_push_test.go b/internal/runtime/runtime_push_test.go index 7933e72..a1dd460 100644 --- a/internal/runtime/runtime_push_test.go +++ b/internal/runtime/runtime_push_test.go @@ -59,7 +59,7 @@ func newPipeSandbox(t *testing.T, guestProtocol int) (*Sandbox, *fakeGuest) { return s, g } -func TestPushSecretsOrderAndSplit(t *testing.T) { +func TestPushSecretsProto2SendsOnlyModelCredential(t *testing.T) { s, g := newPipeSandbox(t, 2) var order []string var mu sync.Mutex @@ -79,17 +79,6 @@ func TestPushSecretsOrderAndSplit(t *testing.T) { if len(p.Secrets) != 1 { t.Errorf("set_model carries more than the model credential: %v", p.Secrets) } - case "set_mcp_tokens": - p, err := protocol.DecodeParams[protocol.SetMCPTokensParams](frame.Params) - if err != nil { - t.Errorf("set_mcp_tokens params: %v", err) - } - if p.Secrets["ABOX_MCP_GH_TOKEN"] != "mt" { - t.Errorf("set_mcp_tokens missing mcp token: %v", p.Secrets) - } - if _, ok := p.Secrets["XAI_API_KEY"]; ok { - t.Errorf("model credential leaked into set_mcp_tokens: %v", p.Secrets) - } default: t.Errorf("unexpected method %q", frame.Method) } @@ -106,7 +95,7 @@ func TestPushSecretsOrderAndSplit(t *testing.T) { } mu.Lock() defer mu.Unlock() - if len(order) != 2 || order[0] != "set_model" || order[1] != "set_mcp_tokens" { + if len(order) != 1 || order[0] != "set_model" { t.Fatalf("order %v", order) } } @@ -114,28 +103,7 @@ func TestPushSecretsOrderAndSplit(t *testing.T) { func TestPushSecretsProto3FiltersModelCredential(t *testing.T) { s, g := newPipeSandbox(t, 3) g.onRequest = func(frame protocol.Frame, reply func(protocol.Frame)) { - switch frame.Method { - case "set_model": - p, err := protocol.DecodeParams[protocol.SetModelParams](frame.Params) - if err != nil { - t.Errorf("params: %v", err) - } - if len(p.Secrets) != 0 { - t.Errorf("proto-3 set_model must carry no secrets: %v", p.Secrets) - } - case "set_mcp_tokens": - p, _ := protocol.DecodeParams[protocol.SetMCPTokensParams](frame.Params) - if p.Secrets["ABOX_MCP_GH_TOKEN"] != "mt" { - t.Errorf("mcp token missing: %v", p.Secrets) - } - if _, ok := p.Secrets["XAI_API_KEY"]; ok { - t.Errorf("model credential leaked to proto-3 guest: %v", p.Secrets) - } - default: - t.Errorf("unexpected method %q", frame.Method) - } - ok, _ := protocol.EncodeParams(map[string]bool{"ok": true}) - reply(protocol.Frame{ID: frame.ID, Result: ok}) + t.Errorf("protocol-3 guest received secret push %q", frame.Method) } model := config.Model{Name: "grok-default", Provider: "xai", CredentialEnv: "XAI_API_KEY"} err := s.PushSecrets(context.Background(), model, map[string]string{ @@ -186,7 +154,7 @@ func TestSetModelDropsSecretsOnProto3(t *testing.T) { } func TestGuestCallMidTurn(t *testing.T) { - s, g := newPipeSandbox(t, 3) + s, g := newPipeSandbox(t, 4) turnStarted := make(chan string, 1) gotOpen := make(chan protocol.ProviderOpenResult, 1) g.onRequest = func(frame protocol.Frame, reply func(protocol.Frame)) { @@ -349,7 +317,7 @@ func TestGuestCallConcurrencyReturnsBusy(t *testing.T) { } s.startReading() for i := 0; i < regularSlots; i++ { - g.write(protocol.Frame{V: protocol.Version, ID: fmt.Sprintf("g-%d", i), Method: "hold", Params: []byte(`{}`)}) + g.write(protocol.Frame{V: protocol.Version, ID: fmt.Sprintf("g-%d", i), Method: "provider_open", Params: []byte(`{}`)}) select { case <-started: case <-time.After(2 * time.Second): @@ -365,7 +333,7 @@ func TestGuestCallConcurrencyReturnsBusy(t *testing.T) { case <-time.After(2 * time.Second): t.Fatal("provider cancellation was blocked by regular guest calls") } - g.write(protocol.Frame{V: protocol.Version, ID: "g-busy", Method: "hold", Params: []byte(`{}`)}) + g.write(protocol.Frame{V: protocol.Version, ID: "g-busy", Method: "provider_open", Params: []byte(`{}`)}) select { case frame := <-busy: if frame.ID != "g-busy" || frame.Error.Code != "busy" { @@ -386,7 +354,7 @@ func TestGuestCallContextEndsWithConnection(t *testing.T) { return nil, &protocol.Error{Code: "canceled", Message: ctx.Err().Error()} }) s.startReading() - g.write(protocol.Frame{V: protocol.Version, ID: "g-life", Method: "hold", Params: []byte(`{}`)}) + g.write(protocol.Frame{V: protocol.Version, ID: "g-life", Method: "provider_open", Params: []byte(`{}`)}) s.failConnection(errors.New("test disconnect")) select { case <-canceled: diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index fc05140..85a3b25 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -42,7 +42,7 @@ func TestPrepareResumeDoesNotClobberRoot(t *testing.T) { if err := os.WriteFile(golden, []byte("GOLDEN"), 0o600); err != nil { t.Fatal(err) } - err = Prepare(s, golden, config.Model{Name: "grok", Provider: "xai", Model: "grok-4"}, nil, true) + err = Prepare(s, golden, config.Model{Name: "grok", Provider: "xai", Model: "grok-4"}, true) if err != nil { t.Fatal(err) } @@ -70,7 +70,7 @@ func TestPrepareResumeRewritesReadOnlyConfig(t *testing.T) { if err := os.WriteFile(s.ConfigDisk(), make([]byte, 1<<20), 0o400); err != nil { t.Fatal(err) } - err = Prepare(s, "", config.Model{Name: "grok", Provider: "xai", Model: "grok-4"}, nil, true) + err = Prepare(s, "", config.Model{Name: "grok", Provider: "xai", Model: "grok-4"}, true) if err != nil { t.Fatal(err) } diff --git a/internal/runtime/runtime_turn_test.go b/internal/runtime/runtime_turn_test.go index 3f1a55e..d9f4202 100644 --- a/internal/runtime/runtime_turn_test.go +++ b/internal/runtime/runtime_turn_test.go @@ -27,7 +27,7 @@ func TestUserTurnCtxRejectsV1ForRich(t *testing.T) { func TestUserTurnPlainOnPipe(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} var wg sync.WaitGroup wg.Add(1) @@ -64,7 +64,7 @@ func TestUserTurnPlainOnPipe(t *testing.T) { func TestUserTurnDoesNotDropBurstFrames(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} const eventCount = 200 written := make(chan error, 1) @@ -129,7 +129,7 @@ func TestFrameQueueBudgets(t *testing.T) { func TestUserTurnQueueOverflowFailsClearly(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} guestDone := make(chan error, 1) go func() { @@ -180,7 +180,7 @@ func TestUserTurnQueueOverflowFailsClearly(t *testing.T) { func TestCallWriteHonorsContextUnderBackpressure(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) defer cancel() done := make(chan error, 1) @@ -200,7 +200,7 @@ func TestCallWriteHonorsContextUnderBackpressure(t *testing.T) { func TestUserTurnCtxCancelWritesCancelTurn(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} ctx, cancel := context.WithCancel(context.Background()) done := make(chan string, 1) @@ -238,7 +238,7 @@ func TestUserTurnCtxCancelWritesCancelTurn(t *testing.T) { func TestUserTurnCtxDeadlineWaitsForCanceledResponse(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} guestDone := make(chan error, 1) go func() { @@ -281,7 +281,7 @@ func TestUserTurnCtxDeadlineWaitsForCanceledResponse(t *testing.T) { func TestUserTurnCtxForwardsOptionsAndResult(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} guestDone := make(chan error, 1) go func() { @@ -331,7 +331,7 @@ func TestUserTurnCtxForwardsOptionsAndResult(t *testing.T) { func TestCallSkipsLateCancelResponse(t *testing.T) { host, guest := net.Pipe() t.Cleanup(func() { host.Close(); guest.Close() }) - s := &Sandbox{conn: host, GuestProtocol: 2} + s := &Sandbox{conn: host, GuestProtocol: 4} var guestWriteMu sync.Mutex writeGuest := func(frame protocol.Frame) error { guestWriteMu.Lock() diff --git a/internal/session/scrub.go b/internal/session/scrub.go index 5f9def5..4659f13 100644 --- a/internal/session/scrub.go +++ b/internal/session/scrub.go @@ -12,8 +12,8 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" ) -// ScrubSecrets removes the "secrets" key from guest-config.json and -// config.raw. It never deletes sessions. +// ScrubSecrets removes legacy credentials and MCP endpoint policy from guest +// configuration artifacts. It never deletes sessions. func ScrubSecrets(root string) (int, error) { entries, err := os.ReadDir(root) if err != nil { @@ -144,10 +144,13 @@ func scrubbedJSONObject(data []byte) ([]byte, error) { if obj == nil { return nil, fmt.Errorf("expected JSON object") } - if _, ok := obj["secrets"]; !ok { + _, hasSecrets := obj["secrets"] + _, hasMCP := obj["mcp_servers"] + if !hasSecrets && !hasMCP { return nil, nil } delete(obj, "secrets") + delete(obj, "mcp_servers") out, err := json.MarshalIndent(obj, "", " ") if err != nil { return nil, fmt.Errorf("scrub: %w", err) diff --git a/internal/session/session.go b/internal/session/session.go index 6439417..56a28e3 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -172,23 +172,13 @@ func WriteTranscript(path string, lines []string) error { return os.WriteFile(path, data, 0o600) } -func (s *Session) WriteGuestConfig(model config.Model, servers []config.MCPServer) error { - var gs []protocol.GuestMCPServer - for _, srv := range servers { - gs = append(gs, protocol.GuestMCPServer{ - Name: srv.Name, - URL: srv.URL, - TokenEnv: config.TokenEnv(srv), - Allowlist: srv.ToolAllowlist, - }) - } +func (s *Session) WriteGuestConfig(model config.Model) error { cfg := protocol.GuestConfig{ SessionID: s.ID, Capability: s.Capability, VsockPort: protocol.RPCPort, RepoDir: protocol.GuestRepoDir, Model: model.ToGuest(), - MCPServers: gs, } data, err := json.MarshalIndent(cfg, "", " ") if err != nil { diff --git a/internal/session/session_test.go b/internal/session/session_test.go index 0b25df9..8b89df9 100644 --- a/internal/session/session_test.go +++ b/internal/session/session_test.go @@ -10,15 +10,13 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" ) -func TestWriteGuestConfigIncludesMCPNoSecrets(t *testing.T) { +func TestWriteGuestConfigExcludesMCPAndSecrets(t *testing.T) { t.Setenv("HOME", t.TempDir()) s, err := Create("/repo", "deadbeef") if err != nil { t.Fatal(err) } - err = s.WriteGuestConfig(config.Model{Name: "grok"}, []config.MCPServer{ - {Name: "gh", URL: "https://api.githubcopilot.com/mcp/", CredentialEnv: "ABOX_MCP_GH_TOKEN"}, - }) + err = s.WriteGuestConfig(config.Model{Name: "grok"}) if err != nil { t.Fatal(err) } @@ -27,8 +25,8 @@ func TestWriteGuestConfigIncludesMCPNoSecrets(t *testing.T) { t.Fatal(err) } body := string(data) - if !strings.Contains(body, "api.githubcopilot.com") || !strings.Contains(body, "ABOX_MCP_GH_TOKEN") { - t.Fatalf("guest config missing mcp: %s", body) + if strings.Contains(body, "mcp_servers") || strings.Contains(body, "api.githubcopilot.com") || strings.Contains(body, "ABOX_MCP_GH_TOKEN") { + t.Fatalf("guest config leaked MCP policy: %s", body) } if strings.Contains(body, `"secrets"`) { t.Fatalf("guest config leaked a secrets key: %s", body) diff --git a/internal/tui/approval_test.go b/internal/tui/approval_test.go new file mode 100644 index 0000000..52fda31 --- /dev/null +++ b/internal/tui/approval_test.go @@ -0,0 +1,154 @@ +package tui + +import ( + "context" + "strconv" + "strings" + "testing" + "time" + + tea "charm.land/bubbletea/v2" + + "github.com/AdminTurnedDevOps/ABox/internal/config" + "github.com/AdminTurnedDevOps/ABox/protocol" +) + +func TestApprovalDefaultsToDenyAndShowsExactRequest(t *testing.T) { + command := "printf 'first\\nsecond' && " + strings.Repeat("x", 120) + "-END" + req := approvalRequest(context.Background(), protocol.RunCommandApprovalParams{ + Command: command, WorkDir: "src/path with space", TimeoutSec: 37, + }) + m := New(config.Defaults(), config.Model{}, nil, nil, "ready", nil, "") + m.width = 40 + + updated, _ := m.Update(approvalMsg{req: req}) + m = updated.(model) + if m.mode != modeApproval || m.approvalAllow { + t.Fatalf("mode=%v allow=%v", m.mode, m.approvalAllow) + } + view := m.View().Content + for _, want := range []string{strconv.Quote(command), `guest workdir: "src/path with space"`, "timeout: 37s", "[Deny]"} { + if !strings.Contains(view, want) { + t.Fatalf("approval view missing %q:\n%s", want, view) + } + } + + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyEnter}) + m = updated.(model) + if got := <-req.response; got != protocol.ApprovalDeny { + t.Fatalf("decision=%q", got) + } +} + +func TestApprovalSelectionAllowsOnce(t *testing.T) { + req := approvalRequest(context.Background(), protocol.RunCommandApprovalParams{Command: "go test ./..."}) + m := New(config.Defaults(), config.Model{}, nil, nil, "ready", nil, "") + m.mode = modeApproval + m.approvalReq = req + + updated, _ := m.Update(tea.KeyPressMsg{Code: tea.KeyRight}) + m = updated.(model) + if !m.approvalAllow { + t.Fatal("right did not select allow once") + } + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyLeft}) + m = updated.(model) + if m.approvalAllow { + t.Fatal("left did not select deny") + } + updated, _ = m.Update(tea.KeyPressMsg{Code: 'j', Text: "j"}) + m = updated.(model) + if !m.approvalAllow { + t.Fatal("j did not select allow once") + } + updated, _ = m.Update(tea.KeyPressMsg{Code: 'k', Text: "k"}) + m = updated.(model) + if m.approvalAllow { + t.Fatal("k did not select deny") + } + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyRight}) + m = updated.(model) + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyEnter}) + m = updated.(model) + if got := <-req.response; got != protocol.ApprovalAllowOnce { + t.Fatalf("decision=%q", got) + } + if m.mode != modeChat || m.approvalReq != nil { + t.Fatalf("mode=%v request=%v", m.mode, m.approvalReq) + } +} + +func TestApprovalEscapeDeniesAndListensAgain(t *testing.T) { + first := approvalRequest(context.Background(), protocol.RunCommandApprovalParams{Command: "first"}) + m := New(config.Defaults(), config.Model{}, nil, nil, "ready", nil, "") + m.mode = modeApproval + m.approvalReq = first + + updated, next := m.Update(tea.KeyPressMsg{Code: tea.KeyEscape}) + m = updated.(model) + if got := <-first.response; got != protocol.ApprovalDeny { + t.Fatalf("escape decision=%q", got) + } + if next == nil { + t.Fatal("escape did not resume approval listening") + } + + msgCh := make(chan tea.Msg, 1) + go func() { msgCh <- next() }() + decisionCh := make(chan protocol.ApprovalDecision, 1) + go func() { + decision, _ := m.approveRunCommand(context.Background(), protocol.RunCommandApprovalParams{Command: "second"}) + decisionCh <- decision + }() + + var msg tea.Msg + select { + case msg = <-msgCh: + case <-time.After(time.Second): + t.Fatal("later approval was not delivered") + } + updated, _ = m.Update(msg) + m = updated.(model) + updated, _ = m.Update(tea.KeyPressMsg{Code: tea.KeyEnter}) + if got := <-decisionCh; got != protocol.ApprovalDeny { + t.Fatalf("later decision=%q", got) + } +} + +func TestApprovalBridgeHonorsCancellation(t *testing.T) { + m := model{approvals: make(chan *runCommandApprovalRequest)} + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { + decision, err := m.approveRunCommand(ctx, protocol.RunCommandApprovalParams{Command: "sleep 10"}) + if decision != protocol.ApprovalDeny { + result <- &unexpectedDecisionError{decision: decision} + return + } + result <- err + }() + req := <-m.approvals + if cap(req.response) != 1 { + t.Fatalf("response channel capacity=%d", cap(req.response)) + } + cancel() + if err := <-result; err != context.Canceled { + t.Fatalf("error=%v", err) + } +} + +type unexpectedDecisionError struct { + decision protocol.ApprovalDecision +} + +func (e *unexpectedDecisionError) Error() string { + return "unexpected approval decision " + string(e.decision) +} + +func approvalRequest(ctx context.Context, params protocol.RunCommandApprovalParams) *runCommandApprovalRequest { + return &runCommandApprovalRequest{ + ctx: ctx, params: params, + response: make(chan protocol.ApprovalDecision, 1), + settled: make(chan struct{}), + } +} diff --git a/internal/tui/commands_test.go b/internal/tui/commands_test.go index d8a045a..80e09bf 100644 --- a/internal/tui/commands_test.go +++ b/internal/tui/commands_test.go @@ -9,7 +9,7 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/internal/credsource" - "github.com/AdminTurnedDevOps/ABox/internal/runtime" + "github.com/AdminTurnedDevOps/ABox/internal/hostbroker" "github.com/AdminTurnedDevOps/ABox/protocol" ) @@ -183,13 +183,20 @@ func TestApplyCloudCredentialAddsMissingProfile(t *testing.T) { func TestUpdateHostBrokerUsesCurrentModelConfig(t *testing.T) { cfg := config.Defaults() - cfg.Models = []config.Model{{ + initial := cfg.Models[0] + updated := config.Model{ Name: "updated", Provider: "openai", Model: "gpt-current", CredentialEnv: "CURRENT_API_KEY", - }} - sb := &runtime.Sandbox{} - m := model{sandbox: sb, resolver: credsource.NewResolver()} - t.Cleanup(func() { _ = m.resolver.Close() }) - m.updateHostBroker(cfg) + } + cfg.Models = []config.Model{updated} + resolver := credsource.NewResolver() + t.Cleanup(func() { _ = resolver.Close() }) + broker, err := hostbroker.New(config.Defaults(), initial, resolver) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = broker.Close() }) + m := model{hostBroker: broker} + m.updateHostBroker(cfg, updated) raw, err := json.Marshal(protocol.ProviderOpenParams{Model: "updated"}) if err != nil { @@ -197,7 +204,7 @@ func TestUpdateHostBrokerUsesCurrentModelConfig(t *testing.T) { } ctx, cancel := context.WithCancel(context.Background()) defer cancel() - if _, perr := sb.OnGuestCall.Handle(ctx, "provider_open", raw, nil); perr != nil { + if _, perr := broker.Handle(ctx, "provider_open", raw, nil); perr != nil { t.Fatalf("updated broker rejected current model: %v", perr) } } diff --git a/internal/tui/tui.go b/internal/tui/tui.go index 67f4636..64e5eda 100644 --- a/internal/tui/tui.go +++ b/internal/tui/tui.go @@ -3,6 +3,7 @@ package tui import ( "context" "fmt" + "strconv" "strings" "time" @@ -14,7 +15,7 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/internal/credsource" - "github.com/AdminTurnedDevOps/ABox/internal/llmbroker" + "github.com/AdminTurnedDevOps/ABox/internal/hostbroker" "github.com/AdminTurnedDevOps/ABox/internal/runtime" "github.com/AdminTurnedDevOps/ABox/internal/session" "github.com/AdminTurnedDevOps/ABox/protocol" @@ -31,12 +32,14 @@ const ( modeCredSourcePick modeCredModelPick modeCredName + modeApproval ) type model struct { cfg config.File sel config.Model sandbox *runtime.Sandbox + hostBroker *hostbroker.Broker ta textarea.Model keyIn textinput.Model mode uiMode @@ -60,11 +63,23 @@ type model struct { selKeyStatus string provKeyStatus map[string]string mcpKeyStatus map[string]string + approvalReq *runCommandApprovalRequest + approvalAllow bool + approvals chan *runCommandApprovalRequest } type evMsg protocol.AgentEvent type errMsg error type doneMsg struct{} +type approvalMsg struct{ req *runCommandApprovalRequest } +type approvalCanceledMsg struct{ req *runCommandApprovalRequest } + +type runCommandApprovalRequest struct { + ctx context.Context + params protocol.RunCommandApprovalParams + response chan protocol.ApprovalDecision + settled chan struct{} +} // Presence is cached: the render path must not shell out to keychain or HTTP. type credStatusMsg struct { @@ -74,7 +89,7 @@ type credStatusMsg struct { partial bool } -func New(cfg config.File, sel config.Model, sb *runtime.Sandbox, vmState string, log []string, transcriptPath string) model { +func New(cfg config.File, sel config.Model, sb *runtime.Sandbox, broker *hostbroker.Broker, vmState string, log []string, transcriptPath string) model { ta := textarea.New() ta.Placeholder = "Ask ABox Anything" ta.Focus() @@ -89,11 +104,15 @@ func New(cfg config.File, sel config.Model, sb *runtime.Sandbox, vmState string, ki.EchoCharacter = '•' ki.Placeholder = "paste API key" ki.Prompt = "key> " - return model{cfg: cfg, sel: sel, sandbox: sb, ta: ta, keyIn: ki, vmState: vmState, log: log, transcriptPath: transcriptPath} + return model{ + cfg: cfg, sel: sel, sandbox: sb, hostBroker: broker, ta: ta, keyIn: ki, + vmState: vmState, log: log, transcriptPath: transcriptPath, + approvals: make(chan *runCommandApprovalRequest), + } } func (m model) Init() tea.Cmd { - return tea.Batch(textarea.Blink, checkCredStatus(m.cfg, m.sel, m.resolver, false)) + return tea.Batch(textarea.Blink, checkCredStatus(m.cfg, m.sel, m.resolver, false), waitApproval(m.approvals)) } func checkCredStatus(cfg config.File, sel config.Model, r *credsource.Resolver, mcp bool) tea.Cmd { @@ -140,6 +159,13 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case tea.KeyPressMsg: switch msg.String() { case "ctrl+c": + if m.mode == modeApproval { + m.resolveApproval(protocol.ApprovalDeny) + if m.cancel != nil { + m.cancel() + } + return m, tea.Quit + } if m.mode != modeChat { m.cancelCredInput() m.mode = modeChat @@ -155,6 +181,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case "ctrl+d": return m, tea.Quit case "esc": + if m.mode == modeApproval { + m.resolveApproval(protocol.ApprovalDeny) + return m, waitApproval(m.approvals) + } if m.mode != modeChat { m.cancelCredInput() m.mode = modeChat @@ -210,6 +240,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return m, nil } case "k": + if m.mode == modeApproval { + m.approvalAllow = false + return m, nil + } if m.mode == modeMCPPick { if m.mcpSel > 0 { m.mcpSel-- @@ -229,6 +263,10 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return m, nil } case "j": + if m.mode == modeApproval { + m.approvalAllow = true + return m, nil + } if m.mode == modeMCPPick { if m.mcpSel < len(mcpServers(m.cfg))-1 { m.mcpSel++ @@ -247,7 +285,25 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { } return m, nil } + case "left": + if m.mode == modeApproval { + m.approvalAllow = false + return m, nil + } + case "right": + if m.mode == modeApproval { + m.approvalAllow = true + return m, nil + } case "enter", "ctrl+m": + if m.mode == modeApproval { + decision := protocol.ApprovalDeny + if m.approvalAllow { + decision = protocol.ApprovalAllowOnce + } + m.resolveApproval(decision) + return m, waitApproval(m.approvals) + } if m.busy { return m, nil } @@ -316,6 +372,22 @@ func (m model) Update(msg tea.Msg) (tea.Model, tea.Cmd) { m.busy = false m.saveTranscript() return m, nil + case approvalMsg: + if msg.req.ctx.Err() != nil { + m.resolveRequest(msg.req, protocol.ApprovalDeny) + return m, waitApproval(m.approvals) + } + m.approvalReq = msg.req + m.approvalAllow = false + m.mode = modeApproval + m.ta.Blur() + return m, waitApprovalCancellation(msg.req) + case approvalCanceledMsg: + if m.approvalReq == msg.req { + m.resolveApproval(protocol.ApprovalDeny) + return m, waitApproval(m.approvals) + } + return m, nil } if m.mode == modeProviderKey || m.mode == modeMCPKey || m.mode == modeCredName { var cmd tea.Cmd @@ -359,7 +431,12 @@ func (m model) submit() (tea.Model, tea.Cmd) { ch := make(chan protocol.AgentEvent, 32) m.events = ch go func() { - _ = m.sandbox.UserTurn(ctx, text, func(e protocol.AgentEvent) { ch <- e }) + _, _ = m.sandbox.UserTurnCtx(ctx, text, runtime.TurnOptions{}, func(e protocol.AgentEvent) { + select { + case ch <- e: + case <-ctx.Done(): + } + }) close(ch) }() return m, waitEvent(ch) @@ -470,12 +547,12 @@ func (m model) saveCloudCredential() (tea.Model, tea.Cmd) { } m.cfg = cfg if m.sandbox != nil { - m.updateHostBroker(cfg) if err := m.sandbox.SetModel(context.Background(), sel, nil); err != nil { m.err = "saved on host but guest agent update failed: " + err.Error() return m, nil } } + m.updateHostBroker(cfg, sel) m.sel = sel m.err = "" m.log = append(m.log, "credential "+m.provPick.Label+" -> "+m.credSource.Label+" "+name+" ("+note+")") @@ -530,11 +607,18 @@ func (m model) saveMCPKey() (tea.Model, tea.Cmd) { } m.cfg = cfg m.mcpKeyStatus = map[string]string{m.mcpPick.Name: "key ok"} - if m.sandbox != nil { - if err := m.sandbox.SetMCPTokens(context.Background(), map[string]string{env: key}); err != nil { - m.err = "saved on host but guest MCP update failed: " + err.Error() + if m.hostBroker != nil { + if err := m.hostBroker.UpdateMCP(cfg); err != nil { + m.err = "saved on host but host MCP update failed: " + err.Error() return m, nil } + if err := m.hostBroker.SetMCPTokens(map[string]string{env: key}); err != nil { + m.err = "saved on host but host MCP token refresh failed: " + err.Error() + return m, nil + } + } else if m.sandbox != nil { + m.err = "saved on host but host MCP broker is unavailable" + return m, nil } m.err = "" m.log = append(m.log, "mcp "+m.mcpPick.Name+" token saved ("+note+") (OAuth: abox mcp login "+m.mcpPick.Name+")") @@ -559,14 +643,13 @@ func (m model) saveProviderKey() (tea.Model, tea.Cmd) { } m.cfg = cfg if m.sandbox != nil { - // Refresh the broker even if the guest update fails so it is not stuck on startup config. - m.updateHostBroker(cfg) secrets := map[string]string{m.provPick.Env: key} if err := m.sandbox.SetModel(context.Background(), sel, secrets); err != nil { m.err = "saved on host but guest agent update failed: " + err.Error() return m, nil } } + m.updateHostBroker(cfg, sel) m.sel = sel m.selKeyStatus = "key ok" if m.provKeyStatus == nil { @@ -579,9 +662,75 @@ func (m model) saveProviderKey() (tea.Model, tea.Cmd) { return m, nil } -func (m model) updateHostBroker(cfg config.File) { - if m.sandbox != nil { - m.sandbox.OnGuestCall = llmbroker.New(cfg, m.resolver) +func (m model) updateHostBroker(cfg config.File, sel config.Model) { + if m.hostBroker != nil { + m.hostBroker.UpdateModel(cfg, sel) + } +} + +func (m model) approveRunCommand(ctx context.Context, params protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) { + req := &runCommandApprovalRequest{ + ctx: ctx, + params: params, + response: make(chan protocol.ApprovalDecision, 1), + settled: make(chan struct{}), + } + select { + case m.approvals <- req: + case <-ctx.Done(): + return protocol.ApprovalDeny, ctx.Err() + } + select { + case decision := <-req.response: + if ctx.Err() != nil { + return protocol.ApprovalDeny, ctx.Err() + } + return decision, nil + case <-ctx.Done(): + return protocol.ApprovalDeny, ctx.Err() + } +} + +func (m *model) resolveApproval(decision protocol.ApprovalDecision) { + if m.approvalReq == nil { + return + } + if m.approvalReq.ctx.Err() != nil { + decision = protocol.ApprovalDeny + } + m.resolveRequest(m.approvalReq, decision) + m.approvalReq = nil + m.approvalAllow = false + m.mode = modeChat + m.ta.Focus() +} + +func (m *model) resolveRequest(req *runCommandApprovalRequest, decision protocol.ApprovalDecision) { + select { + case req.response <- decision: + default: + } + select { + case <-req.settled: + default: + close(req.settled) + } +} + +func waitApproval(ch <-chan *runCommandApprovalRequest) tea.Cmd { + return func() tea.Msg { + return approvalMsg{req: <-ch} + } +} + +func waitApprovalCancellation(req *runCommandApprovalRequest) tea.Cmd { + return func() tea.Msg { + select { + case <-req.ctx.Done(): + return approvalCanceledMsg{req: req} + case <-req.settled: + return nil + } } } @@ -727,6 +876,21 @@ func (m model) View() tea.View { composer = b.String() case modeMCPKey: composer = "Bearer token for " + m.mcpPick.Name + "\n" + m.keyIn.View() + case modeApproval: + if m.approvalReq != nil { + workdir := m.approvalReq.params.WorkDir + if workdir == "" { + workdir = "." + } + deny, allow := "[Deny]", " Allow once " + if m.approvalAllow { + deny, allow = " Deny ", "[Allow once]" + } + composer = fmt.Sprintf( + "Approve run_command?\ncommand: %s\nguest workdir: %s\ntimeout: %ds\n\n%s %s\nleft/right or j/k select; enter confirms; esc denies", + strconv.Quote(m.approvalReq.params.Command), strconv.Quote(workdir), m.approvalReq.params.TimeoutSec, deny, allow, + ) + } default: if m.showingSlash() { var b strings.Builder @@ -820,12 +984,21 @@ func max(a, b int) int { return b } -func Run(cfg config.File, sel config.Model, sb *runtime.Sandbox, vmState string, log []string, resolver *credsource.Resolver, transcriptPath string) error { +func Run(cfg config.File, sel config.Model, sb *runtime.Sandbox, broker *hostbroker.Broker, vmState string, log []string, resolver *credsource.Resolver, transcriptPath string) error { if resolver == nil { resolver = credsource.NewResolver() } - m := New(cfg, sel, sb, vmState, log, transcriptPath) + m := New(cfg, sel, sb, broker, vmState, log, transcriptPath) m.resolver = resolver + if sb != nil { + if broker == nil { + return fmt.Errorf("host broker is required when the VM is ready") + } + sb.SetGuestCallHandler(broker) + if err := sb.SetRunCommandApprover(runtime.RunCommandApproverFunc(m.approveRunCommand)); err != nil { + return fmt.Errorf("configure run_command approval: %w", err) + } + } p := tea.NewProgram(m) _, err := p.Run() return err diff --git a/pkg/abox/abox.go b/pkg/abox/abox.go index 08c5fb8..51c1761 100644 --- a/pkg/abox/abox.go +++ b/pkg/abox/abox.go @@ -11,7 +11,7 @@ import ( "github.com/AdminTurnedDevOps/ABox/internal/config" "github.com/AdminTurnedDevOps/ABox/internal/credsource" - "github.com/AdminTurnedDevOps/ABox/internal/llmbroker" + "github.com/AdminTurnedDevOps/ABox/internal/hostbroker" "github.com/AdminTurnedDevOps/ABox/internal/repository" "github.com/AdminTurnedDevOps/ABox/internal/runtime" "github.com/AdminTurnedDevOps/ABox/internal/session" @@ -54,10 +54,6 @@ func Resume(ctx context.Context, sessionID string, opts Options) (*Session, erro func open(ctx context.Context, opts Options, resume bool, resumeID string) (*Session, error) { opts = opts.withDefaults() - cfg, cfgPath, err := config.Load() - if err != nil { - return nil, fmt.Errorf("load config: %w", err) - } n, scrubErr := session.ScrubSecretsEverywhere() if n > 0 { fmt.Fprintf(os.Stderr, "abox: scrubbed plaintext secrets from %d old session(s)\n", n) @@ -65,6 +61,10 @@ func open(ctx context.Context, opts Options, resume bool, resumeID string) (*Ses if scrubErr != nil { return nil, fmt.Errorf("scrub legacy session secrets: %w", scrubErr) } + cfg, cfgPath, err := config.Load() + if err != nil { + return nil, fmt.Errorf("load config: %w", err) + } resolver := credsource.NewResolver() sel, ok := cfg.ModelNamed(opts.Model) if !ok { @@ -110,12 +110,7 @@ func open(ctx context.Context, opts Options, resume bool, resumeID string) (*Ses if image == "" { image = cfg.Runtime.Image } - mcpServers, err := cfg.ResolvedMCPServers() - if err != nil { - resolver.Close() - return nil, err - } - if err := runtime.Prepare(sess, image, sel, mcpServers, resume); err != nil { + if err := runtime.Prepare(sess, image, sel, resume); err != nil { resolver.Close() return nil, err } @@ -149,18 +144,18 @@ func open(ctx context.Context, opts Options, resume bool, resumeID string) (*Ses } return nil, fmt.Errorf("%w: protocol-1 guest cannot use the secretless config; rebuild the guest image", ErrGuestTooOld) } - sb.OnGuestCall = llmbroker.New(cfg, resolver) - - // Push whatever resolved before returning a partial-resolution error. - pushCtx, pushCancel := context.WithTimeout(ctx, 30*time.Second) - defer pushCancel() - secrets, resolveErr := credsource.ResolveSelected(pushCtx, resolver, cfg, sel) - pushErr := sb.PushSecrets(pushCtx, sel, secrets) - if err := credentialStartupError(resolveErr, pushErr); err != nil { + if sb.GuestProtocol < 4 { + sb.Stop() + resolver.Close() + return nil, fmt.Errorf("%w: guest protocol %d cannot enforce brokered MCP and command approvals; rebuild the guest image and start a new session", ErrGuestTooOld, sb.GuestProtocol) + } + broker, err := hostbroker.New(cfg, sel, resolver) + if err != nil { sb.Stop() resolver.Close() return nil, err } + sb.SetGuestCallHandler(broker) if !resume { archive, err := repository.ArchiveHEAD(snap.Root) if err != nil { @@ -174,7 +169,7 @@ func open(ctx context.Context, opts Options, resume bool, resumeID string) (*Ses return nil, fmt.Errorf("transfer repo: %w", err) } } - return &Session{cfg: cfg, sess: sess, sb: sb, sel: sel, resolver: resolver}, nil + return &Session{cfg: cfg, sess: sess, sb: sb, sel: sel, resolver: resolver, broker: broker}, nil } func credentialStartupError(resolveErr, pushErr error) error { @@ -210,6 +205,33 @@ type Capabilities struct { Cancel bool RichEvents bool TurnOptions bool + Approvals bool + MCPBroker bool +} + +type ApprovalDecision uint8 + +const ( + ApprovalDeny ApprovalDecision = iota + ApprovalAllowOnce +) + +type ApprovalRequest struct { + Tool string + ToolID string + Command string + WorkDir string + TimeoutSec int +} + +type Approver interface { + Approve(context.Context, ApprovalRequest) (ApprovalDecision, error) +} + +type ApproverFunc func(context.Context, ApprovalRequest) (ApprovalDecision, error) + +func (f ApproverFunc) Approve(ctx context.Context, request ApprovalRequest) (ApprovalDecision, error) { + return f(ctx, request) } type TurnOpts struct { @@ -231,6 +253,7 @@ type Session struct { sb *runtime.Sandbox sel config.Model resolver *credsource.Resolver + broker *hostbroker.Broker } func (s *Session) ID() string { return s.sess.ID } @@ -242,6 +265,8 @@ func (s *Session) Capabilities() Capabilities { Cancel: p >= 2, RichEvents: p >= 2, TurnOptions: p >= 2, + Approvals: p >= 4, + MCPBroker: p >= 4, } } @@ -324,15 +349,37 @@ func (s *Session) SetModel(ctx context.Context, model string) error { if err := s.sb.SetModel(ctx, sel, secrets); err != nil { return err } - // Broker snapshots config at construction; replace it after the guest accepts the model. - s.sb.OnGuestCall = llmbroker.New(cfg, s.resolver) + s.broker.UpdateModel(cfg, sel) s.cfg = cfg s.sel = sel return nil } func (s *Session) SetMCPTokens(ctx context.Context, secrets map[string]string) error { - return s.sb.SetMCPTokens(ctx, secrets) + if err := ctx.Err(); err != nil { + return err + } + return s.broker.SetMCPTokens(secrets) +} + +func (s *Session) SetApprover(approver Approver) error { + if !s.mu.TryLock() { + return fmt.Errorf("cannot set approver while a turn or model update is in progress") + } + defer s.mu.Unlock() + if approver == nil { + return s.sb.SetRunCommandApprover(nil) + } + return s.sb.SetRunCommandApprover(runtime.RunCommandApproverFunc(func(ctx context.Context, params protocol.RunCommandApprovalParams) (protocol.ApprovalDecision, error) { + decision, err := approver.Approve(ctx, ApprovalRequest{ + Tool: "run_command", ToolID: params.ToolID, Command: params.Command, + WorkDir: params.WorkDir, TimeoutSec: params.TimeoutSec, + }) + if err != nil || decision != ApprovalAllowOnce { + return protocol.ApprovalDeny, err + } + return protocol.ApprovalAllowOnce, nil + })) } func (s *Session) ExportPatch(ctx context.Context) (patch, summary string, err error) { @@ -364,6 +411,9 @@ func (s *Session) Close() error { return nil } err := s.sb.Stop() + if s.broker != nil { + err = errors.Join(err, s.broker.Close()) + } s.resolver.Close() return err } diff --git a/protocol/protocol.go b/protocol/protocol.go index 8dc1b13..ee19bec 100644 --- a/protocol/protocol.go +++ b/protocol/protocol.go @@ -10,7 +10,7 @@ import ( ) const ( - Version = 3 // host provider broker; LLM credentials stay on the host + Version = 4 // host provider/MCP brokers and model-command approval MaxFrameBytes = 1 << 20 MaxArchiveChunk = 256 << 10 @@ -27,6 +27,12 @@ const ( MaxProviderEvents = 1 << 20 // events in one provider stream MaxProviderStreams = 2 // concurrent provider streams per session MaxGuestCalls = 8 // concurrent guest-initiated host RPCs + + MaxModelCommandBytes = 16 << 10 + MaxMCPTools = MaxProviderTools - 5 + MaxMCPSchemaBytes = 64 << 10 + MaxMCPArgsBytes = 512 << 10 + MaxMCPResultBytes = 512 << 10 ) // Frame is a length-prefixed JSON message. @@ -148,6 +154,25 @@ type RunCommandResult struct { Trunc bool `json:"truncated"` } +type ApprovalDecision string + +const ( + ApprovalDeny ApprovalDecision = "deny" + ApprovalAllowOnce ApprovalDecision = "allow_once" +) + +type RunCommandApprovalParams struct { + TurnID string `json:"turn_id"` + ToolID string `json:"tool_id,omitempty"` + Command string `json:"command"` + WorkDir string `json:"workdir,omitempty"` + TimeoutSec int `json:"timeout_sec,omitempty"` +} + +type RunCommandApprovalResult struct { + Decision ApprovalDecision `json:"decision"` +} + type ArchiveChunkParams struct { Offset int64 `json:"offset"` Last bool `json:"last"` @@ -216,6 +241,37 @@ type SetMCPTokensParams struct { Secrets map[string]string `json:"secrets"` } +type MCPTool struct { + Server string `json:"server"` + Name string `json:"name"` + Prefixed string `json:"prefixed"` + Description string `json:"description,omitempty"` + Parameters map[string]any `json:"parameters"` +} + +type MCPListParams struct{} + +type MCPListResult struct { + Tools []MCPTool `json:"tools"` +} + +type MCPCallParams struct { + CallID string `json:"call_id"` + Server string `json:"server"` + Tool string `json:"tool"` + Arguments json.RawMessage `json:"arguments,omitempty"` +} + +type MCPCallResult struct { + Text string `json:"text,omitempty"` + IsError bool `json:"is_error,omitempty"` + Truncated bool `json:"truncated,omitempty"` +} + +type MCPCancelParams struct { + CallID string `json:"call_id"` +} + // Model is a configured alias, never a URL, header, or credential name. type ProviderOpenParams struct { Model string `json:"model"` @@ -372,4 +428,21 @@ func DecodeParams[T any](raw json.RawMessage) (T, error) { return v, err } +// GuestMethodMinVersion returns the minimum negotiated protocol for a +// guest-initiated host RPC. Unknown methods are rejected by the runtime. +func GuestMethodMinVersion(method string) (int, bool) { + switch method { + case "provider_open", "provider_send", "provider_cancel": + return 3, true + case "mcp_list", "mcp_call", "mcp_cancel", "request_run_command_approval": + return 4, true + default: + return 0, false + } +} + +func GuestMethodIsCancellation(method string) bool { + return method == "provider_cancel" || method == "mcp_cancel" +} + const DefaultRPCTimeout = 60 * time.Second