Skip to content
Open
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
447 changes: 410 additions & 37 deletions pkg/ai/callertools/runtime.go

Large diffs are not rendered by default.

188 changes: 188 additions & 0 deletions pkg/ai/callertools/runtime_ginkgo_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import (
"context"
"errors"
"net/http"
"net/http/httptest"
"net/url"
"sync/atomic"
"time"

Expand All @@ -15,6 +17,7 @@ import (

. "github.com/onsi/ginkgo/v2"
. "github.com/onsi/gomega"
. "github.com/onsi/gomega/gstruct"
)

var _ = Describe("Authenticated caller-tool runtime", func() {
Expand Down Expand Up @@ -78,6 +81,7 @@ var _ = Describe("Authenticated caller-tool runtime", func() {
Expect(request.Tool).To(Equal("invoice_update"))
Expect(request.SessionID).To(Equal("captain-session-2"))
Expect(request.ToolUseID).To(Equal("approval-call-1"))
Expect(request.Delegated).To(BeFalse())
return api.PermissionDecision{Allow: true, UpdatedInput: map[string]any{"status": "approved"}}, nil
},
SessionID: "captain-session-2",
Expand Down Expand Up @@ -308,6 +312,167 @@ var _ = Describe("Authenticated caller-tool runtime", func() {
}
Expect(values).To(ConsistOf("first", "second"))
})

It("exposes and executes only the tools selected for a remote task", func(ctx SpecContext) {
remote := httptest.NewServer(callertools.RemoteHandler())
DeferCleanup(remote.Close)
var hiddenCalls atomic.Int32
runtime, err := callertools.New(callertools.Options{
Definitions: []api.ToolDefinition{
{
Name: "version", DefaultPermission: api.ToolPolicyAllow,
Handler: func(context.Context, map[string]any) (any, error) {
return map[string]any{"version": "test"}, nil
},
},
{
Name: "contexts", DefaultPermission: api.ToolPolicyAllow,
Handler: func(context.Context, map[string]any) (any, error) { return []string{"default"}, nil },
},
{
Name: "whoami", DefaultPermission: api.ToolPolicyAllow,
Handler: func(context.Context, map[string]any) (any, error) {
hiddenCalls.Add(1)
return "captain", nil
},
},
},
SessionID: "remote-session",
})
Expect(err).NotTo(HaveOccurred())
DeferCleanup(runtime.Close)

delegation, err := runtime.Endpoint().Delegate(ctx, api.CallerToolBinding{
TaskID: "task-1", Agent: "agent-1", ExpiresAt: time.Now().Add(time.Minute),
ToolNames: []string{"version", "contexts"},
})
Expect(err).NotTo(HaveOccurred())
DeferCleanup(delegation.Revoke)
client := authenticatedClient(ctx, servedDelegation(remote.URL, delegation.Endpoint))
DeferCleanup(client.Close)

tools, err := client.ListTools(ctx, mcp.ListToolsRequest{})
Expect(err).NotTo(HaveOccurred())
names := make([]string, 0, len(tools.Tools))
for _, tool := range tools.Tools {
names = append(names, tool.Name)
}
Expect(names).To(ConsistOf("version", "contexts"))

request := mcp.CallToolRequest{}
request.Params.Name = "version"
result, err := client.CallTool(ctx, request)
Expect(err).NotTo(HaveOccurred())
Expect(result.IsError).To(BeFalse())
Expect(result.StructuredContent).To(Equal(map[string]any{"version": "test"}))

request.Params.Name = "whoami"
_, err = client.CallTool(ctx, request)
Expect(err).To(HaveOccurred())
Expect(hiddenCalls.Load()).To(BeZero())
})

It("rejects remote credentials with the wrong bearer or binding and after expiry or revocation", func(ctx SpecContext) {
remote := httptest.NewServer(callertools.RemoteHandler())
DeferCleanup(remote.Close)
runtime := newRuntime("remote-auth-session", "remote")
DeferCleanup(runtime.Close)
issue := func(expiry time.Time) *api.CallerToolDelegation {
delegation, err := runtime.Endpoint().Delegate(ctx, api.CallerToolBinding{
TaskID: "task-auth", Agent: "agent-auth", ExpiresAt: expiry, ToolNames: []string{"identity"},
})
Expect(err).NotTo(HaveOccurred())
delegation.Endpoint = servedDelegation(remote.URL, delegation.Endpoint)
return delegation
}

active := issue(time.Now().Add(time.Minute))
DeferCleanup(active.Revoke)
invalidBearer := cloneEndpoint(active.Endpoint)
invalidBearer.Headers["Authorization"] = "Bearer invalid"
Expect(authenticatedStatus(invalidBearer)).To(Equal(http.StatusUnauthorized))
wrongTask := cloneEndpoint(active.Endpoint)
wrongTask.Headers[callertools.TaskHeader] = "another-task"
Expect(authenticatedStatus(wrongTask)).To(Equal(http.StatusForbidden))
wrongAgent := cloneEndpoint(active.Endpoint)
wrongAgent.Headers[callertools.AgentHeader] = "another-agent"
Expect(authenticatedStatus(wrongAgent)).To(Equal(http.StatusForbidden))

expiring := issue(time.Now().Add(25 * time.Millisecond))
Eventually(func() int { return authenticatedStatus(expiring.Endpoint) }).Should(Equal(http.StatusUnauthorized))
revoked := issue(time.Now().Add(time.Minute))
revoked.Revoke()
Expect(authenticatedStatus(revoked.Endpoint)).To(Equal(http.StatusUnauthorized))
})

It("returns terminal remote tool results for approval denial and broker failure", func(ctx SpecContext) {
remote := httptest.NewServer(callertools.RemoteHandler())
DeferCleanup(remote.Close)
for _, test := range []struct {
name string
decision api.PermissionDecision
err error
message string
reason string
}{
{name: "denied", decision: api.PermissionDecision{Message: "operator denied the call"}, message: "operator denied the call", reason: "approval_denied"},
{name: "failed", err: errors.New("approval service unavailable"), message: "approval service unavailable", reason: "approval_failed"},
} {
var calls atomic.Int32
events := make(chan api.Event, 2)
audits := make(chan callertools.AuditEvent, 4)
runtime, err := callertools.New(callertools.Options{
Definitions: []api.ToolDefinition{{
Name: "version", DefaultPermission: api.ToolPolicyAsk,
Handler: func(context.Context, map[string]any) (any, error) {
calls.Add(1)
return "must not execute", nil
},
}},
SessionID: "remote-approval-" + test.name,
CanUseTool: func(_ context.Context, request api.PermissionRequest) (api.PermissionDecision, error) {
Expect(request.Delegated).To(BeTrue())
Expect(request.ToolUseIDGenerated).To(BeTrue())
return test.decision, test.err
},
ObserveDelegatedTool: func(_ context.Context, event api.Event) error {
events <- event
return nil
},
Audit: func(event callertools.AuditEvent) { audits <- event },
})
Expect(err).NotTo(HaveOccurred())
delegation, err := runtime.Endpoint().Delegate(ctx, api.CallerToolBinding{
TaskID: "task-" + test.name, Agent: "agent-1", ExpiresAt: time.Now().Add(time.Minute),
ToolNames: []string{"version"},
})
Expect(err).NotTo(HaveOccurred())
client := authenticatedClient(ctx, servedDelegation(remote.URL, delegation.Endpoint))

request := mcp.CallToolRequest{}
request.Params.Name = "version"
result, err := client.CallTool(ctx, request)
Expect(err).NotTo(HaveOccurred())
Expect(result.IsError).To(BeTrue())
Expect(toolResultText(result)).To(ContainSubstring(test.message))
Expect(calls.Load()).To(BeZero())
var use, terminal api.Event
Eventually(events).Should(Receive(&use))
Eventually(events).Should(Receive(&terminal))
Expect(use.Kind).To(Equal(api.EventToolUse))
Expect(terminal).To(MatchFields(IgnoreExtras, Fields{
"Kind": Equal(api.EventToolResult), "ToolCallID": Equal(use.ToolCallID),
"Success": BeFalse(), "Text": ContainSubstring(test.message), "Delegated": BeTrue(),
}))
Eventually(audits).Should(Receive(MatchFields(IgnoreExtras, Fields{
"Action": Equal("call"), "Result": Equal("denied"), "Reason": Equal(test.reason),
})))

Expect(client.Close()).To(Succeed())
delegation.Revoke()
Expect(runtime.Close()).To(Succeed())
}
})
})

func authenticatedClient(ctx context.Context, endpoint api.CallerToolEndpoint) *mcpclient.Client {
Expand Down Expand Up @@ -337,6 +502,29 @@ func authenticatedStatus(endpoint api.CallerToolEndpoint) int {
return response.StatusCode
}

func servedDelegation(serverURL string, endpoint api.CallerToolEndpoint) api.CallerToolEndpoint {
parsed, err := url.Parse(endpoint.URL)
Expect(err).NotTo(HaveOccurred())
endpoint.URL = serverURL + parsed.RequestURI()
return endpoint
}

func cloneEndpoint(endpoint api.CallerToolEndpoint) api.CallerToolEndpoint {
cloned := endpoint
cloned.Headers = make(map[string]string, len(endpoint.Headers))
for name, value := range endpoint.Headers {
cloned.Headers[name] = value
}
return cloned
}

func toolResultText(result *mcp.CallToolResult) string {
Expect(result.Content).NotTo(BeEmpty())
text, ok := mcp.AsTextContent(result.Content[0])
Expect(ok).To(BeTrue())
return text.Text
}

func newRuntime(sessionID, marker string) *callertools.Runtime {
runtime, err := callertools.New(callertools.Options{
Definitions: []api.ToolDefinition{{
Expand Down
12 changes: 11 additions & 1 deletion pkg/ai/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,17 @@ func NewProvider(cfg Config) (Provider, error) {
}

func newResolvedProvider(cfg Config) (Provider, error) {
p, err := api.NewProvider(cfg)
var p Provider
var err error
if cfg.SandboxSelection != nil {
if descriptor, ok := api.SandboxFor(cfg.SandboxSelection.Kind); ok && descriptor.Has(api.CapabilityRemoteExec) {
p, err = newRemoteProvider(cfg)
} else {
p, err = api.NewProvider(cfg)
}
} else {
p, err = api.NewProvider(cfg)
}
if err != nil {
return nil, err
}
Expand Down
106 changes: 106 additions & 0 deletions pkg/ai/remote_provider.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
package ai

import (
"context"
"fmt"
"sync"

"github.com/flanksource/captain/pkg/api"
)

// remoteProvider adapts a whole-run remote sandbox to the provider contract.
// Runtime-only caller-tool authority is attached to the sandbox selection and
// never projected into the serializable request.
type remoteProvider struct {
executor api.RemoteExecutor
sandbox api.Sandbox
model string
backend api.Backend
tools bool

prepareOnce sync.Once
prepareErr error
closeOnce sync.Once
closeErr error
}

func newRemoteProvider(cfg Config) (Provider, error) {
selection := *cfg.SandboxSelection
descriptor, ok := api.SandboxFor(selection.Kind)
if !ok {
return nil, fmt.Errorf("unknown sandbox kind %q", selection.Kind)
}
if err := descriptor.ValidateMode(cfg.Model.Backend.Mode()); err != nil {
return nil, err
}
if len(selection.CallerTools) > 0 && cfg.CallerTools == nil {
return nil, fmt.Errorf("delegated caller tools require a supervisor caller-tool endpoint")
}
if len(selection.CallerTools) > 0 && !api.SupportsCallerTools(cfg.Model.Backend) {
return nil, fmt.Errorf("remote backend %q does not support delegated caller tools", cfg.Model.Backend)
}
selection.CallerToolEndpoint = cfg.CallerTools
sandbox, err := api.NewSandbox(selection)
if err != nil {
return nil, err
}
executor, ok := api.SandboxAs[api.RemoteExecutor](sandbox)
if !ok {
_ = sandbox.Close()
return nil, fmt.Errorf("sandbox %q declares remote execution but provides none", selection.Kind)
}
return &remoteProvider{
executor: executor, sandbox: sandbox, model: cfg.Model.Name,
backend: cfg.Model.Backend, tools: cfg.CallerTools != nil,
}, nil
}

func (provider *remoteProvider) Execute(ctx context.Context, request Request) (*Response, error) {
provider.prepareOnce.Do(func() {
_, provider.prepareErr = provider.sandbox.Prepare(ctx, &request)
})
if provider.prepareErr != nil {
return nil, provider.prepareErr
}
return provider.executor.Execute(ctx, request)
}

func (provider *remoteProvider) ExecuteStream(ctx context.Context, request Request) (<-chan Event, error) {
events := make(chan Event, 3)
go func() {
defer close(events)
response, err := provider.Execute(ctx, request)
if err != nil {
emitRemoteEvent(ctx, events, Event{Kind: EventError, Error: err.Error(), Model: provider.model})
emitRemoteEvent(ctx, events, Event{Kind: EventResult, Success: false, Error: err.Error(), Model: provider.model})
return
}
if response.Text != "" {
emitRemoteEvent(ctx, events, Event{Kind: EventText, Text: response.Text, Model: provider.model})
}
emitRemoteEvent(ctx, events, Event{
Kind: EventResult, Success: true, Model: provider.model,
Usage: &response.Usage, CostUSD: response.CostUSD,
})
}()
return events, nil
}

func emitRemoteEvent(ctx context.Context, events chan<- Event, event Event) {
select {
case events <- event:
case <-ctx.Done():
}
}

func (provider *remoteProvider) GetModel() string { return provider.model }
func (provider *remoteProvider) GetBackend() api.Backend { return provider.backend }

func (provider *remoteProvider) SupportsCallerTools() bool { return provider.tools }

func (provider *remoteProvider) Close() error {
provider.closeOnce.Do(func() {
provider.closeErr = provider.sandbox.Close()
})
return provider.closeErr
}
5 changes: 5 additions & 0 deletions pkg/aichat/approval_execution.go
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,11 @@ func (s *Service) resumeToolApproval(ctx context.Context, threadID string, conti
config.SessionID = continuation.Spec.SessionID
config.CaptainSessionID = execution.CaptainSessionID()
config.Tools = definitions
config.CallerTools = execution.CallerTools()
config, err = s.applySandboxSelection(ctx, config, continuation.Spec.Sandbox)
if err != nil {
return false, err
}
config, err = s.prepareProviderConfig(ctx, config)
if err != nil {
return false, err
Expand Down
Loading
Loading