Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 20 additions & 47 deletions cmd/abox-guest/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,6 @@ import (
"fmt"
"io"
"net"
"net/url"
"os"
"strings"
"sync"
Expand All @@ -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"
Expand All @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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}
Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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 {
Expand All @@ -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 {
Expand Down
4 changes: 2 additions & 2 deletions cmd/abox-vmm/start_darwin_arm64.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
43 changes: 28 additions & 15 deletions cmd/abox/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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:])
}
Expand Down Expand Up @@ -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()

Expand Down Expand Up @@ -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
}
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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) {
Expand Down
Loading
Loading