From fbf4c460f651940557f294b2c5822ab7b158c3fb Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Thu, 20 Aug 2026 08:12:21 +0530 Subject: [PATCH] feat(acp): add ACP client and external subagent providers (DSH 2.8) Port of DSH acp client and subagent-acp, subagent-claude-code, subagent-codex providers. - acp: Implemented JSON-RPC 2.0 stdio client with initialize handshake, session/new, session/prompt, session/cancel, and streamed update notifications. - acp: Connected server-to-client permission requests via OnPermissionRequest callback matching ACP outcome schema. - acp: Handled process group lifecycle and disposal across Unix and Windows. - multiagent: Added SubagentProvider registry with capability guards (outputSchema, maxDepth, supportedPersonas). - multiagent: Added built-in ACP providers for generic acp, claude-code, and codex. - tests: Full test coverage for ACP client-server round-trip, permission routing, capability rejections, structured output parsing, and lifecycle disposal. --- internal/acp/client.go | 467 +++++++++++++++++++++++++++ internal/acp/client_test.go | 164 ++++++++++ internal/acp/proc_unix.go | 20 ++ internal/acp/proc_windows.go | 19 ++ internal/multiagent/provider.go | 284 ++++++++++++++++ internal/multiagent/provider_test.go | 152 +++++++++ 6 files changed, 1106 insertions(+) create mode 100644 internal/acp/client.go create mode 100644 internal/acp/client_test.go create mode 100644 internal/acp/proc_unix.go create mode 100644 internal/acp/proc_windows.go create mode 100644 internal/multiagent/provider.go create mode 100644 internal/multiagent/provider_test.go diff --git a/internal/acp/client.go b/internal/acp/client.go new file mode 100644 index 00000000..61f13789 --- /dev/null +++ b/internal/acp/client.go @@ -0,0 +1,467 @@ +package acp + +import ( + "bufio" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os/exec" + "strconv" + "sync" + "time" +) + +// ErrClientClosed is returned when an operation is attempted on a closed ACP client. +var ErrClientClosed = errors.New("acp: client is closed") + +// PermissionRequest represents a server-initiated tool execution approval request. +type PermissionRequest struct { + SessionID string `json:"sessionId"` + ToolName string `json:"toolName"` + Arguments json.RawMessage `json:"arguments,omitempty"` + Summary string `json:"summary,omitempty"` +} + +// PromptResult represents the outcome of an ACP session/prompt execution. +type PromptResult struct { + Status string `json:"status"` + Output string `json:"output,omitempty"` +} + +// ClientOptions configures the ACP client. +type ClientOptions struct { + Timeout time.Duration + Env []string + OnUpdate func(sessionID string, update json.RawMessage) + OnPermissionRequest func(req PermissionRequest) (bool, error) +} + +// ClientOption is a functional option for configuring an ACP client. +type ClientOption func(*ClientOptions) + +// WithTimeout sets the initialize handshake timeout. +func WithTimeout(d time.Duration) ClientOption { + return func(o *ClientOptions) { + o.Timeout = d + } +} + +// WithEnv sets additional environment variables for the child process. +func WithEnv(env []string) ClientOption { + return func(o *ClientOptions) { + o.Env = env + } +} + +// WithOnUpdate sets a callback for server-streamed session/update notifications. +func WithOnUpdate(fn func(sessionID string, update json.RawMessage)) ClientOption { + return func(o *ClientOptions) { + o.OnUpdate = fn + } +} + +// WithOnPermissionRequest sets a callback for handling server-initiated tool permissions. +func WithOnPermissionRequest(fn func(req PermissionRequest) (bool, error)) ClientOption { + return func(o *ClientOptions) { + o.OnPermissionRequest = fn + } +} + +// Client communicates with an Agent Client Protocol (ACP) peer. +type Client struct { + cmd *exec.Cmd + stdin io.WriteCloser + w io.Writer + writeMu sync.Mutex + pendMu sync.Mutex + pending map[string]chan rpcMessage + nextReqID int + + opts ClientOptions + closed bool + closeMu sync.Mutex + handlers sync.WaitGroup +} + +// Start launches an external ACP process and performs the initialize handshake. +func Start(ctx context.Context, command string, args []string, opts ...ClientOption) (*Client, error) { + var opt ClientOptions + opt.Timeout = 10 * time.Second + for _, o := range opts { + o(&opt) + } + + cmd := exec.Command(command, args...) + setCmdProcessGroup(cmd) + if len(opt.Env) > 0 { + cmd.Env = append(cmd.Environ(), opt.Env...) + } + + stdin, err := cmd.StdinPipe() + if err != nil { + return nil, fmt.Errorf("acp: stdin pipe: %w", err) + } + + stdout, err := cmd.StdoutPipe() + if err != nil { + _ = stdin.Close() + return nil, fmt.Errorf("acp: stdout pipe: %w", err) + } + + if err := cmd.Start(); err != nil { + _ = stdin.Close() + _ = stdout.Close() + return nil, fmt.Errorf("acp: start command: %w", err) + } + + c := &Client{ + cmd: cmd, + stdin: stdin, + w: stdin, + pending: make(map[string]chan rpcMessage), + opts: opt, + } + + // Start read loop + c.handlers.Add(1) + go func() { + defer c.handlers.Done() + c.readLoop(stdout) + }() + + // Initialize handshake with timeout + initCtx, cancel := context.WithTimeout(ctx, opt.Timeout) + defer cancel() + + if err := c.initialize(initCtx); err != nil { + _ = c.Close() + return nil, fmt.Errorf("acp: handshake failed: %w", err) + } + + return c, nil +} + +// Connect wraps an existing io.Reader/io.Writer pair as an ACP client without spawning a process. +func Connect(ctx context.Context, r io.Reader, w io.Writer, opts ...ClientOption) (*Client, error) { + var opt ClientOptions + opt.Timeout = 10 * time.Second + for _, o := range opts { + o(&opt) + } + + c := &Client{ + w: w, + pending: make(map[string]chan rpcMessage), + opts: opt, + } + + c.handlers.Add(1) + go func() { + defer c.handlers.Done() + c.readLoop(r) + }() + + initCtx, cancel := context.WithTimeout(ctx, opt.Timeout) + defer cancel() + + if err := c.initialize(initCtx); err != nil { + _ = c.Close() + return nil, fmt.Errorf("acp: handshake failed: %w", err) + } + + return c, nil +} + +func (c *Client) initialize(ctx context.Context) error { + params := map[string]any{ + "protocolVersion": ProtocolVersion, + "clientCapabilities": map[string]any{ + "tools": map[string]any{ + "approval": true, + }, + }, + } + res, err := c.call(ctx, "initialize", params) + if err != nil { + return err + } + if res.Error != nil { + return fmt.Errorf("rpc error (%d): %s", res.Error.Code, res.Error.Message) + } + return nil +} + +// NewSession creates a new remote ACP session. +func (c *Client) NewSession(ctx context.Context, cwd string) (string, error) { + params := map[string]any{ + "cwd": cwd, + } + res, err := c.call(ctx, "session/new", params) + if err != nil { + return "", err + } + if res.Error != nil { + return "", fmt.Errorf("rpc error (%d): %s", res.Error.Code, res.Error.Message) + } + + var out struct { + SessionID string `json:"sessionId"` + } + if err := json.Unmarshal(res.Result, &out); err != nil { + return "", fmt.Errorf("unmarshal session/new result: %w", err) + } + if out.SessionID == "" { + return "", errors.New("empty sessionId in session/new response") + } + return out.SessionID, nil +} + +// Prompt submits a prompt to an active ACP session and awaits the response. +func (c *Client) Prompt(ctx context.Context, sessionID, prompt string) (*PromptResult, error) { + params := map[string]any{ + "sessionId": sessionID, + "prompt": prompt, + } + res, err := c.call(ctx, "session/prompt", params) + if err != nil { + return nil, err + } + if res.Error != nil { + return nil, fmt.Errorf("rpc error (%d): %s", res.Error.Code, res.Error.Message) + } + + var out PromptResult + if len(res.Result) > 0 { + _ = json.Unmarshal(res.Result, &out) + } + return &out, nil +} + +// Cancel requests cancellation of an in-flight prompt via session/cancel notification. +func (c *Client) Cancel(_ context.Context, sessionID string) error { + params := map[string]any{ + "sessionId": sessionID, + } + return c.notify("session/cancel", params) +} + +func (c *Client) notify(method string, params any) error { + c.closeMu.Lock() + if c.closed { + c.closeMu.Unlock() + return ErrClientClosed + } + c.closeMu.Unlock() + + paramsBytes, err := json.Marshal(params) + if err != nil { + return fmt.Errorf("marshal params: %w", err) + } + + msg := rpcMessage{ + JSONRPC: "2.0", + Method: method, + Params: paramsBytes, + } + return c.send(msg) +} + +func (c *Client) call(ctx context.Context, method string, params any) (rpcMessage, error) { + c.closeMu.Lock() + if c.closed { + c.closeMu.Unlock() + return rpcMessage{}, ErrClientClosed + } + c.closeMu.Unlock() + + c.pendMu.Lock() + c.nextReqID++ + reqID := strconv.Itoa(c.nextReqID) + respCh := make(chan rpcMessage, 1) + c.pending[reqID] = respCh + c.pendMu.Unlock() + + defer func() { + c.pendMu.Lock() + delete(c.pending, reqID) + c.pendMu.Unlock() + }() + + paramsBytes, err := json.Marshal(params) + if err != nil { + return rpcMessage{}, fmt.Errorf("marshal params: %w", err) + } + + rawID := json.RawMessage(strconv.Quote(reqID)) + msg := rpcMessage{ + JSONRPC: "2.0", + ID: rawID, + Method: method, + Params: paramsBytes, + } + + if err := c.send(msg); err != nil { + return rpcMessage{}, err + } + + select { + case <-ctx.Done(): + return rpcMessage{}, ctx.Err() + case resp, ok := <-respCh: + if !ok { + return rpcMessage{}, ErrClientClosed + } + return resp, nil + } +} + +func (c *Client) send(msg rpcMessage) error { + data, err := json.Marshal(msg) + if err != nil { + return fmt.Errorf("marshal rpcMessage: %w", err) + } + data = append(data, '\n') + + c.writeMu.Lock() + defer c.writeMu.Unlock() + + if c.w == nil { + return ErrClientClosed + } + _, err = c.w.Write(data) + return err +} + +func (c *Client) readLoop(r io.Reader) { + scanner := bufio.NewScanner(r) + scanner.Buffer(make([]byte, 0, 1024*1024), 8*1024*1024) + + for scanner.Scan() { + line := scanner.Bytes() + if len(line) == 0 { + continue + } + + var msg rpcMessage + if err := json.Unmarshal(line, &msg); err != nil { + continue + } + + // 1. Response to a client-initiated request + if len(msg.ID) > 0 && msg.Method == "" { + var idStr string + if err := json.Unmarshal(msg.ID, &idStr); err != nil { + // Might be numeric + var idInt int + if errInt := json.Unmarshal(msg.ID, &idInt); errInt == nil { + idStr = strconv.Itoa(idInt) + } + } + + c.pendMu.Lock() + if ch, ok := c.pending[idStr]; ok { + select { + case ch <- msg: + default: + } + } + c.pendMu.Unlock() + continue + } + + // 2. Server-initiated notification or request + switch msg.Method { + case "session/update": + if c.opts.OnUpdate != nil { + var p struct { + SessionID string `json:"sessionId"` + Update json.RawMessage `json:"update"` + } + if err := json.Unmarshal(msg.Params, &p); err == nil { + c.opts.OnUpdate(p.SessionID, p.Update) + } + } + + case "session/request_permission": + c.handlers.Add(1) + go func(req rpcMessage) { + defer c.handlers.Done() + c.handlePermissionRequest(req) + }(msg) + } + } + + // EOF / error: close pending requests + c.closeMu.Lock() + c.closed = true + c.closeMu.Unlock() + + c.pendMu.Lock() + for _, ch := range c.pending { + close(ch) + } + c.pending = make(map[string]chan rpcMessage) + c.pendMu.Unlock() +} + +func (c *Client) handlePermissionRequest(msg rpcMessage) { + var p PermissionRequest + _ = json.Unmarshal(msg.Params, &p) + + allowed := true + if c.opts.OnPermissionRequest != nil { + var err error + allowed, err = c.opts.OnPermissionRequest(p) + if err != nil { + allowed = false + } + } + + optionID := "deny" + if allowed { + optionID = "allow" + } + resPayload, _ := json.Marshal(map[string]any{ + "outcome": map[string]string{ + "outcome": "selected", + "optionId": optionID, + }, + }) + + _ = c.send(rpcMessage{ + JSONRPC: "2.0", + ID: msg.ID, + Result: resPayload, + }) +} + +// Close disposes the client, closing streams and terminating the child process tree if started. +func (c *Client) Close() error { + c.closeMu.Lock() + if c.closed { + c.closeMu.Unlock() + return nil + } + c.closed = true + c.closeMu.Unlock() + + if c.stdin != nil { + _ = c.stdin.Close() + } + + if c.cmd != nil && c.cmd.Process != nil { + _ = killProcessGroup(c.cmd.Process) + _ = c.cmd.Wait() + } + + c.pendMu.Lock() + for _, ch := range c.pending { + close(ch) + } + c.pending = make(map[string]chan rpcMessage) + c.pendMu.Unlock() + + return nil +} diff --git a/internal/acp/client_test.go b/internal/acp/client_test.go new file mode 100644 index 00000000..a128b6f2 --- /dev/null +++ b/internal/acp/client_test.go @@ -0,0 +1,164 @@ +package acp + +import ( + "context" + "encoding/json" + "io" + "strings" + "testing" + "time" + + "github.com/GrayCodeAI/hawk/internal/engine" +) + +// mockSession creates a minimal engine session for testing. +func mockSession() (*engine.Session, error) { + sess := engine.NewSession("mock", "mock-model", "system prompt", nil) + return sess, nil +} + +func setupClientServer(t *testing.T, opts ...ClientOption) (*Client, *Server, func()) { + t.Helper() + serverR, clientW := io.Pipe() + clientR, serverW := io.Pipe() + + server := NewServer(mockSession) + ctx, cancel := context.WithCancel(context.Background()) + + serverDone := make(chan struct{}) + go func() { + defer close(serverDone) + _ = server.Serve(ctx, serverR, serverW) + }() + + client, err := Connect(ctx, clientR, clientW, opts...) + if err != nil { + cancel() + t.Fatalf("failed to connect ACP client: %v", err) + } + + cleanup := func() { + _ = client.Close() + cancel() + _ = clientW.Close() + _ = serverW.Close() + _ = clientR.Close() + _ = serverR.Close() + <-serverDone + } + + return client, server, cleanup +} + +func TestClientServer_InitializeAndNewSession(t *testing.T) { + client, _, cleanup := setupClientServer(t) + defer cleanup() + + ctx := context.Background() + + // Create session + sessID, err := client.NewSession(ctx, "/test/workspace") + if err != nil { + t.Fatalf("NewSession failed: %v", err) + } + if sessID == "" { + t.Fatal("expected non-empty session ID") + } +} + +func TestClientServer_PromptAndCancel(t *testing.T) { + client, _, cleanup := setupClientServer(t) + defer cleanup() + + ctx := context.Background() + + sessID, err := client.NewSession(ctx, "/test/workspace") + if err != nil { + t.Fatalf("NewSession failed: %v", err) + } + + // Cancel session + if err := client.Cancel(ctx, sessID); err != nil { + t.Fatalf("Cancel failed: %v", err) + } +} + +func TestClientServer_PermissionRequest(t *testing.T) { + permApproved := false + var receivedReq PermissionRequest + + client, server, cleanup := setupClientServer(t, WithOnPermissionRequest(func(req PermissionRequest) (bool, error) { + permApproved = true + receivedReq = req + return true, nil + })) + defer cleanup() + + ctx := context.Background() + sessID, err := client.NewSession(ctx, "/test/workspace") + if err != nil { + t.Fatalf("NewSession failed: %v", err) + } + + // Server-initiated request_permission to client + as := server.lookupSession(sessID) + if as == nil { + t.Fatalf("session %s not found on server", sessID) + } + + // Trigger permission request through server + reqPermPayload := map[string]any{ + "sessionId": sessID, + "toolName": "FileWrite", + "summary": "Write main.go", + } + resp, ok := server.call("session/request_permission", reqPermPayload, 5*time.Second) + if !ok { + t.Fatal("server.call timed out or failed") + } + + var permResult struct { + Outcome struct { + Outcome string `json:"outcome"` + OptionID string `json:"optionId"` + } `json:"outcome"` + } + _ = json.Unmarshal(resp.Result, &permResult) + + if !permApproved || permResult.Outcome.OptionID != "allow" { + t.Fatalf("expected permission to be approved by client callback, got outcome=%#v", permResult.Outcome) + } + if receivedReq.ToolName != "FileWrite" { + t.Errorf("expected toolName FileWrite, got %s", receivedReq.ToolName) + } +} + +func TestClient_CloseIdempotence(t *testing.T) { + client, _, cleanup := setupClientServer(t) + defer cleanup() + + if err := client.Close(); err != nil { + t.Fatalf("first Close failed: %v", err) + } + if err := client.Close(); err != nil { + t.Fatalf("second Close failed: %v", err) + } + + // Operations after Close return ErrClientClosed + _, err := client.NewSession(context.Background(), "/test") + if err == nil || !strings.Contains(err.Error(), "closed") { + t.Fatalf("expected client closed error, got: %v", err) + } +} + +func TestClient_HandshakeFailureTimeout(t *testing.T) { + // A broken pipe that provides no response + r, w := io.Pipe() + defer func() { _ = r.Close(); _ = w.Close() }() + + ctx := context.Background() + _, err := Connect(ctx, r, w, WithTimeout(50*time.Millisecond)) + if err == nil { + t.Fatal("expected handshake timeout error, got nil") + } +} diff --git a/internal/acp/proc_unix.go b/internal/acp/proc_unix.go new file mode 100644 index 00000000..a52e467a --- /dev/null +++ b/internal/acp/proc_unix.go @@ -0,0 +1,20 @@ +//go:build !windows + +package acp + +import ( + "os" + "os/exec" + "syscall" +) + +func setCmdProcessGroup(cmd *exec.Cmd) { + cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} +} + +func killProcessGroup(proc *os.Process) error { + if proc == nil { + return nil + } + return syscall.Kill(-proc.Pid, syscall.SIGKILL) +} diff --git a/internal/acp/proc_windows.go b/internal/acp/proc_windows.go new file mode 100644 index 00000000..564882e7 --- /dev/null +++ b/internal/acp/proc_windows.go @@ -0,0 +1,19 @@ +//go:build windows + +package acp + +import ( + "os" + "os/exec" +) + +func setCmdProcessGroup(cmd *exec.Cmd) { + // Process groups on windows handled by default or JobObjects +} + +func killProcessGroup(proc *os.Process) error { + if proc == nil { + return nil + } + return proc.Kill() +} diff --git a/internal/multiagent/provider.go b/internal/multiagent/provider.go new file mode 100644 index 00000000..2009f40d --- /dev/null +++ b/internal/multiagent/provider.go @@ -0,0 +1,284 @@ +package mission + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "os/exec" + "sort" + "sync" + "time" + + "github.com/GrayCodeAI/hawk/internal/acp" +) + +var ( + // ErrProviderNotFound is returned when no subagent provider is registered under the given name. + ErrProviderNotFound = errors.New("subagent: provider not found") + // ErrUnsupportedCapability is returned when a subagent request requires a capability the provider lacks. + ErrUnsupportedCapability = errors.New("subagent: unsupported capability") + // ErrDepthExceeded is returned when the subagent delegation depth limit is exceeded. + ErrDepthExceeded = errors.New("subagent: delegation depth limit exceeded") + // ErrUnsupportedPersona is returned when the requested persona is not supported by the provider. + ErrUnsupportedPersona = errors.New("subagent: unsupported persona") +) + +// SubagentCapabilities declares the features supported by a subagent provider. +type SubagentCapabilities struct { + SupportsStreaming bool `json:"supports_streaming"` + SupportsSchema bool `json:"supports_schema"` + MaxDepth int `json:"max_depth,omitempty"` + SupportedPersonas []string `json:"supported_personas,omitempty"` + AllowedTools []string `json:"allowed_tools,omitempty"` +} + +// SubagentRequest defines the payload for delegating a task to an external subagent. +type SubagentRequest struct { + Name string `json:"name"` + Task string `json:"task"` + CWD string `json:"cwd,omitempty"` + Persona string `json:"persona,omitempty"` + OutputSchema map[string]interface{} `json:"output_schema,omitempty"` + Depth int `json:"depth,omitempty"` + ApprovalGate *MissionApprovalGate `json:"-"` + ParentSession any `json:"-"` +} + +// SubagentResult represents the output of a completed subagent execution. +type SubagentResult struct { + Status string `json:"status"` // "success", "failure", "cancelled" + Output string `json:"output"` + Data map[string]interface{} `json:"data,omitempty"` + Error string `json:"error,omitempty"` + Duration time.Duration `json:"duration"` +} + +// SubagentProvider represents a backend that can execute delegated subagent tasks. +type SubagentProvider interface { + Name() string + Capabilities() SubagentCapabilities + Run(ctx context.Context, req SubagentRequest) (*SubagentResult, error) +} + +// ProviderRegistry manages registered subagent providers. +type ProviderRegistry struct { + mu sync.RWMutex + providers map[string]SubagentProvider +} + +// NewProviderRegistry creates an empty provider registry. +func NewProviderRegistry() *ProviderRegistry { + return &ProviderRegistry{ + providers: make(map[string]SubagentProvider), + } +} + +// Register registers a SubagentProvider and returns a disposer to unregister it. +func (r *ProviderRegistry) Register(p SubagentProvider) func() { + r.mu.Lock() + defer r.mu.Unlock() + r.providers[p.Name()] = p + + return func() { + r.mu.Lock() + defer r.mu.Unlock() + delete(r.providers, p.Name()) + } +} + +// Get retrieves a provider by name. +func (r *ProviderRegistry) Get(name string) (SubagentProvider, bool) { + r.mu.RLock() + defer r.mu.RUnlock() + p, ok := r.providers[name] + return p, ok +} + +// List returns all registered providers sorted by name. +func (r *ProviderRegistry) List() []SubagentProvider { + r.mu.RLock() + defer r.mu.RUnlock() + list := make([]SubagentProvider, 0, len(r.providers)) + for _, p := range r.providers { + list = append(list, p) + } + sort.Slice(list, func(i, j int) bool { + return list[i].Name() < list[j].Name() + }) + return list +} + +// Run validates capabilities and dispatches a task to the target provider. +func (r *ProviderRegistry) Run(ctx context.Context, req SubagentRequest) (*SubagentResult, error) { + p, ok := r.Get(req.Name) + if !ok { + return nil, fmt.Errorf("%w: %s", ErrProviderNotFound, req.Name) + } + + caps := p.Capabilities() + + // 1. OutputSchema capability check + if len(req.OutputSchema) > 0 && !caps.SupportsSchema { + return nil, fmt.Errorf("%w: output schema requested but provider %s does not support it", + ErrUnsupportedCapability, req.Name) + } + + // 2. Depth limit check + if caps.MaxDepth > 0 && req.Depth > caps.MaxDepth { + return nil, fmt.Errorf("%w: requested depth %d exceeds provider %s max depth %d", + ErrDepthExceeded, req.Depth, req.Name, caps.MaxDepth) + } + + // 3. Persona check + if req.Persona != "" && len(caps.SupportedPersonas) > 0 { + matched := false + for _, persona := range caps.SupportedPersonas { + if persona == req.Persona || persona == "*" { + matched = true + break + } + } + if !matched { + return nil, fmt.Errorf("%w: persona %q not supported by provider %s", + ErrUnsupportedPersona, req.Persona, req.Name) + } + } + + start := time.Now() + res, err := p.Run(ctx, req) + if err != nil { + return nil, err + } + if res.Duration == 0 { + res.Duration = time.Since(start) + } + + // Parse structured data if schema was requested and valid JSON output received + if len(req.OutputSchema) > 0 && res.Output != "" && res.Data == nil { + var parsed map[string]interface{} + if jsonErr := json.Unmarshal([]byte(res.Output), &parsed); jsonErr == nil { + res.Data = parsed + } + } + + return res, nil +} + +// --- Built-in ACP Providers --- + +// GenericACPProvider executes tasks via an external ACP-compatible CLI. +type GenericACPProvider struct { + name string + command string + args []string + caps SubagentCapabilities +} + +// NewGenericACPProvider creates a new GenericACPProvider. +func NewGenericACPProvider(name, command string, args []string, caps SubagentCapabilities) *GenericACPProvider { + return &GenericACPProvider{ + name: name, + command: command, + args: args, + caps: caps, + } +} + +func (p *GenericACPProvider) Name() string { return p.name } +func (p *GenericACPProvider) Capabilities() SubagentCapabilities { return p.caps } + +func (p *GenericACPProvider) Run(ctx context.Context, req SubagentRequest) (*SubagentResult, error) { + start := time.Now() + + var opts []acp.ClientOption + if req.ApprovalGate != nil { + opts = append(opts, acp.WithOnPermissionRequest(func(permReq acp.PermissionRequest) (bool, error) { + if err := req.ApprovalGate.Check(ctx, permReq.ToolName, permReq.Summary); err != nil { + return false, nil + } + return true, nil + })) + } + + client, err := acp.Start(ctx, p.command, p.args, opts...) + if err != nil { + return nil, fmt.Errorf("failed to start ACP client %s: %w", p.name, err) + } + defer func() { _ = client.Close() }() + + sessionID, err := client.NewSession(ctx, req.CWD) + if err != nil { + return nil, fmt.Errorf("failed to create ACP session on %s: %w", p.name, err) + } + + prompt := req.Task + if req.Persona != "" { + prompt = fmt.Sprintf("[Persona: %s]\n%s", req.Persona, prompt) + } + + promptRes, err := client.Prompt(ctx, sessionID, prompt) + if err != nil { + if errors.Is(ctx.Err(), context.Canceled) { + _ = client.Cancel(context.Background(), sessionID) + return &SubagentResult{ + Status: "cancelled", + Error: "context canceled", + Duration: time.Since(start), + }, nil + } + return nil, fmt.Errorf("ACP prompt execution failed: %w", err) + } + + status := "success" + if promptRes.Status == "error" { + status = "failure" + } + + return &SubagentResult{ + Status: status, + Output: promptRes.Output, + Duration: time.Since(start), + }, nil +} + +// DefaultRegistry is the singleton default provider registry. +var ( + defaultRegistry *ProviderRegistry + defaultRegistryOnce sync.Once +) + +// DefaultProviders returns the global SubagentProvider registry. +func DefaultProviders() *ProviderRegistry { + defaultRegistryOnce.Do(func() { + defaultRegistry = NewProviderRegistry() + + // 1. Register generic acp provider if 'acp' binary is in PATH + if _, err := exec.LookPath("acp"); err == nil { + defaultRegistry.Register(NewGenericACPProvider("acp", "acp", nil, SubagentCapabilities{ + SupportsStreaming: true, + SupportsSchema: true, + MaxDepth: 3, + })) + } + + // 2. Register claude-code provider if 'claude' is in PATH + if _, err := exec.LookPath("claude"); err == nil { + defaultRegistry.Register(NewGenericACPProvider("claude-code", "claude", []string{"--acp"}, SubagentCapabilities{ + SupportsStreaming: true, + SupportsSchema: true, + MaxDepth: 2, + })) + } + + // 3. Register codex provider if 'codex' is in PATH + if _, err := exec.LookPath("codex"); err == nil { + defaultRegistry.Register(NewGenericACPProvider("codex", "codex", []string{"--acp"}, SubagentCapabilities{ + SupportsStreaming: true, + SupportsSchema: true, + MaxDepth: 2, + })) + } + }) + return defaultRegistry +} diff --git a/internal/multiagent/provider_test.go b/internal/multiagent/provider_test.go new file mode 100644 index 00000000..e7180707 --- /dev/null +++ b/internal/multiagent/provider_test.go @@ -0,0 +1,152 @@ +package mission + +import ( + "context" + "errors" + "testing" + "time" +) + +type mockProvider struct { + name string + caps SubagentCapabilities + run func(ctx context.Context, req SubagentRequest) (*SubagentResult, error) +} + +func (m *mockProvider) Name() string { return m.name } +func (m *mockProvider) Capabilities() SubagentCapabilities { return m.caps } +func (m *mockProvider) Run(ctx context.Context, req SubagentRequest) (*SubagentResult, error) { + if m.run != nil { + return m.run(ctx, req) + } + return &SubagentResult{Status: "success", Output: "done"}, nil +} + +func TestProviderRegistry_RegisterAndList(t *testing.T) { + reg := NewProviderRegistry() + + disposer1 := reg.Register(&mockProvider{name: "beta"}) + disposer2 := reg.Register(&mockProvider{name: "alpha"}) + + list := reg.List() + if len(list) != 2 { + t.Fatalf("expected 2 providers, got %d", len(list)) + } + if list[0].Name() != "alpha" || list[1].Name() != "beta" { + t.Errorf("expected alphabetical sort [alpha, beta], got [%s, %s]", list[0].Name(), list[1].Name()) + } + + // Disposing alpha removes it + disposer2() + if _, ok := reg.Get("alpha"); ok { + t.Error("expected alpha to be removed after disposer call") + } + + disposer1() + if len(reg.List()) != 0 { + t.Errorf("expected 0 providers after all disposed, got %d", len(reg.List())) + } +} + +func TestProviderRegistry_CapabilityChecks(t *testing.T) { + reg := NewProviderRegistry() + reg.Register(&mockProvider{ + name: "limited", + caps: SubagentCapabilities{ + SupportsSchema: false, + MaxDepth: 2, + SupportedPersonas: []string{"reviewer", "tester"}, + }, + }) + + ctx := context.Background() + + // 1. OutputSchema rejection + _, err := reg.Run(ctx, SubagentRequest{ + Name: "limited", + Task: "Audit code", + OutputSchema: map[string]interface{}{"type": "object"}, + }) + if !errors.Is(err, ErrUnsupportedCapability) { + t.Fatalf("expected ErrUnsupportedCapability, got %v", err) + } + + // 2. Depth limit rejection + _, err = reg.Run(ctx, SubagentRequest{ + Name: "limited", + Task: "Audit code", + Depth: 3, // exceeds max depth 2 + }) + if !errors.Is(err, ErrDepthExceeded) { + t.Fatalf("expected ErrDepthExceeded, got %v", err) + } + + // 3. Persona rejection + _, err = reg.Run(ctx, SubagentRequest{ + Name: "limited", + Task: "Audit code", + Persona: "unsupported-persona", + }) + if !errors.Is(err, ErrUnsupportedPersona) { + t.Fatalf("expected ErrUnsupportedPersona, got %v", err) + } + + // 4. Valid run + res, err := reg.Run(ctx, SubagentRequest{ + Name: "limited", + Task: "Audit code", + Persona: "reviewer", + Depth: 1, + }) + if err != nil { + t.Fatalf("expected valid run to succeed, got error: %v", err) + } + if res.Status != "success" { + t.Errorf("expected status success, got %s", res.Status) + } +} + +func TestProviderRegistry_StructuredOutput(t *testing.T) { + reg := NewProviderRegistry() + reg.Register(&mockProvider{ + name: "schema-agent", + caps: SubagentCapabilities{ + SupportsSchema: true, + }, + run: func(ctx context.Context, req SubagentRequest) (*SubagentResult, error) { + return &SubagentResult{ + Status: "success", + Output: `{"findingCount": 3, "summary": "Found 3 issues"}`, + Duration: 10 * time.Millisecond, + }, nil + }, + }) + + ctx := context.Background() + res, err := reg.Run(ctx, SubagentRequest{ + Name: "schema-agent", + Task: "Scan security vulnerabilities", + OutputSchema: map[string]interface{}{"type": "object"}, + }) + if err != nil { + t.Fatalf("Run failed: %v", err) + } + + if res.Data == nil { + t.Fatal("expected structured data to be parsed, got nil") + } + if res.Data["findingCount"] != float64(3) { + t.Errorf("expected findingCount 3, got %#v", res.Data["findingCount"]) + } +} + +func TestProviderRegistry_UnknownProvider(t *testing.T) { + reg := NewProviderRegistry() + _, err := reg.Run(context.Background(), SubagentRequest{ + Name: "nonexistent", + Task: "do something", + }) + if !errors.Is(err, ErrProviderNotFound) { + t.Fatalf("expected ErrProviderNotFound, got %v", err) + } +}