diff --git a/go.mod b/go.mod index 6b797af1c..e98a13026 100644 --- a/go.mod +++ b/go.mod @@ -5,7 +5,7 @@ go 1.26.1 require ( codeberg.org/miekg/dns v0.6.65 github.com/99designs/gqlgen v0.17.87 - github.com/anthropics/anthropic-sdk-go v1.26.0 + github.com/anthropics/anthropic-sdk-go v1.25.1 github.com/aws/aws-sdk-go-v2 v1.41.3 github.com/aws/aws-sdk-go-v2/credentials v1.19.10 github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.18 @@ -87,7 +87,7 @@ require ( github.com/gobwas/ws v1.4.0 // indirect github.com/goccy/go-yaml v1.19.2 // indirect github.com/gogs/chardet v0.0.0-20211120154057-b7413eaefb8f // indirect - github.com/google/jsonschema-go v0.4.2 // indirect + github.com/google/jsonschema-go v0.4.2 github.com/google/s2a-go v0.1.9 // indirect github.com/google/shlex v0.0.0-20191202100458-e7afc7fbc510 // indirect github.com/google/uuid v1.6.0 // indirect diff --git a/go.sum b/go.sum index 9d008c521..df1ecdecd 100644 --- a/go.sum +++ b/go.sum @@ -12,8 +12,8 @@ github.com/agnivade/levenshtein v1.2.1 h1:EHBY3UOn1gwdy/VbFwgo4cxecRznFk7fKWN1KO github.com/agnivade/levenshtein v1.2.1/go.mod h1:QVVI16kDrtSuwcpd0p1+xMC6Z/VfhtCyDIjcwga4/DU= github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883 h1:bvNMNQO63//z+xNgfBlViaCIJKLlCJ6/fmUseuG0wVQ= github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883/go.mod h1:rCTlJbsFo29Kk6CurOXKm700vrz8f0KW0JNfpkRJY/8= -github.com/anthropics/anthropic-sdk-go v1.26.0 h1:oUTzFaUpAevfuELAP1sjL6CQJ9HHAfT7CoSYSac11PY= -github.com/anthropics/anthropic-sdk-go v1.26.0/go.mod h1:qUKmaW+uuPB64iy1l+4kOSvaLqPXnHTTBKH6RVZ7q5Q= +github.com/anthropics/anthropic-sdk-go v1.25.1 h1:pI8ls4kuiDCUq6Hw+fB7eJx/nfgetLqcT43o/FlZGN0= +github.com/anthropics/anthropic-sdk-go v1.25.1/go.mod h1:WTz31rIUHUHqai2UslPpw5CwXrQP3geYBioRV4WOLvE= github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0 h1:jfIu9sQUG6Ig+0+Ap1h4unLjW6YQJpKZVmUzxsD4E/Q= github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0/go.mod h1:t2tdKJDJF9BV14lnkjHmOQgcvEKgtqs5a1N3LNdJhGE= github.com/aws/aws-sdk-go-v2 v1.41.3 h1:4kQ/fa22KjDt13QCy1+bYADvdgcxpfH18f0zP542kZA= @@ -85,8 +85,6 @@ github.com/digitorus/pkcs7 v0.0.0-20230713084857-e76b763bdc49 h1:h+XMRXf+WLY0h/3 github.com/digitorus/pkcs7 v0.0.0-20230713084857-e76b763bdc49/go.mod h1:SKVExuS+vpu2l9IoOc0RwqE7NYnb0JlcFHFnEJkVDzc= github.com/digitorus/timestamp v0.0.0-20250524132541-c45532741eea h1:ALRwvjsSP53QmnN3Bcj0NpR8SsFLnskny/EIMebAk1c= github.com/digitorus/timestamp v0.0.0-20250524132541-c45532741eea/go.mod h1:GvWntX9qiTlOud0WkQ6ewFm0LPy5JUR1Xo0Ngbd1w6Y= -github.com/dnaeon/go-vcr v1.2.0 h1:zHCHvJYTMh1N7xnV7zf1m1GPBF9Ad0Jk/whtQ1663qI= -github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ= github.com/dnephin/pflag v1.0.7 h1:oxONGlWxhmUct0YzKTgrpQv9AUA1wtPBn7zuSjJqptk= github.com/dnephin/pflag v1.0.7/go.mod h1:uxE91IoWURlOiTUIA8Mq5ZZkAv3dPUfZNaT80Zm7OQE= github.com/fatih/color v1.18.0 h1:S8gINlzdQ840/4pfAwic/ZE0djQEH3wM94VfqLTZcOM= diff --git a/pkg/agent/agent.go b/pkg/agent/agent.go new file mode 100644 index 000000000..2fe3c9750 --- /dev/null +++ b/pkg/agent/agent.go @@ -0,0 +1,368 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "fmt" + "io" + + "go.gearno.de/kit/log" + "go.probo.inc/probo/pkg/llm" +) + +const DefaultMaxTurns = 10 + +type ( + Option func(*Agent) + + Agent struct { + name string + handoffDescription string + instructions string + instructionsFunc func(ctx context.Context, a *Agent) string + model string + modelSettings ModelSettings + tools []Tool + handoffs []*Handoff + mcpServers []*MCPServer + maxTurns int + client *llm.Client + logger *log.Logger + hooks []RunHooks + agentHooks AgentHooks + inputGuardrails []InputGuardrail + outputGuardrails []OutputGuardrail + session Session + sessionID string + outputType *OutputType + toolUseBehavior ToolUseBehavior + resetToolChoice bool + responseFormat *llm.ResponseFormat + approval *ApprovalConfig + } +) + +func New(name string, client *llm.Client, opts ...Option) *Agent { + a := &Agent{ + name: name, + client: client, + maxTurns: DefaultMaxTurns, + toolUseBehavior: RunLLMAgain(), + resetToolChoice: true, + logger: log.NewLogger(log.WithOutput(io.Discard)), + } + + for _, opt := range opts { + opt(a) + } + + a.logger = a.logger.Named("agent").With(log.String("agent", name)) + + return a +} + +func (a *Agent) Name() string { + return a.name +} + +func (a *Agent) AsTool(name, description string) Tool { + return newAgentTool(a, name, description) +} + +func (a *Agent) HandoffDescription() string { + return a.handoffDescription +} + +// Clone creates a shallow copy of the agent with the given options applied. +func (a *Agent) Clone(opts ...Option) *Agent { + cp := *a + + cp.tools = make([]Tool, len(a.tools)) + copy(cp.tools, a.tools) + + cp.handoffs = make([]*Handoff, len(a.handoffs)) + copy(cp.handoffs, a.handoffs) + + cp.mcpServers = make([]*MCPServer, len(a.mcpServers)) + copy(cp.mcpServers, a.mcpServers) + + cp.hooks = make([]RunHooks, len(a.hooks)) + copy(cp.hooks, a.hooks) + + cp.inputGuardrails = make([]InputGuardrail, len(a.inputGuardrails)) + copy(cp.inputGuardrails, a.inputGuardrails) + + cp.outputGuardrails = make([]OutputGuardrail, len(a.outputGuardrails)) + copy(cp.outputGuardrails, a.outputGuardrails) + + if a.approval != nil { + newApproval := ApprovalConfig{ + ShouldApprove: a.approval.ShouldApprove, + } + if len(a.approval.ToolNames) > 0 { + newApproval.ToolNames = make([]string, len(a.approval.ToolNames)) + copy(newApproval.ToolNames, a.approval.ToolNames) + newApproval.toolNameSet = buildToolNameSet(newApproval.ToolNames) + } + cp.approval = &newApproval + } + + for _, opt := range opts { + opt(&cp) + } + + return &cp +} + +// clone creates a shallow copy suitable for overriding pointer fields +// (e.g. responseFormat) without affecting the original. Slice fields remain +// shared; use the exported Clone method when slice mutations are needed. +func (a *Agent) clone() *Agent { + cp := *a + return &cp +} + +func WithInstructions(s string) Option { + return func(a *Agent) { + a.instructions = s + a.instructionsFunc = nil + } +} + +func WithInstructionsFunc(fn func(ctx context.Context, a *Agent) string) Option { + return func(a *Agent) { + a.instructionsFunc = fn + a.instructions = "" + } +} + +func WithHandoffDescription(desc string) Option { + return func(a *Agent) { + a.handoffDescription = desc + } +} + +func WithModel(m string) Option { + return func(a *Agent) { + a.model = m + } +} + +func WithModelSettings(s ModelSettings) Option { + return func(a *Agent) { + a.modelSettings = s + } +} + +func WithTools(tools ...Tool) Option { + return func(a *Agent) { + a.tools = append(a.tools, tools...) + } +} + +func WithHandoffs(agents ...*Agent) Option { + return func(a *Agent) { + for _, ag := range agents { + a.handoffs = append(a.handoffs, &Handoff{Agent: ag}) + } + } +} + +func WithHandoffConfigs(handoffs ...*Handoff) Option { + return func(a *Agent) { + a.handoffs = append(a.handoffs, handoffs...) + } +} + +func WithMaxTurns(n int) Option { + return func(a *Agent) { + if n < 1 { + n = 1 + } + a.maxTurns = n + } +} + +func WithTemperature(t float64) Option { + return func(a *Agent) { + a.modelSettings.Temperature = &t + } +} + +func WithTopP(p float64) Option { + return func(a *Agent) { + a.modelSettings.TopP = &p + } +} + +func WithFrequencyPenalty(p float64) Option { + return func(a *Agent) { + a.modelSettings.FrequencyPenalty = &p + } +} + +func WithPresencePenalty(p float64) Option { + return func(a *Agent) { + a.modelSettings.PresencePenalty = &p + } +} + +func WithMaxTokens(n int) Option { + return func(a *Agent) { + a.modelSettings.MaxTokens = &n + } +} + +func WithToolChoice(tc llm.ToolChoice) Option { + return func(a *Agent) { + a.modelSettings.ToolChoice = &tc + } +} + +func WithParallelToolCalls(enabled bool) Option { + return func(a *Agent) { + a.modelSettings.ParallelToolCalls = &enabled + } +} + +func WithLogger(l *log.Logger) Option { + return func(a *Agent) { + a.logger = l + } +} + +func WithHooks(hooks ...RunHooks) Option { + return func(a *Agent) { + a.hooks = append(a.hooks, hooks...) + } +} + +func WithAgentHooks(hooks AgentHooks) Option { + return func(a *Agent) { + a.agentHooks = hooks + } +} + +func WithInputGuardrails(guards ...InputGuardrail) Option { + return func(a *Agent) { + a.inputGuardrails = append(a.inputGuardrails, guards...) + } +} + +func WithOutputGuardrails(guards ...OutputGuardrail) Option { + return func(a *Agent) { + a.outputGuardrails = append(a.outputGuardrails, guards...) + } +} + +func WithSession(s Session, sessionID string) Option { + return func(a *Agent) { + a.session = s + a.sessionID = sessionID + } +} + +func WithOutputType(t *OutputType) Option { + return func(a *Agent) { + a.outputType = t + } +} + +func WithToolUseBehavior(b ToolUseBehavior) Option { + return func(a *Agent) { + a.toolUseBehavior = b + } +} + +func WithResetToolChoice(reset bool) Option { + return func(a *Agent) { + a.resetToolChoice = reset + } +} + +func WithMCPServers(servers ...*MCPServer) Option { + return func(a *Agent) { + a.mcpServers = append(a.mcpServers, servers...) + } +} + +func WithApproval(config ApprovalConfig) Option { + config.toolNameSet = buildToolNameSet(config.ToolNames) + return func(a *Agent) { + a.approval = &config + } +} + +func (a *Agent) resolveTools(ctx context.Context) ([]ToolDescriptor, map[string]ToolDescriptor, error) { + var all []ToolDescriptor + + for _, t := range a.tools { + all = append(all, t) + } + + for _, h := range a.handoffs { + all = append(all, h.tool()) + } + + for _, s := range a.mcpServers { + mcpTools, err := s.Tools(ctx) + if err != nil { + return nil, nil, fmt.Errorf("cannot resolve MCP tools from %q: %w", s.name, err) + } + for _, t := range mcpTools { + all = append(all, t) + } + } + + toolMap := make(map[string]ToolDescriptor, len(all)) + for _, t := range all { + name := t.Name() + if _, exists := toolMap[name]; exists { + return nil, nil, fmt.Errorf("cannot resolve tools: duplicate tool name %q", name) + } + toolMap[name] = t + } + + return all, toolMap, nil +} + +func (a *Agent) buildSystemPrompt(ctx context.Context) string { + instr := a.instructions + if a.instructionsFunc != nil { + instr = a.instructionsFunc(ctx, a) + } + + data := systemPromptData{ + Instructions: instr, + } + + for _, h := range a.handoffs { + desc := h.toolDescription() + if len([]rune(desc)) > 200 { + desc = string([]rune(desc)[:200]) + "..." + } + + data.Handoffs = append( + data.Handoffs, + systemPromptHandoff{ + Name: h.Agent.name, + Description: desc, + }, + ) + } + + return buildSystemPrompt(data) +} diff --git a/pkg/agent/agent_test.go b/pkg/agent/agent_test.go new file mode 100644 index 000000000..46a89f5f8 --- /dev/null +++ b/pkg/agent/agent_test.go @@ -0,0 +1,2918 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +type mockProvider struct { + responses []*llm.ChatCompletionResponse + calls int +} + +func (m *mockProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + if m.calls >= len(m.responses) { + return nil, errors.New("no more mock responses") + } + resp := m.responses[m.calls] + m.calls++ + return resp, nil +} + +func (m *mockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) { + return nil, errors.New("not implemented") +} + +type mockChatStream struct { + events []llm.ChatCompletionStreamEvent + pos int +} + +func (s *mockChatStream) Next() bool { + return s.pos < len(s.events) +} + +func (s *mockChatStream) Event() llm.ChatCompletionStreamEvent { + ev := s.events[s.pos] + s.pos++ + return ev +} + +func (s *mockChatStream) Err() error { return nil } +func (s *mockChatStream) Close() error { return nil } + +type mockStreamProvider struct { + stream llm.ChatCompletionStream + calls int +} + +func (p *mockStreamProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + return nil, errors.New("not implemented") +} + +func (p *mockStreamProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) { + p.calls++ + return p.stream, nil +} + +type mockMultiStreamProvider struct { + streams []llm.ChatCompletionStream + calls int +} + +func (p *mockMultiStreamProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + return nil, errors.New("not implemented") +} + +func (p *mockMultiStreamProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) { + if p.calls >= len(p.streams) { + return nil, errors.New("no more mock streams") + } + s := p.streams[p.calls] + p.calls++ + return s, nil +} + +type blockingGuardrail struct { + keyword string +} + +func (g *blockingGuardrail) Name() string { return "blocker" } + +func (g *blockingGuardrail) Check(_ context.Context, messages []llm.Message) (*agent.GuardrailResult, error) { + for _, m := range messages { + if m.Role == llm.RoleUser { + for _, p := range m.Parts { + if tp, ok := p.(llm.TextPart); ok { + if tp.Text == g.keyword { + return &agent.GuardrailResult{ + Tripwire: true, + Message: "blocked content detected", + }, nil + } + } + } + } + } + return nil, nil +} + +type outputBlocker struct{} + +func (g *outputBlocker) Name() string { return "output_blocker" } + +func (g *outputBlocker) Check(_ context.Context, message llm.Message) (*agent.GuardrailResult, error) { + if message.Text() == "bad response" { + return &agent.GuardrailResult{ + Tripwire: true, + Message: "output blocked", + }, nil + } + return nil, nil +} + +type recordingHook struct { + agent.NoOpHooks + runStarted bool + runEnded bool + toolNames []string + handoffs []string +} + +func (h *recordingHook) OnRunStart(_ context.Context, _ *agent.Agent, _ []llm.Message) { + h.runStarted = true +} + +func (h *recordingHook) OnRunEnd(_ context.Context, _ *agent.Agent, _ *agent.Result, _ error) { + h.runEnded = true +} + +func (h *recordingHook) OnToolEnd(_ context.Context, _ *agent.Agent, tool agent.Tool, _ agent.ToolResult, _ error) { + h.toolNames = append(h.toolNames, tool.Name()) +} + +func (h *recordingHook) OnHandoff(_ context.Context, from *agent.Agent, to *agent.Agent) { + h.handoffs = append(h.handoffs, from.Name()+"->"+to.Name()) +} + +type recordingAgentHook struct { + agent.NoOpAgentHooks + started bool + ended bool + handoffed bool +} + +func (h *recordingAgentHook) OnStart(_ context.Context, _ *agent.Agent) { + h.started = true +} + +func (h *recordingAgentHook) OnEnd(_ context.Context, _ *agent.Agent, _ string) { + h.ended = true +} + +func (h *recordingAgentHook) OnHandoff(_ context.Context, _ *agent.Agent, _ *agent.Agent) { + h.handoffed = true +} + +type testSession struct { + messages map[string][]llm.Message +} + +func newTestSession() *testSession { + return &testSession{messages: make(map[string][]llm.Message)} +} + +func (s *testSession) Load(_ context.Context, sessionID string) ([]llm.Message, error) { + msgs := s.messages[sessionID] + cp := make([]llm.Message, len(msgs)) + copy(cp, msgs) + return cp, nil +} + +func (s *testSession) Save(_ context.Context, sessionID string, messages []llm.Message) error { + cp := make([]llm.Message, len(messages)) + copy(cp, messages) + s.messages[sessionID] = cp + return nil +} + +type failingSession struct{} + +func (s *failingSession) Load(_ context.Context, _ string) ([]llm.Message, error) { + return nil, nil +} + +func (s *failingSession) Save(_ context.Context, _ string, _ []llm.Message) error { + return errors.New("disk full") +} + +type errorLoadSession struct{} + +func (s *errorLoadSession) Load(_ context.Context, _ string) ([]llm.Message, error) { + return nil, errors.New("storage unavailable") +} + +func (s *errorLoadSession) Save(_ context.Context, _ string, _ []llm.Message) error { + return nil +} + +type errorGuardrail struct { + err error +} + +func (g *errorGuardrail) Name() string { return "error_guardrail" } + +func (g *errorGuardrail) Check(_ context.Context, _ []llm.Message) (*agent.GuardrailResult, error) { + return nil, g.err +} + +type errorOutputGuardrail struct { + err error +} + +func (g *errorOutputGuardrail) Name() string { return "error_output_guardrail" } + +func (g *errorOutputGuardrail) Check(_ context.Context, _ llm.Message) (*agent.GuardrailResult, error) { + return nil, g.err +} + +type errorChatStream struct { + err error +} + +func (s *errorChatStream) Next() bool { return false } +func (s *errorChatStream) Event() llm.ChatCompletionStreamEvent { return llm.ChatCompletionStreamEvent{} } +func (s *errorChatStream) Err() error { return s.err } +func (s *errorChatStream) Close() error { return nil } + +type errorStreamProvider struct { + err error +} + +func (p *errorStreamProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + return nil, errors.New("not implemented") +} + +func (p *errorStreamProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) { + return nil, p.err +} + +func newTestClient(provider llm.Provider) *llm.Client { + return llm.NewClient(provider, "test") +} + +func userMessage(text string) llm.Message { + return llm.Message{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: text}}, + } +} + +func assistantMessage(text string) llm.Message { + return llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: text}}, + } +} + +func stopResponse(text string) *llm.ChatCompletionResponse { + return &llm.ChatCompletionResponse{ + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: text}}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + FinishReason: llm.FinishReasonStop, + } +} + +func toolCallResponse(toolCalls ...llm.ToolCall) *llm.ChatCompletionResponse { + return &llm.ChatCompletionResponse{ + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + ToolCalls: toolCalls, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + FinishReason: llm.FinishReasonToolCalls, + } +} + +func finishReasonPtr(r llm.FinishReason) *llm.FinishReason { + return &r +} + +func TestRun(t *testing.T) { + t.Parallel() + + t.Run( + "simple completion", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Hello!"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithInstructions("You are helpful."), + agent.WithModel("test-model"), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + require.NoError(t, err) + assert.Equal(t, 1, result.Turns) + assert.Equal(t, "Hello!", result.FinalMessage().Text()) + assert.Equal(t, 10, result.Usage.InputTokens) + assert.Equal(t, 5, result.Usage.OutputTokens) + assert.Equal(t, "assistant", result.LastAgent.Name()) + }, + ) + + t.Run( + "tool call", + func(t *testing.T) { + t.Parallel() + + type Params struct { + City string `json:"city"` + } + + weatherTool := agent.FunctionTool[Params]( + "get_weather", + "Get weather for a city", + func(_ context.Context, p Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "Sunny, 22°C in " + p.City}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "get_weather", + Arguments: `{"city":"Paris"}`, + }, + }), + stopResponse("It's sunny and 22°C in Paris!"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(weatherTool), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("What's the weather in Paris?")}, + ) + + require.NoError(t, err) + assert.Equal(t, 2, result.Turns) + assert.Equal(t, "It's sunny and 22°C in Paris!", result.FinalMessage().Text()) + assert.Equal(t, 20, result.Usage.InputTokens) + }, + ) + + t.Run( + "max turns exceeded", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, + }), + toolCallResponse(llm.ToolCall{ + ID: "tc_2", + Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, + }), + }, + } + + type Params struct{} + noopTool := agent.FunctionTool[Params]( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(noopTool), + agent.WithMaxTurns(2), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("loop")}, + ) + + require.Error(t, err) + var maxTurnsErr *agent.MaxTurnsExceededError + require.ErrorAs(t, err, &maxTurnsErr) + assert.Equal(t, 2, maxTurnsErr.MaxTurns) + }, + ) + + t.Run( + "duplicate tool names", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + makeTool := func(name string) agent.Tool { + return agent.FunctionTool[Params]( + name, + "desc", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + } + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(makeTool("search"), makeTool("search")), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "duplicate tool name") + }, + ) + + t.Run( + "context cancellation", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("should not reach"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := ag.Run(ctx, []llm.Message{userMessage("test")}) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot complete") + assert.Equal(t, 0, provider.calls) + }, + ) + + t.Run( + "finish reason length treated as stop", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + { + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: "Truncated response"}}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 50}, + FinishReason: llm.FinishReasonLength, + }, + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Write a long essay")}, + ) + + require.NoError(t, err) + assert.Equal(t, 1, result.Turns) + assert.Equal(t, "Truncated response", result.FinalMessage().Text()) + assert.Equal(t, 50, result.Usage.OutputTokens) + }, + ) + + t.Run( + "parallel tool execution preserves ordering", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool1 := agent.FunctionTool[Params]( + "first", + "First tool", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "result_1"}, nil + }, + ) + tool2 := agent.FunctionTool[Params]( + "second", + "Second tool", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "result_2"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse( + llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "first", Arguments: `{}`}, + }, + llm.ToolCall{ + ID: "tc_2", + Function: llm.FunctionCall{Name: "second", Arguments: `{}`}, + }, + ), + stopResponse("Both done."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool1, tool2), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("do both")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Both done.", result.FinalMessage().Text()) + + var toolMsgs []llm.Message + for _, m := range result.Messages { + if m.Role == llm.RoleTool { + toolMsgs = append(toolMsgs, m) + } + } + require.Len(t, toolMsgs, 2) + assert.Equal(t, "tc_1", toolMsgs[0].ToolCallID) + assert.Equal(t, "result_1", toolMsgs[0].Text()) + assert.Equal(t, "tc_2", toolMsgs[1].ToolCallID) + assert.Equal(t, "result_2", toolMsgs[1].Text()) + }, + ) + + t.Run( + "parallel tool execution with partial failure", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + successTool := agent.FunctionTool[Params]( + "succeed", + "Always succeeds", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "success_result"}, nil + }, + ) + failTool := agent.FunctionTool[Params]( + "fail", + "Always fails", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, errors.New("tool exploded") + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse( + llm.ToolCall{ + ID: "tc_ok", + Function: llm.FunctionCall{Name: "succeed", Arguments: `{}`}, + }, + llm.ToolCall{ + ID: "tc_fail", + Function: llm.FunctionCall{Name: "fail", Arguments: `{}`}, + }, + ), + stopResponse("Handled both."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(successTool, failTool), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("do both")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Handled both.", result.FinalMessage().Text()) + + var toolMsgs []llm.Message + for _, m := range result.Messages { + if m.Role == llm.RoleTool { + toolMsgs = append(toolMsgs, m) + } + } + require.Len(t, toolMsgs, 2) + assert.Equal(t, "tc_ok", toolMsgs[0].ToolCallID) + assert.Equal(t, "success_result", toolMsgs[0].Text()) + assert.Equal(t, "tc_fail", toolMsgs[1].ToolCallID) + assert.Contains(t, toolMsgs[1].Text(), "Error:") + }, + ) + + t.Run( + "tool accesses run context during execution", + func(t *testing.T) { + t.Parallel() + + type RequestContext struct { + TenantID string + } + + var capturedTenantID string + + type Params struct{} + tool := agent.FunctionTool[Params]( + "check_tenant", + "Check current tenant", + func(ctx context.Context, _ Params) (agent.ToolResult, error) { + rc := agent.RunContextFrom[*RequestContext](ctx) + capturedTenantID = rc.TenantID + return agent.ToolResult{Content: "tenant: " + rc.TenantID}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "check_tenant", Arguments: `{}`}, + }), + stopResponse("Done"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool), + ) + + ctx := agent.WithRunContext( + context.Background(), + &RequestContext{TenantID: "t_456"}, + ) + + result, err := ag.Run( + ctx, + []llm.Message{userMessage("Check my tenant")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Done", result.FinalMessage().Text()) + assert.Equal(t, "t_456", capturedTenantID) + }, + ) +} + +func TestRun_Handoff(t *testing.T) { + t.Parallel() + + t.Run( + "basic handoff", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_billing", Arguments: `{}`}, + }), + stopResponse("Your invoice is $42."), + }, + } + + client := newTestClient(provider) + + billing := agent.New( + "billing", + client, + agent.WithInstructions("You handle billing questions."), + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithInstructions("Route users to the right agent."), + agent.WithModel("test-model"), + agent.WithHandoffs(billing), + ) + + result, err := triage.Run( + context.Background(), + []llm.Message{userMessage("How much is my invoice?")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Your invoice is $42.", result.FinalMessage().Text()) + assert.Equal(t, "billing", result.LastAgent.Name()) + }, + ) + + t.Run( + "custom tool name and description", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "ask_billing", Arguments: `{}`}, + }), + stopResponse("Your invoice is $42."), + }, + } + + client := newTestClient(provider) + + billing := agent.New( + "billing", + client, + agent.WithModel("test-model"), + agent.WithHandoffDescription("Handles all billing questions."), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + billing, + agent.WithHandoffToolName("ask_billing"), + agent.WithHandoffToolDescription("Route to billing"), + ), + ), + ) + + result, err := triage.Run( + context.Background(), + []llm.Message{userMessage("How much is my invoice?")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Your invoice is $42.", result.FinalMessage().Text()) + assert.Equal(t, "billing", result.LastAgent.Name()) + }, + ) + + t.Run( + "on_handoff callback fires", + func(t *testing.T) { + t.Parallel() + + var callbackFired bool + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_billing", Arguments: `{}`}, + }), + stopResponse("Done."), + }, + } + + client := newTestClient(provider) + + billing := agent.New( + "billing", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + billing, + agent.WithOnHandoff(func(_ context.Context) error { + callbackFired = true + return nil + }), + ), + ), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.True(t, callbackFired) + }, + ) + + t.Run( + "input filter filters messages", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }), + stopResponse("Filtered."), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + router := agent.New( + "router", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + specialist, + agent.WithHandoffInputFilter(func(data agent.HandoffInputData) []llm.Message { + var filtered []llm.Message + for _, m := range data.NewItems { + if m.Role == llm.RoleUser { + filtered = append(filtered, m) + } + } + return filtered + }), + ), + ), + ) + + result, err := router.Run( + context.Background(), + []llm.Message{userMessage("help me")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Filtered.", result.FinalMessage().Text()) + }, + ) + + t.Run( + "tool name sanitized from agent name", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_billing_support", Arguments: `{}`}, + }), + stopResponse("Done."), + }, + } + + client := newTestClient(provider) + + billing := agent.New( + "billing support", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffs(billing), + ) + + result, err := triage.Run( + context.Background(), + []llm.Message{userMessage("help")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Done.", result.FinalMessage().Text()) + assert.Equal(t, "billing support", result.LastAgent.Name()) + }, + ) + + t.Run( + "handoff description getter", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "billing", + newTestClient(&mockProvider{}), + agent.WithHandoffDescription("Handles billing and invoicing."), + ) + + assert.Equal(t, "Handles billing and invoicing.", ag.HandoffDescription()) + }, + ) +} + +func TestRun_Guardrails(t *testing.T) { + t.Parallel() + + t.Run( + "input guardrail trips", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("should not reach"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInputGuardrails(&blockingGuardrail{keyword: "forbidden"}), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("forbidden")}, + ) + + require.Error(t, err) + var tripErr *agent.InputGuardrailTrippedError + require.ErrorAs(t, err, &tripErr) + assert.Equal(t, "blocker", tripErr.Guardrail) + assert.Equal(t, 0, provider.calls) + }, + ) + + t.Run( + "output guardrail trips", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("bad response"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithOutputGuardrails(&outputBlocker{}), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + var tripErr *agent.OutputGuardrailTrippedError + require.ErrorAs(t, err, &tripErr) + assert.Equal(t, "output_blocker", tripErr.Guardrail) + }, + ) +} + +func TestRun_Hooks(t *testing.T) { + t.Parallel() + + t.Run( + "run hooks fire during tool use", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + noopTool := agent.FunctionTool[Params]( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, + }), + stopResponse("done"), + }, + } + + hook := &recordingHook{} + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(noopTool), + agent.WithHooks(hook), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.Equal(t, "done", result.FinalMessage().Text()) + assert.True(t, hook.runStarted) + assert.True(t, hook.runEnded) + assert.Equal(t, []string{"noop"}, hook.toolNames) + }, + ) + + t.Run( + "agent hooks fire during run", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Hello!"), + }, + } + + hook := &recordingAgentHook{} + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithAgentHooks(hook), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + require.NoError(t, err) + assert.True(t, hook.started) + assert.True(t, hook.ended) + }, + ) + + t.Run( + "target agent OnHandoff fires on handoff", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_billing", Arguments: `{}`}, + }), + stopResponse("Done."), + }, + } + + client := newTestClient(provider) + + billingHook := &recordingAgentHook{} + billing := agent.New( + "billing", + client, + agent.WithModel("test-model"), + agent.WithAgentHooks(billingHook), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffs(billing), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("invoice?")}, + ) + + require.NoError(t, err) + assert.True(t, billingHook.handoffed) + assert.True(t, billingHook.started) + }, + ) + + t.Run( + "OnRunEnd fires even on session save failure", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Hello!"), + }, + } + + hook := &recordingHook{} + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithHooks(hook), + agent.WithSession(&failingSession{}, "sess-1"), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot save session") + assert.True(t, hook.runStarted, "OnRunStart should have fired") + assert.True(t, hook.runEnded, "OnRunEnd should fire even when session save fails") + }, + ) +} + +func TestRun_Session(t *testing.T) { + t.Parallel() + + t.Run( + "round trip across two runs", + func(t *testing.T) { + t.Parallel() + + store := newTestSession() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Hi Alice!"), + stopResponse("Your name is Alice."), + }, + } + + client := newTestClient(provider) + + ag1 := agent.New( + "assistant", + client, + agent.WithModel("test-model"), + agent.WithInstructions("You are a helpful assistant."), + agent.WithSession(store, "session-1"), + ) + + _, err := ag1.Run( + context.Background(), + []llm.Message{userMessage("I'm Alice")}, + ) + require.NoError(t, err) + + ag2 := agent.New( + "assistant", + client, + agent.WithModel("test-model"), + agent.WithInstructions("You are a helpful assistant."), + agent.WithSession(store, "session-1"), + ) + + result, err := ag2.Run( + context.Background(), + []llm.Message{userMessage("What's my name?")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Your name is Alice.", result.FinalMessage().Text()) + assert.Len(t, result.Messages, 4) + }, + ) +} + +func TestRun_DynamicInstructions(t *testing.T) { + t.Parallel() + + t.Run( + "instructions from function with run context", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("Hello, Alice!"), + }, + } + + type userInfo struct { + Name string + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInstructionsFunc(func(ctx context.Context, _ *agent.Agent) string { + info := agent.RunContextFrom[*userInfo](ctx) + return "You are helping " + info.Name + ". Be concise." + }), + ) + + ctx := agent.WithRunContext( + context.Background(), + &userInfo{Name: "Alice"}, + ) + result, err := ag.Run( + ctx, + []llm.Message{userMessage("Hi")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Hello, Alice!", result.FinalMessage().Text()) + }, + ) + + t.Run( + "instructionsFunc overrides static instructions", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("dynamic"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInstructions("static"), + agent.WithInstructionsFunc(func(_ context.Context, _ *agent.Agent) string { + return "dynamic" + }), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.Equal(t, "dynamic", result.FinalMessage().Text()) + }, + ) +} + +func TestRun_ToolUseBehavior(t *testing.T) { + t.Parallel() + + t.Run( + "stop on first tool", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + tool := agent.FunctionTool[Params]( + "compute", + "Compute something", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "computed_result"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "compute", Arguments: `{}`}, + }), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool), + agent.WithToolUseBehavior(agent.StopOnFirstTool()), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("compute")}, + ) + + require.NoError(t, err) + assert.Equal(t, "computed_result", result.FinalMessage().Text()) + assert.Equal(t, 1, result.Turns) + assert.Equal(t, 1, provider.calls) + }, + ) + + t.Run( + "stop at specific tools", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool1 := agent.FunctionTool[Params]( + "search", + "Search", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "search_result"}, nil + }, + ) + tool2 := agent.FunctionTool[Params]( + "submit", + "Submit", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "submitted"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "submit", Arguments: `{}`}, + }), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool1, tool2), + agent.WithToolUseBehavior(agent.StopAtTools("submit")), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("submit it")}, + ) + + require.NoError(t, err) + assert.Equal(t, "submitted", result.FinalMessage().Text()) + }, + ) + + t.Run( + "default run_llm_again continues loop", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + tool := agent.FunctionTool[Params]( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, + }), + stopResponse("Final answer."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Final answer.", result.FinalMessage().Text()) + assert.Equal(t, 2, result.Turns) + }, + ) + + t.Run( + "custom behavior error propagation", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + tool := agent.FunctionTool[Params]( + "compute", + "Compute something", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "result"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "compute", Arguments: `{}`}, + }), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool), + agent.WithToolUseBehavior(agent.ToolUseBehavior(func(_ context.Context, _ []agent.ToolCallResult) (string, bool, error) { + return "", false, errors.New("custom behavior failed") + })), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("compute")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "custom behavior failed") + }, + ) +} + +func TestRun_OutputType(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + Country string `json:"country"` + } + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse(`{"name":"Probo","country":"FR"}`), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithOutputType(agent.NewOutputType[Info]("info")), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Tell me about Probo")}, + ) + + require.NoError(t, err) + assert.Contains(t, result.FinalMessage().Text(), "Probo") +} + +func TestRun_Approval(t *testing.T) { + t.Parallel() + + t.Run( + "tool requiring approval interrupts the run", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, + }), + }, + } + + deleteTool := agent.FunctionTool[struct{}]( + "delete_account", + "Deletes the user account", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + return agent.ToolResult{Content: "deleted"}, nil + }, + ) + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(deleteTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"delete_account"}, + }), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Delete my account")}, + ) + + require.Error(t, err) + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.Len(t, interrupted.ToolCalls, 1) + assert.Equal(t, "delete_account", interrupted.ToolCalls[0].Function.Name) + assert.Len(t, interrupted.PendingApprovals, 1) + assert.Equal(t, "delete_account", interrupted.PendingApprovals[0].Function.Name) + assert.Equal(t, 1, interrupted.Turns) + }, + ) + + t.Run( + "resume with approval executes the tool", + func(t *testing.T) { + t.Parallel() + + var toolExecuted bool + + deleteTool := agent.FunctionTool[struct{}]( + "delete_account", + "Deletes the user account", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + toolExecuted = true + return agent.ToolResult{Content: "account deleted"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, + }), + stopResponse("Your account has been deleted."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(deleteTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"delete_account"}, + }), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Delete my account")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.False(t, toolExecuted) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: true}, + }, + }, + ) + + require.NoError(t, err) + assert.True(t, toolExecuted) + assert.Equal(t, "Your account has been deleted.", result.FinalMessage().Text()) + }, + ) + + t.Run( + "resume with rejection denies the tool", + func(t *testing.T) { + t.Parallel() + + deleteTool := agent.FunctionTool[struct{}]( + "delete_account", + "Deletes the user account", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + t.Fatal("tool should not be executed") + return agent.ToolResult{}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, + }), + stopResponse("OK, I won't delete your account."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(deleteTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"delete_account"}, + }), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Delete my account")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: false, Message: "User cancelled the operation."}, + }, + }, + ) + + require.NoError(t, err) + assert.Equal(t, "OK, I won't delete your account.", result.FinalMessage().Text()) + }, + ) + + t.Run( + "ShouldApprove function takes priority", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "safe_tool", Arguments: `{}`}, + }), + stopResponse("Done"), + }, + } + + safeTool := agent.FunctionTool[struct{}]( + "safe_tool", + "A safe tool", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + return agent.ToolResult{Content: "safe result"}, nil + }, + ) + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(safeTool), + agent.WithApproval(agent.ApprovalConfig{ + ShouldApprove: func(_ context.Context, tc llm.ToolCall) bool { + return tc.Function.Name == "dangerous_tool" + }, + }), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Do safe thing")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Done", result.FinalMessage().Text()) + }, + ) + + t.Run( + "multi-tool batch with partial approval", + func(t *testing.T) { + t.Parallel() + + var safeExecuted, dangerExecuted bool + + safeTool := agent.FunctionTool[struct{}]( + "safe_action", + "A safe action", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + safeExecuted = true + return agent.ToolResult{Content: "safe done"}, nil + }, + ) + + dangerTool := agent.FunctionTool[struct{}]( + "danger_action", + "A dangerous action", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + dangerExecuted = true + return agent.ToolResult{Content: "danger done"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse( + llm.ToolCall{ + ID: "tc_safe", + Function: llm.FunctionCall{Name: "safe_action", Arguments: `{}`}, + }, + llm.ToolCall{ + ID: "tc_danger", + Function: llm.FunctionCall{Name: "danger_action", Arguments: `{}`}, + }, + ), + stopResponse("Both actions completed."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(safeTool, dangerTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"danger_action"}, + }), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Do both")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.Len(t, interrupted.ToolCalls, 2) + assert.Len(t, interrupted.PendingApprovals, 1) + assert.Equal(t, "danger_action", interrupted.PendingApprovals[0].Function.Name) + assert.False(t, safeExecuted) + assert.False(t, dangerExecuted) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc_danger": {Approved: true}, + }, + }, + ) + + require.NoError(t, err) + assert.True(t, safeExecuted) + assert.True(t, dangerExecuted) + assert.Equal(t, "Both actions completed.", result.FinalMessage().Text()) + }, + ) +} + +func TestResume(t *testing.T) { + t.Parallel() + + t.Run( + "handoff tool call with approval", + func(t *testing.T) { + t.Parallel() + + var handoffCallbackFired bool + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }), + stopResponse("Specialist handled it."), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + specialist, + agent.WithOnHandoff(func(_ context.Context) error { + handoffCallbackFired = true + return nil + }), + ), + ), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"transfer_to_specialist"}, + }), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("Help me")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.False(t, handoffCallbackFired) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: true}, + }, + }, + ) + + require.NoError(t, err) + assert.True(t, handoffCallbackFired) + assert.Equal(t, "Specialist handled it.", result.FinalMessage().Text()) + assert.Equal(t, "specialist", result.LastAgent.Name()) + }, + ) + + t.Run( + "carries forward usage and turns", + func(t *testing.T) { + t.Parallel() + + deleteTool := agent.FunctionTool[struct{}]( + "delete_account", + "Deletes the user account", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + return agent.ToolResult{Content: "deleted"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, + }), + stopResponse("Account deleted."), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(deleteTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"delete_account"}, + }), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("Delete my account")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.Equal(t, 1, interrupted.Turns) + assert.Equal(t, 10, interrupted.Usage.InputTokens) + assert.Equal(t, 5, interrupted.Usage.OutputTokens) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: true}, + }, + }, + ) + + require.NoError(t, err) + assert.Equal(t, 2, result.Turns, "turns should include the interrupted turn") + assert.Equal(t, 20, result.Usage.InputTokens, "usage should include pre-interruption tokens") + assert.Equal(t, 10, result.Usage.OutputTokens, "usage should include pre-interruption tokens") + }, + ) + + t.Run( + "rejected handoff stays with current agent", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }), + stopResponse("OK, I will handle it myself."), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffs(specialist), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"transfer_to_specialist"}, + }), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("Help me")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: false, Message: "User declined the transfer."}, + }, + }, + ) + + require.NoError(t, err) + assert.Equal(t, "OK, I will handle it myself.", result.FinalMessage().Text()) + assert.Equal(t, "triage", result.LastAgent.Name()) + }, + ) +} + +func TestRunTyped(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + Country string `json:"country"` + } + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse(`{"name":"Stripe","country":"US"}`), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInstructions("Return structured info."), + ) + + result, err := agent.RunTyped[Info]( + context.Background(), + ag, + []llm.Message{userMessage("Tell me about Stripe")}, + ) + + require.NoError(t, err) + assert.Equal(t, "Stripe", result.Output.Name) + assert.Equal(t, "US", result.Output.Country) +} + +func TestRunStreamed(t *testing.T) { + t.Parallel() + + t.Run( + "streams delta events and completes", + func(t *testing.T) { + t.Parallel() + + mockStream := &mockChatStream{ + events: []llm.ChatCompletionStreamEvent{ + {Delta: llm.MessageDelta{Content: "Hello"}}, + {Delta: llm.MessageDelta{Content: " world"}}, + { + Delta: llm.MessageDelta{Content: "!"}, + Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3}, + FinishReason: finishReasonPtr(llm.FinishReasonStop), + }, + }, + } + + streamProvider := &mockStreamProvider{stream: mockStream} + client := llm.NewClient(streamProvider, "test") + + ag := agent.New( + "assistant", + client, + agent.WithModel("test-model"), + agent.WithInstructions("Be brief."), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + var deltas []string + var gotComplete bool + + for ev := range sr.Events { + switch ev.Type { + case agent.StreamEventLLMDelta: + deltas = append(deltas, ev.Delta) + case agent.StreamEventComplete: + gotComplete = true + } + } + + result, err := sr.Wait() + require.NoError(t, err) + assert.True(t, gotComplete) + assert.Equal(t, []string{"Hello", " world", "!"}, deltas) + assert.Equal(t, 1, result.Turns) + assert.Equal(t, "Hello world!", result.FinalMessage().Text()) + }, + ) + + t.Run( + "with tool calls", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + tool := agent.FunctionTool[Params]( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + stream1 := &mockChatStream{ + events: []llm.ChatCompletionStreamEvent{ + { + Delta: llm.MessageDelta{ + ToolCalls: []llm.ToolCallDelta{ + {Index: 0, ID: "tc_1", Name: "noop", Arguments: `{}`}, + }, + }, + Usage: &llm.Usage{InputTokens: 10, OutputTokens: 5}, + FinishReason: finishReasonPtr(llm.FinishReasonToolCalls), + }, + }, + } + + stream2 := &mockChatStream{ + events: []llm.ChatCompletionStreamEvent{ + {Delta: llm.MessageDelta{Content: "Done!"}}, + { + Delta: llm.MessageDelta{}, + Usage: &llm.Usage{InputTokens: 15, OutputTokens: 3}, + FinishReason: finishReasonPtr(llm.FinishReasonStop), + }, + }, + } + + streamProvider := &mockMultiStreamProvider{ + streams: []llm.ChatCompletionStream{stream1, stream2}, + } + + client := llm.NewClient(streamProvider, "test") + + ag := agent.New( + "assistant", + client, + agent.WithModel("test-model"), + agent.WithTools(tool), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + var gotToolStart, gotToolEnd, gotComplete bool + for ev := range sr.Events { + switch ev.Type { + case agent.StreamEventToolStart: + gotToolStart = true + case agent.StreamEventToolEnd: + gotToolEnd = true + case agent.StreamEventComplete: + gotComplete = true + } + } + + result, err := sr.Wait() + require.NoError(t, err) + assert.True(t, gotToolStart) + assert.True(t, gotToolEnd) + assert.True(t, gotComplete) + assert.Equal(t, 2, result.Turns) + assert.Equal(t, "Done!", result.FinalMessage().Text()) + }, + ) + + t.Run( + "session save failure emits error not complete", + func(t *testing.T) { + t.Parallel() + + mockStream := &mockChatStream{ + events: []llm.ChatCompletionStreamEvent{ + {Delta: llm.MessageDelta{Content: "Hello"}}, + { + Delta: llm.MessageDelta{Content: "!"}, + Usage: &llm.Usage{InputTokens: 10, OutputTokens: 2}, + FinishReason: finishReasonPtr(llm.FinishReasonStop), + }, + }, + } + + streamProvider := &mockStreamProvider{stream: mockStream} + + ag := agent.New( + "assistant", + llm.NewClient(streamProvider, "test"), + agent.WithModel("test-model"), + agent.WithSession(&failingSession{}, "sess-1"), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("Hi")}, + ) + + var gotComplete, gotError bool + for ev := range sr.Events { + switch ev.Type { + case agent.StreamEventComplete: + gotComplete = true + case agent.StreamEventError: + gotError = true + } + } + + _, err := sr.Wait() + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot save session") + assert.False(t, gotComplete, "StreamEventComplete should not be emitted when session save fails") + assert.True(t, gotError, "StreamEventError should be emitted when session save fails") + }, + ) +} + +func TestClone(t *testing.T) { + t.Parallel() + + t.Run( + "preserves name and overrides instructions", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + original := agent.New( + "original", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInstructions("original instructions"), + ) + + cloned := original.Clone( + agent.WithInstructions("cloned instructions"), + ) + + assert.Equal(t, "original", cloned.Name()) + }, + ) + + t.Run( + "does not mutate original tool list", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + tool1 := agent.FunctionTool[Params]( + "t1", + "desc", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + tool2 := agent.FunctionTool[Params]( + "t2", + "desc", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + stopResponse("ok"), + }, + } + + original := agent.New( + "original", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool1), + ) + + cloned := original.Clone(agent.WithTools(tool2)) + + originalResult, err := original.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + require.NoError(t, err) + assert.Equal(t, "ok", originalResult.FinalMessage().Text()) + + clonedResult, err := cloned.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + require.NoError(t, err) + assert.Equal(t, "ok", clonedResult.FinalMessage().Text()) + }, + ) +} + +func TestWithMaxTurns(t *testing.T) { + t.Parallel() + + t.Run( + "zero clamps to minimum", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithMaxTurns(0), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.Equal(t, 1, result.Turns) + }, + ) + + t.Run( + "negative clamps to minimum", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithMaxTurns(-5), + ) + + result, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.NoError(t, err) + assert.Equal(t, 1, result.Turns) + }, + ) +} + +func TestGenerateSchema_EmbeddedStruct(t *testing.T) { + t.Parallel() + + type Base struct { + ID string `json:"id" jsonschema:"unique identifier"` + Kind string `json:"kind"` + } + type Params struct { + Base + Name string `json:"name"` + } + + tool := agent.FunctionTool[Params]( + "create", + "Create item", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + var schema map[string]any + require.NoError(t, json.Unmarshal(tool.Definition().Parameters, &schema)) + assert.Equal(t, "object", schema["type"]) + + props := schema["properties"].(map[string]any) + assert.Contains(t, props, "id") + assert.Contains(t, props, "kind") + assert.Contains(t, props, "name") + assert.NotContains(t, props, "Base") + + idProp := props["id"].(map[string]any) + assert.Equal(t, "string", idProp["type"]) + assert.Equal(t, "unique identifier", idProp["description"]) + + required := schema["required"].([]any) + assert.Contains(t, required, "id") + assert.Contains(t, required, "kind") + assert.Contains(t, required, "name") +} + +func TestRun_InputGuardrailError(t *testing.T) { + t.Parallel() + + guardrail := &errorGuardrail{err: errors.New("check failed")} + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("should not reach"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInputGuardrails(guardrail), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot run input guardrail") + assert.Contains(t, err.Error(), "check failed") + assert.Equal(t, 0, provider.calls) +} + +func TestRun_OutputGuardrailError(t *testing.T) { + t.Parallel() + + guardrail := &errorOutputGuardrail{err: errors.New("output check failed")} + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("hello"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithOutputGuardrails(guardrail), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot run output guardrail") + assert.Contains(t, err.Error(), "output check failed") +} + +func TestRun_HandoffCallbackError(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "transfer_to_billing", Arguments: `{}`}, + }), + }, + } + + client := newTestClient(provider) + + billing := agent.New( + "billing", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + billing, + agent.WithOnHandoff(func(_ context.Context) error { + return errors.New("handoff setup failed") + }), + ), + ), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("help")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot execute handoff callback") + assert.Contains(t, err.Error(), "handoff setup failed") +} + +func TestRun_ContentFilterFinishReason(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + { + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: ""}}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 0}, + FinishReason: llm.FinishReasonContentFilter, + }, + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("bad prompt")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "content was filtered") +} + +func TestRun_UnknownToolCall(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "nonexistent_tool", Arguments: `{}`}, + }), + }, + } + + type Params struct{} + tool := agent.FunctionTool[Params]( + "real_tool", + "A real tool", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tool), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "unknown tool") + assert.Contains(t, err.Error(), "nonexistent_tool") +} + +func TestRun_SessionLoadFailure(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("should not reach"), + }, + } + + ag := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithSession(&errorLoadSession{}, "sess-1"), + ) + + _, err := ag.Run( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot load session") + assert.Equal(t, 0, provider.calls) +} + +func TestClone_WithApprovalConfig(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "delete", Arguments: `{}`}, + }), + toolCallResponse(llm.ToolCall{ + ID: "tc2", + Function: llm.FunctionCall{Name: "delete", Arguments: `{}`}, + }), + }, + } + + deleteTool := agent.FunctionTool[struct{}]( + "delete", + "Delete something", + func(_ context.Context, _ struct{}) (agent.ToolResult, error) { + return agent.ToolResult{Content: "deleted"}, nil + }, + ) + + original := agent.New( + "assistant", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(deleteTool), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"delete"}, + }), + ) + + cloned := original.Clone() + + _, err := cloned.Run( + context.Background(), + []llm.Message{userMessage("delete it")}, + ) + + require.Error(t, err) + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + assert.Len(t, interrupted.PendingApprovals, 1) + + _, err = original.Run( + context.Background(), + []llm.Message{userMessage("delete it")}, + ) + + require.Error(t, err) + require.ErrorAs(t, err, &interrupted) + assert.Len(t, interrupted.PendingApprovals, 1) +} + +func TestRun_HandoffWithPreHandoffTools(t *testing.T) { + t.Parallel() + + t.Run( + "tools before handoff are executed in order", + func(t *testing.T) { + t.Parallel() + + var executionOrder []string + + type Params struct{} + tool1 := agent.FunctionTool[Params]( + "prepare", + "Prepare data", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + executionOrder = append(executionOrder, "prepare") + return agent.ToolResult{Content: "prepared"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse( + llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "prepare", Arguments: `{}`}, + }, + llm.ToolCall{ + ID: "tc_2", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }, + ), + stopResponse("Specialist here."), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + router := agent.New( + "router", + client, + agent.WithModel("test-model"), + agent.WithTools(tool1), + agent.WithHandoffs(specialist), + ) + + result, err := router.Run( + context.Background(), + []llm.Message{userMessage("prepare and transfer")}, + ) + + require.NoError(t, err) + assert.Equal(t, []string{"prepare"}, executionOrder) + assert.Equal(t, "Specialist here.", result.FinalMessage().Text()) + assert.Equal(t, "specialist", result.LastAgent.Name()) + }, + ) + + t.Run( + "pre-handoff tool error aborts handoff", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + failingTool := agent.FunctionTool[Params]( + "prepare", + "Prepare data", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, errors.New("preparation failed") + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse( + llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "prepare", Arguments: `{}`}, + }, + llm.ToolCall{ + ID: "tc_2", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }, + ), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + router := agent.New( + "router", + client, + agent.WithModel("test-model"), + agent.WithTools(failingTool), + agent.WithHandoffs(specialist), + ) + + _, err := router.Run( + context.Background(), + []llm.Message{userMessage("prepare and transfer")}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot execute tool") + assert.Contains(t, err.Error(), "preparation failed") + }, + ) +} + +func TestRunStreamed_StreamError(t *testing.T) { + t.Parallel() + + errStream := &errorChatStream{err: errors.New("stream broke")} + streamProvider := &mockStreamProvider{stream: errStream} + + ag := agent.New( + "assistant", + llm.NewClient(streamProvider, "test"), + agent.WithModel("test-model"), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + for range sr.Events { + } + + _, err := sr.Wait() + require.Error(t, err) + assert.Contains(t, err.Error(), "stream broke") +} + +func TestRunStreamed_StreamProviderError(t *testing.T) { + t.Parallel() + + errProvider := &errorStreamProvider{err: errors.New("connection refused")} + + ag := agent.New( + "assistant", + llm.NewClient(errProvider, "test"), + agent.WithModel("test-model"), + ) + + sr := ag.RunStreamed( + context.Background(), + []llm.Message{userMessage("test")}, + ) + + for range sr.Events { + } + + _, err := sr.Wait() + require.Error(t, err) + assert.Contains(t, err.Error(), "connection refused") +} + +func TestResume_HandoffWithInputFilter(t *testing.T) { + t.Parallel() + + var receivedMessages int + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc1", + Function: llm.FunctionCall{Name: "transfer_to_specialist", Arguments: `{}`}, + }), + stopResponse("Filtered specialist here."), + }, + } + + client := newTestClient(provider) + + specialist := agent.New( + "specialist", + client, + agent.WithModel("test-model"), + ) + + triage := agent.New( + "triage", + client, + agent.WithModel("test-model"), + agent.WithHandoffConfigs( + agent.HandoffTo( + specialist, + agent.WithHandoffInputFilter(func(data agent.HandoffInputData) []llm.Message { + receivedMessages = len(data.InputHistory) + len(data.NewItems) + return data.NewItems + }), + ), + ), + agent.WithApproval(agent.ApprovalConfig{ + ToolNames: []string{"transfer_to_specialist"}, + }), + ) + + _, err := triage.Run( + context.Background(), + []llm.Message{userMessage("Help me")}, + ) + + var interrupted *agent.InterruptedError + require.ErrorAs(t, err, &interrupted) + + result, err := agent.Resume( + context.Background(), + interrupted, + agent.ResumeInput{ + Approvals: map[string]agent.ApprovalResult{ + "tc1": {Approved: true}, + }, + }, + ) + + require.NoError(t, err) + assert.Greater(t, receivedMessages, 0) + assert.Equal(t, "Filtered specialist here.", result.FinalMessage().Text()) + assert.Equal(t, "specialist", result.LastAgent.Name()) +} diff --git a/pkg/agent/agent_tool.go b/pkg/agent/agent_tool.go new file mode 100644 index 000000000..54b32adb1 --- /dev/null +++ b/pkg/agent/agent_tool.go @@ -0,0 +1,87 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "encoding/json" + "fmt" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + agentTool struct { + agent *Agent + toolName string + description string + schema json.RawMessage + } + + agentToolParams struct { + Input string `json:"input" jsonschema:"The input to send to the agent"` + } +) + +var ( + agentToolParamsSchema = jsonSchemaFor[agentToolParams]() +) + +func newAgentTool(agent *Agent, name, description string) *agentTool { + return &agentTool{ + agent: agent, + toolName: name, + description: description, + schema: agentToolParamsSchema, + } +} + +func (t *agentTool) Name() string { return t.toolName } + +func (t *agentTool) Definition() llm.Tool { + return llm.Tool{ + Name: t.toolName, + Description: t.description, + Parameters: t.schema, + } +} + +func (t *agentTool) Execute(ctx context.Context, arguments string) (ToolResult, error) { + var params agentToolParams + + if err := json.Unmarshal([]byte(arguments), ¶ms); err != nil { + return ToolResult{ + Content: fmt.Sprintf("Invalid parameters: %s", err.Error()), + IsError: true, + }, nil + } + + result, err := t.agent.Run( + ctx, + []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{ + llm.TextPart{Text: params.Input}, + }, + }, + }, + ) + if err != nil { + return ToolResult{}, err + } + + return ToolResult{Content: result.FinalMessage().Text()}, nil +} diff --git a/pkg/agent/agent_tool_test.go b/pkg/agent/agent_tool_test.go new file mode 100644 index 000000000..2164ec5fd --- /dev/null +++ b/pkg/agent/agent_tool_test.go @@ -0,0 +1,391 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +func TestAgentTool_Name(t *testing.T) { + t.Parallel() + + t.Run( + "returns custom tool name not agent name", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "geography", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("geo_expert", "Ask geography questions") + assert.Equal(t, "geo_expert", tool.Name()) + }, + ) + + t.Run( + "different AsTool calls return different names", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "helper", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool1 := ag.AsTool("tool_a", "First tool") + tool2 := ag.AsTool("tool_b", "Second tool") + + assert.Equal(t, "tool_a", tool1.Name()) + assert.Equal(t, "tool_b", tool2.Name()) + }, + ) +} + +func TestAgentTool_Definition(t *testing.T) { + t.Parallel() + + t.Run( + "returns name and description", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "sub", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("my_tool", "Does something useful.") + def := tool.Definition() + + assert.Equal(t, "my_tool", def.Name) + assert.Equal(t, "Does something useful.", def.Description) + }, + ) + + t.Run( + "schema contains input string property", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "sub", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("delegate", "Delegate work.") + def := tool.Definition() + + require.NotNil(t, def.Parameters) + + var schema map[string]any + require.NoError(t, json.Unmarshal(def.Parameters, &schema)) + + assert.Equal(t, "object", schema["type"]) + + props, ok := schema["properties"].(map[string]any) + require.True(t, ok) + assert.Contains(t, props, "input") + + inputProp := props["input"].(map[string]any) + assert.Equal(t, "string", inputProp["type"]) + }, + ) + + t.Run( + "schema requires input field", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "sub", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("ask", "Ask a question.") + def := tool.Definition() + + var schema map[string]any + require.NoError(t, json.Unmarshal(def.Parameters, &schema)) + + required, ok := schema["required"].([]any) + require.True(t, ok) + assert.Contains(t, required, "input") + }, + ) +} + +func TestAgentTool_Execute(t *testing.T) { + t.Parallel() + + t.Run( + "runs sub-agent and returns final message", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("The capital of France is Paris."), + }, + } + + ag := agent.New( + "geography", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithInstructions("You are a geography expert."), + ) + + tool := ag.AsTool("geo_expert", "Ask geography questions.") + result, err := tool.Execute( + context.Background(), + `{"input":"What is the capital of France?"}`, + ) + + require.NoError(t, err) + assert.Equal(t, "The capital of France is Paris.", result.Content) + assert.False(t, result.IsError) + assert.Equal(t, 1, provider.calls) + }, + ) + + t.Run( + "invalid JSON returns tool error not Go error", + func(t *testing.T) { + t.Parallel() + + ag := agent.New( + "sub", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("sub_tool", "A sub-agent tool.") + result, err := tool.Execute(context.Background(), `{bad json}`) + + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Invalid parameters") + }, + ) + + t.Run( + "empty JSON object returns tool error", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + ag := agent.New( + "sub", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("sub_tool", "A sub-agent tool.") + result, err := tool.Execute(context.Background(), `{}`) + + require.NoError(t, err) + assert.False(t, result.IsError) + }, + ) + + t.Run( + "sub-agent error propagates as Go error", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{}, + } + + ag := agent.New( + "sub", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("helper", "A helper agent.") + _, err := tool.Execute( + context.Background(), + `{"input":"hello"}`, + ) + + require.Error(t, err) + }, + ) + + t.Run( + "context is forwarded to sub-agent", + func(t *testing.T) { + t.Parallel() + + type AppCtx struct { + TenantID string + } + + var captured string + + type Params struct{} + tenantTool := agent.FunctionTool[Params]( + "get_tenant", + "Get tenant", + func(ctx context.Context, _ Params) (agent.ToolResult, error) { + rc := agent.RunContextFrom[*AppCtx](ctx) + captured = rc.TenantID + return agent.ToolResult{Content: rc.TenantID}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "get_tenant", Arguments: `{}`}, + }), + stopResponse("tenant is t_789"), + }, + } + + subAgent := agent.New( + "tenant_agent", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(tenantTool), + ) + + tool := subAgent.AsTool("check_tenant", "Check tenant info.") + + ctx := agent.WithRunContext( + context.Background(), + &AppCtx{TenantID: "t_789"}, + ) + + result, err := tool.Execute(ctx, `{"input":"what tenant?"}`) + + require.NoError(t, err) + assert.Equal(t, "tenant is t_789", result.Content) + assert.Equal(t, "t_789", captured) + }, + ) + + t.Run( + "sub-agent with tool calls completes multi-turn", + func(t *testing.T) { + t.Parallel() + + type Params struct { + Expr string `json:"expr"` + } + + calcTool := agent.FunctionTool[Params]( + "calc", + "Calculate expression", + func(_ context.Context, p Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "42"}, nil + }, + ) + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + toolCallResponse(llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "calc", Arguments: `{"expr":"6*7"}`}, + }), + stopResponse("The answer is 42."), + }, + } + + subAgent := agent.New( + "math", + newTestClient(provider), + agent.WithModel("test-model"), + agent.WithTools(calcTool), + ) + + tool := subAgent.AsTool("math_expert", "Ask math questions.") + result, err := tool.Execute( + context.Background(), + `{"input":"What is 6 times 7?"}`, + ) + + require.NoError(t, err) + assert.Equal(t, "The answer is 42.", result.Content) + assert.False(t, result.IsError) + assert.Equal(t, 2, provider.calls) + }, + ) + + t.Run( + "extra JSON fields are ignored", + func(t *testing.T) { + t.Parallel() + + provider := &mockProvider{ + responses: []*llm.ChatCompletionResponse{ + stopResponse("ok"), + }, + } + + ag := agent.New( + "sub", + newTestClient(provider), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("sub_tool", "Sub tool.") + result, err := tool.Execute( + context.Background(), + `{"input":"hello","extra":"ignored"}`, + ) + + require.NoError(t, err) + assert.False(t, result.IsError) + assert.Equal(t, "ok", result.Content) + }, + ) +} + +func TestAgentTool_InterfaceSatisfaction(t *testing.T) { + t.Parallel() + + ag := agent.New( + "sub", + newTestClient(&mockProvider{}), + agent.WithModel("test-model"), + ) + + tool := ag.AsTool("test_tool", "Test tool") + + assert.Implements(t, (*agent.Tool)(nil), tool) + assert.Implements(t, (*agent.ToolDescriptor)(nil), tool) +} diff --git a/pkg/agent/approval.go b/pkg/agent/approval.go new file mode 100644 index 000000000..b3d7b32ef --- /dev/null +++ b/pkg/agent/approval.go @@ -0,0 +1,64 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + ApprovalConfig struct { + ToolNames []string + ShouldApprove func(ctx context.Context, toolCall llm.ToolCall) bool + + toolNameSet map[string]struct{} + } + + ApprovalResult struct { + Approved bool + Message string + } + + ResumeInput struct { + Approvals map[string]ApprovalResult + } +) + +func buildToolNameSet(names []string) map[string]struct{} { + if len(names) == 0 { + return nil + } + + set := make(map[string]struct{}, len(names)) + for _, name := range names { + set[name] = struct{}{} + } + return set +} + +func (c *ApprovalConfig) requiresApproval(ctx context.Context, tc llm.ToolCall) bool { + if c == nil { + return false + } + + if c.ShouldApprove != nil { + return c.ShouldApprove(ctx, tc) + } + + _, ok := c.toolNameSet[tc.Function.Name] + return ok +} diff --git a/pkg/agent/approval_test.go b/pkg/agent/approval_test.go new file mode 100644 index 000000000..7c3817190 --- /dev/null +++ b/pkg/agent/approval_test.go @@ -0,0 +1,203 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "go.probo.inc/probo/pkg/llm" +) + +func TestBuildToolNameSet(t *testing.T) { + t.Parallel() + + t.Run( + "nil input returns nil", + func(t *testing.T) { + t.Parallel() + + assert.Nil(t, buildToolNameSet(nil)) + }, + ) + + t.Run( + "empty slice returns nil", + func(t *testing.T) { + t.Parallel() + + assert.Nil(t, buildToolNameSet([]string{})) + }, + ) + + t.Run( + "single name", + func(t *testing.T) { + t.Parallel() + + set := buildToolNameSet([]string{"delete"}) + + assert.Len(t, set, 1) + _, ok := set["delete"] + assert.True(t, ok) + }, + ) + + t.Run( + "multiple names", + func(t *testing.T) { + t.Parallel() + + set := buildToolNameSet([]string{"delete", "update", "create"}) + + assert.Len(t, set, 3) + for _, name := range []string{"delete", "update", "create"} { + _, ok := set[name] + assert.True(t, ok, "expected set to contain %q", name) + } + }, + ) + + t.Run( + "duplicate names are deduplicated", + func(t *testing.T) { + t.Parallel() + + set := buildToolNameSet([]string{"delete", "delete", "update"}) + + assert.Len(t, set, 2) + }, + ) +} + +func TestApprovalConfig_RequiresApproval(t *testing.T) { + t.Parallel() + + tc := llm.ToolCall{ + ID: "tc_1", + Function: llm.FunctionCall{Name: "delete_user", Arguments: `{}`}, + } + + t.Run( + "nil config returns false", + func(t *testing.T) { + t.Parallel() + + var c *ApprovalConfig + assert.False(t, c.requiresApproval(context.Background(), tc)) + }, + ) + + t.Run( + "empty config returns false", + func(t *testing.T) { + t.Parallel() + + c := &ApprovalConfig{} + assert.False(t, c.requiresApproval(context.Background(), tc)) + }, + ) + + t.Run( + "tool name in set returns true", + func(t *testing.T) { + t.Parallel() + + c := &ApprovalConfig{ + toolNameSet: map[string]struct{}{ + "delete_user": {}, + }, + } + + assert.True(t, c.requiresApproval(context.Background(), tc)) + }, + ) + + t.Run( + "tool name not in set returns false", + func(t *testing.T) { + t.Parallel() + + c := &ApprovalConfig{ + toolNameSet: map[string]struct{}{ + "list_users": {}, + }, + } + + assert.False(t, c.requiresApproval(context.Background(), tc)) + }, + ) + + t.Run( + "ShouldApprove takes precedence over tool name set", + func(t *testing.T) { + t.Parallel() + + c := &ApprovalConfig{ + toolNameSet: map[string]struct{}{ + "delete_user": {}, + }, + ShouldApprove: func(_ context.Context, _ llm.ToolCall) bool { + return false + }, + } + + assert.False(t, c.requiresApproval(context.Background(), tc)) + }, + ) + + t.Run( + "ShouldApprove receives context and tool call", + func(t *testing.T) { + t.Parallel() + + type ctxKey struct{} + ctx := context.WithValue(context.Background(), ctxKey{}, "marker") + + var capturedCtx context.Context + var capturedTC llm.ToolCall + + c := &ApprovalConfig{ + ShouldApprove: func(ctx context.Context, tc llm.ToolCall) bool { + capturedCtx = ctx + capturedTC = tc + return true + }, + } + + result := c.requiresApproval(ctx, tc) + + assert.True(t, result) + assert.Equal(t, "marker", capturedCtx.Value(ctxKey{})) + assert.Equal(t, "delete_user", capturedTC.Function.Name) + }, + ) + + t.Run( + "ShouldApprove returning true requires approval", + func(t *testing.T) { + t.Parallel() + + c := &ApprovalConfig{ + ShouldApprove: func(_ context.Context, _ llm.ToolCall) bool { + return true + }, + } + + assert.True(t, c.requiresApproval(context.Background(), tc)) + }, + ) +} diff --git a/pkg/agent/context.go b/pkg/agent/context.go new file mode 100644 index 000000000..140d60b2a --- /dev/null +++ b/pkg/agent/context.go @@ -0,0 +1,53 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "fmt" +) + +type runContextKey struct{} + +func WithRunContext(ctx context.Context, val any) context.Context { + return context.WithValue(ctx, runContextKey{}, val) +} + +func RunContextFrom[C any](ctx context.Context) C { + val := ctx.Value(runContextKey{}) + if val == nil { + var zero C + panic(fmt.Sprintf("agent: no run context found (expected %T)", zero)) + } + + typed, ok := val.(C) + if !ok { + var zero C + panic(fmt.Sprintf("agent: run context type mismatch: stored %T, requested %T", val, zero)) + } + + return typed +} + +func TryRunContextFrom[C any](ctx context.Context) (C, bool) { + val := ctx.Value(runContextKey{}) + if val == nil { + var zero C + return zero, false + } + + typed, ok := val.(C) + return typed, ok +} diff --git a/pkg/agent/context_test.go b/pkg/agent/context_test.go new file mode 100644 index 000000000..6fb6cc2c7 --- /dev/null +++ b/pkg/agent/context_test.go @@ -0,0 +1,198 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" +) + +type testRunContext struct { + UserID string + Locale string +} + +func TestWithRunContext(t *testing.T) { + t.Parallel() + + t.Run( + "stores value retrievable by RunContextFrom", + func(t *testing.T) { + t.Parallel() + + rc := testRunContext{UserID: "u_123", Locale: "en"} + ctx := agent.WithRunContext(context.Background(), rc) + + got := agent.RunContextFrom[testRunContext](ctx) + assert.Equal(t, "u_123", got.UserID) + assert.Equal(t, "en", got.Locale) + }, + ) + + t.Run( + "overwrites previous run context", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), testRunContext{UserID: "first"}) + ctx = agent.WithRunContext(ctx, testRunContext{UserID: "second"}) + + got := agent.RunContextFrom[testRunContext](ctx) + assert.Equal(t, "second", got.UserID) + }, + ) +} + +func TestRunContextFrom(t *testing.T) { + t.Parallel() + + t.Run( + "returns stored value with correct type", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), "hello") + + got := agent.RunContextFrom[string](ctx) + assert.Equal(t, "hello", got) + }, + ) + + t.Run( + "panics when no run context is set", + func(t *testing.T) { + t.Parallel() + + require.Panics( + t, + func() { + agent.RunContextFrom[string](context.Background()) + }, + ) + }, + ) + + t.Run( + "panic message includes expected type when missing", + func(t *testing.T) { + t.Parallel() + + assert.PanicsWithValue( + t, + "agent: no run context found (expected string)", + func() { + agent.RunContextFrom[string](context.Background()) + }, + ) + }, + ) + + t.Run( + "panics on type mismatch", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), 42) + + require.Panics( + t, + func() { + agent.RunContextFrom[string](ctx) + }, + ) + }, + ) + + t.Run( + "panic message includes both types on mismatch", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), 42) + + assert.PanicsWithValue( + t, + "agent: run context type mismatch: stored int, requested string", + func() { + agent.RunContextFrom[string](ctx) + }, + ) + }, + ) +} + +func TestTryRunContextFrom(t *testing.T) { + t.Parallel() + + t.Run( + "returns value and true when context is set", + func(t *testing.T) { + t.Parallel() + + rc := testRunContext{UserID: "u_456", Locale: "fr"} + ctx := agent.WithRunContext(context.Background(), rc) + + got, ok := agent.TryRunContextFrom[testRunContext](ctx) + + assert.True(t, ok) + assert.Equal(t, "u_456", got.UserID) + assert.Equal(t, "fr", got.Locale) + }, + ) + + t.Run( + "returns zero value and false when no context is set", + func(t *testing.T) { + t.Parallel() + + got, ok := agent.TryRunContextFrom[testRunContext](context.Background()) + + assert.False(t, ok) + assert.Equal(t, testRunContext{}, got) + }, + ) + + t.Run( + "returns zero value and false on type mismatch", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), 42) + + got, ok := agent.TryRunContextFrom[string](ctx) + + assert.False(t, ok) + assert.Equal(t, "", got) + }, + ) + + t.Run( + "works with interface types", + func(t *testing.T) { + t.Parallel() + + ctx := agent.WithRunContext(context.Background(), "a string value") + + got, ok := agent.TryRunContextFrom[any](ctx) + + assert.True(t, ok) + assert.Equal(t, "a string value", got) + }, + ) +} diff --git a/pkg/agent/errors.go b/pkg/agent/errors.go new file mode 100644 index 000000000..df59c79b3 --- /dev/null +++ b/pkg/agent/errors.go @@ -0,0 +1,71 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "fmt" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + MaxTurnsExceededError struct { + MaxTurns int + } + + InputGuardrailTrippedError struct { + Guardrail string + Message string + } + + OutputGuardrailTrippedError struct { + Guardrail string + Message string + } + + InterruptedError struct { + ToolCalls []llm.ToolCall + PendingApprovals []llm.ToolCall + Agent *Agent + Messages []llm.Message + Usage llm.Usage + Turns int + } + + needsApprovalError struct { + allToolCalls []llm.ToolCall + pendingApprovals []llm.ToolCall + } +) + +func (e *MaxTurnsExceededError) Error() string { + return fmt.Sprintf("agent exceeded maximum number of turns (%d)", e.MaxTurns) +} + +func (e *InputGuardrailTrippedError) Error() string { + return fmt.Sprintf("input guardrail %q tripped: %s", e.Guardrail, e.Message) +} + +func (e *OutputGuardrailTrippedError) Error() string { + return fmt.Sprintf("output guardrail %q tripped: %s", e.Guardrail, e.Message) +} + +func (e *InterruptedError) Error() string { + return fmt.Sprintf("run interrupted: %d tool call(s) require approval", len(e.PendingApprovals)) +} + +func (e *needsApprovalError) Error() string { + return fmt.Sprintf("%d tool call(s) require approval", len(e.pendingApprovals)) +} diff --git a/pkg/agent/guardrail.go b/pkg/agent/guardrail.go new file mode 100644 index 000000000..2bf42618e --- /dev/null +++ b/pkg/agent/guardrail.go @@ -0,0 +1,38 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + GuardrailResult struct { + Tripwire bool + Message string + } + + InputGuardrail interface { + Name() string + Check(ctx context.Context, messages []llm.Message) (*GuardrailResult, error) + } + + OutputGuardrail interface { + Name() string + Check(ctx context.Context, message llm.Message) (*GuardrailResult, error) + } +) diff --git a/pkg/agent/handoff.go b/pkg/agent/handoff.go new file mode 100644 index 000000000..cec8c83eb --- /dev/null +++ b/pkg/agent/handoff.go @@ -0,0 +1,141 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "fmt" + "strings" + "unicode" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + HandoffInputData struct { + InputHistory []llm.Message + NewItems []llm.Message + } + + HandoffInputFilter func(data HandoffInputData) []llm.Message + + HandoffOption func(*Handoff) + + Handoff struct { + Agent *Agent + ToolName string + ToolDescription string + InputFilter HandoffInputFilter + OnHandoff func(ctx context.Context) error + } + + handoffParams struct{} +) + +var ( + handoffParamsSchema = jsonSchemaFor[handoffParams]() +) + +func HandoffTo(agent *Agent, opts ...HandoffOption) *Handoff { + h := &Handoff{Agent: agent} + for _, opt := range opts { + opt(h) + } + return h +} + +func WithHandoffToolName(name string) HandoffOption { + return func(h *Handoff) { + h.ToolName = name + } +} + +func WithHandoffToolDescription(desc string) HandoffOption { + return func(h *Handoff) { + h.ToolDescription = desc + } +} + +func WithHandoffInputFilter(fn HandoffInputFilter) HandoffOption { + return func(h *Handoff) { + h.InputFilter = fn + } +} + +func WithOnHandoff(fn func(ctx context.Context) error) HandoffOption { + return func(h *Handoff) { + h.OnHandoff = fn + } +} + +func (h *Handoff) toolName() string { + if h.ToolName != "" { + return h.ToolName + } + return "transfer_to_" + sanitizeToolName(h.Agent.name) +} + +func sanitizeToolName(name string) string { + var b strings.Builder + b.Grow(len(name)) + + for _, r := range name { + switch { + case unicode.IsLetter(r), unicode.IsDigit(r): + b.WriteRune(r) + default: + b.WriteByte('_') + } + } + + return strings.ToLower(b.String()) +} + +func (h *Handoff) toolDescription() string { + if h.ToolDescription != "" { + return h.ToolDescription + } + + desc := fmt.Sprintf( + "Transfer the conversation to %s.", + h.Agent.name, + ) + + if h.Agent.handoffDescription != "" { + desc += " " + h.Agent.handoffDescription + } + + return desc +} + +func (h *Handoff) tool() ToolDescriptor { + return &handoffToolAdapter{handoff: h} +} + +type handoffToolAdapter struct { + handoff *Handoff +} + +func (t *handoffToolAdapter) Name() string { + return t.handoff.toolName() +} + +func (t *handoffToolAdapter) Definition() llm.Tool { + return llm.Tool{ + Name: t.handoff.toolName(), + Description: t.handoff.toolDescription(), + Parameters: handoffParamsSchema, + } +} diff --git a/pkg/agent/handoff_test.go b/pkg/agent/handoff_test.go new file mode 100644 index 000000000..9fee403cd --- /dev/null +++ b/pkg/agent/handoff_test.go @@ -0,0 +1,306 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +func TestHandoffTo(t *testing.T) { + t.Parallel() + + t.Run( + "sets the target agent", + func(t *testing.T) { + t.Parallel() + + target := agent.New( + "billing", + newTestClient(&mockProvider{}), + ) + + h := agent.HandoffTo(target) + + assert.Equal(t, target, h.Agent) + }, + ) + + t.Run( + "defaults have zero values", + func(t *testing.T) { + t.Parallel() + + target := agent.New( + "support", + newTestClient(&mockProvider{}), + ) + + h := agent.HandoffTo(target) + + assert.Empty(t, h.ToolName) + assert.Empty(t, h.ToolDescription) + assert.Nil(t, h.InputFilter) + assert.Nil(t, h.OnHandoff) + }, + ) + + t.Run( + "applies multiple options", + func(t *testing.T) { + t.Parallel() + + target := agent.New( + "escalation", + newTestClient(&mockProvider{}), + ) + + filter := func(data agent.HandoffInputData) []llm.Message { + return data.NewItems + } + + callback := func(_ context.Context) error { + return nil + } + + h := agent.HandoffTo( + target, + agent.WithHandoffToolName("escalate"), + agent.WithHandoffToolDescription("Escalate to senior agent"), + agent.WithHandoffInputFilter(filter), + agent.WithOnHandoff(callback), + ) + + assert.Equal(t, target, h.Agent) + assert.Equal(t, "escalate", h.ToolName) + assert.Equal(t, "Escalate to senior agent", h.ToolDescription) + assert.NotNil(t, h.InputFilter) + assert.NotNil(t, h.OnHandoff) + }, + ) +} + +func TestWithHandoffToolName(t *testing.T) { + t.Parallel() + + target := agent.New("billing", newTestClient(&mockProvider{})) + h := agent.HandoffTo(target, agent.WithHandoffToolName("ask_billing")) + + assert.Equal(t, "ask_billing", h.ToolName) +} + +func TestWithHandoffToolDescription(t *testing.T) { + t.Parallel() + + target := agent.New("billing", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithHandoffToolDescription("Route billing questions"), + ) + + assert.Equal(t, "Route billing questions", h.ToolDescription) +} + +func TestWithHandoffInputFilter(t *testing.T) { + t.Parallel() + + t.Run( + "sets the filter function", + func(t *testing.T) { + t.Parallel() + + target := agent.New("specialist", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithHandoffInputFilter(func(data agent.HandoffInputData) []llm.Message { + return data.NewItems + }), + ) + + assert.NotNil(t, h.InputFilter) + }, + ) + + t.Run( + "filter receives correct data and returns filtered messages", + func(t *testing.T) { + t.Parallel() + + history := []llm.Message{ + userMessage("old message"), + assistantMessage("old reply"), + } + newItems := []llm.Message{ + userMessage("new question"), + } + + target := agent.New("specialist", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithHandoffInputFilter(func(data agent.HandoffInputData) []llm.Message { + var filtered []llm.Message + for _, m := range data.NewItems { + if m.Role == llm.RoleUser { + filtered = append(filtered, m) + } + } + return filtered + }), + ) + + result := h.InputFilter(agent.HandoffInputData{ + InputHistory: history, + NewItems: newItems, + }) + + require.Len(t, result, 1) + assert.Equal(t, llm.RoleUser, result[0].Role) + }, + ) + + t.Run( + "filter can combine history and new items", + func(t *testing.T) { + t.Parallel() + + history := []llm.Message{ + userMessage("context"), + } + newItems := []llm.Message{ + userMessage("question"), + } + + target := agent.New("specialist", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithHandoffInputFilter(func(data agent.HandoffInputData) []llm.Message { + all := make([]llm.Message, 0, len(data.InputHistory)+len(data.NewItems)) + all = append(all, data.InputHistory...) + all = append(all, data.NewItems...) + return all + }), + ) + + result := h.InputFilter(agent.HandoffInputData{ + InputHistory: history, + NewItems: newItems, + }) + + assert.Len(t, result, 2) + }, + ) + + t.Run( + "filter can return empty slice", + func(t *testing.T) { + t.Parallel() + + target := agent.New("specialist", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithHandoffInputFilter(func(_ agent.HandoffInputData) []llm.Message { + return nil + }), + ) + + result := h.InputFilter(agent.HandoffInputData{ + InputHistory: []llm.Message{userMessage("hello")}, + NewItems: []llm.Message{userMessage("world")}, + }) + + assert.Empty(t, result) + }, + ) +} + +func TestWithOnHandoff(t *testing.T) { + t.Parallel() + + t.Run( + "sets the callback", + func(t *testing.T) { + t.Parallel() + + target := agent.New("billing", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithOnHandoff(func(_ context.Context) error { + return nil + }), + ) + + assert.NotNil(t, h.OnHandoff) + }, + ) + + t.Run( + "callback is invocable", + func(t *testing.T) { + t.Parallel() + + var called bool + + target := agent.New("billing", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithOnHandoff(func(_ context.Context) error { + called = true + return nil + }), + ) + + err := h.OnHandoff(context.Background()) + require.NoError(t, err) + assert.True(t, called) + }, + ) + + t.Run( + "callback propagates errors", + func(t *testing.T) { + t.Parallel() + + target := agent.New("billing", newTestClient(&mockProvider{})) + h := agent.HandoffTo( + target, + agent.WithOnHandoff(func(_ context.Context) error { + return assert.AnError + }), + ) + + err := h.OnHandoff(context.Background()) + assert.ErrorIs(t, err, assert.AnError) + }, + ) +} + +func TestHandoffTo_OptionOrder(t *testing.T) { + t.Parallel() + + target := agent.New("billing", newTestClient(&mockProvider{})) + + h := agent.HandoffTo( + target, + agent.WithHandoffToolName("first_name"), + agent.WithHandoffToolName("second_name"), + ) + + assert.Equal(t, "second_name", h.ToolName, "last option wins") +} diff --git a/pkg/agent/hooks.go b/pkg/agent/hooks.go new file mode 100644 index 000000000..2abe31b98 --- /dev/null +++ b/pkg/agent/hooks.go @@ -0,0 +1,72 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + + "go.probo.inc/probo/pkg/llm" +) + +// RunHooks receives callbacks on lifecycle events for the entire agent run. +type RunHooks interface { + OnRunStart(ctx context.Context, agent *Agent, messages []llm.Message) + OnRunEnd(ctx context.Context, agent *Agent, result *Result, err error) + OnLLMStart(ctx context.Context, agent *Agent, messages []llm.Message) + OnLLMEnd(ctx context.Context, agent *Agent, response *llm.ChatCompletionResponse, err error) + OnToolStart(ctx context.Context, agent *Agent, tool Tool, arguments string) + OnToolEnd(ctx context.Context, agent *Agent, tool Tool, result ToolResult, err error) + OnHandoff(ctx context.Context, from *Agent, to *Agent) + OnGuardrailTripped(ctx context.Context, agent *Agent, name string, result *GuardrailResult) +} + +// NoOpHooks is a RunHooks implementation that does nothing. +type NoOpHooks struct{} + +var _ RunHooks = NoOpHooks{} + +func (NoOpHooks) OnRunStart(context.Context, *Agent, []llm.Message) {} +func (NoOpHooks) OnRunEnd(context.Context, *Agent, *Result, error) {} +func (NoOpHooks) OnLLMStart(context.Context, *Agent, []llm.Message) {} +func (NoOpHooks) OnLLMEnd(context.Context, *Agent, *llm.ChatCompletionResponse, error) {} +func (NoOpHooks) OnToolStart(context.Context, *Agent, Tool, string) {} +func (NoOpHooks) OnToolEnd(context.Context, *Agent, Tool, ToolResult, error) {} +func (NoOpHooks) OnHandoff(context.Context, *Agent, *Agent) {} +func (NoOpHooks) OnGuardrailTripped(context.Context, *Agent, string, *GuardrailResult) {} + +// AgentHooks receives callbacks on lifecycle events for a specific agent. +// Set via WithAgentHooks on an individual agent. +type AgentHooks interface { + OnStart(ctx context.Context, agent *Agent) + OnEnd(ctx context.Context, agent *Agent, output string) + OnHandoff(ctx context.Context, agent *Agent, source *Agent) + OnToolStart(ctx context.Context, agent *Agent, tool Tool) + OnToolEnd(ctx context.Context, agent *Agent, tool Tool, result ToolResult) + OnLLMStart(ctx context.Context, agent *Agent, messages []llm.Message) + OnLLMEnd(ctx context.Context, agent *Agent, response *llm.ChatCompletionResponse, err error) +} + +// NoOpAgentHooks is an AgentHooks implementation that does nothing. +type NoOpAgentHooks struct{} + +var _ AgentHooks = NoOpAgentHooks{} + +func (NoOpAgentHooks) OnStart(context.Context, *Agent) {} +func (NoOpAgentHooks) OnEnd(context.Context, *Agent, string) {} +func (NoOpAgentHooks) OnHandoff(context.Context, *Agent, *Agent) {} +func (NoOpAgentHooks) OnToolStart(context.Context, *Agent, Tool) {} +func (NoOpAgentHooks) OnToolEnd(context.Context, *Agent, Tool, ToolResult) {} +func (NoOpAgentHooks) OnLLMStart(context.Context, *Agent, []llm.Message) {} +func (NoOpAgentHooks) OnLLMEnd(context.Context, *Agent, *llm.ChatCompletionResponse, error) {} diff --git a/pkg/agent/mcp.go b/pkg/agent/mcp.go new file mode 100644 index 000000000..be5aef13f --- /dev/null +++ b/pkg/agent/mcp.go @@ -0,0 +1,182 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "encoding/json" + "fmt" + "strings" + "sync" + + "github.com/modelcontextprotocol/go-sdk/mcp" + "go.probo.inc/probo/pkg/llm" +) + +type ( + MCPServer struct { + name string + session *mcp.ClientSession + + mu sync.RWMutex + cachedTools []Tool + toolsCached bool + } + + mcpTool struct { + server *MCPServer + name string + description string + inputSchema json.RawMessage + } +) + +func NewMCPServer(name string, session *mcp.ClientSession) *MCPServer { + return &MCPServer{ + name: name, + session: session, + } +} + +func (s *MCPServer) Name() string { + return s.name +} + +func (s *MCPServer) Tools(ctx context.Context) ([]Tool, error) { + s.mu.RLock() + if s.toolsCached { + cp := make([]Tool, len(s.cachedTools)) + copy(cp, s.cachedTools) + s.mu.RUnlock() + return cp, nil + } + s.mu.RUnlock() + + s.mu.Lock() + defer s.mu.Unlock() + + if s.toolsCached { + cp := make([]Tool, len(s.cachedTools)) + copy(cp, s.cachedTools) + return cp, nil + } + + var allTools []*mcp.Tool + var cursor string + + for { + params := &mcp.ListToolsParams{} + if cursor != "" { + params.Cursor = cursor + } + + result, err := s.session.ListTools(ctx, params) + if err != nil { + return nil, fmt.Errorf("cannot list tools from MCP server %q: %w", s.name, err) + } + + allTools = append(allTools, result.Tools...) + + if result.NextCursor == "" { + break + } + cursor = result.NextCursor + } + + tools := make([]Tool, len(allTools)) + for i, t := range allTools { + schema, err := json.Marshal(t.InputSchema) + if err != nil { + return nil, fmt.Errorf("cannot marshal input schema for tool %q: %w", t.Name, err) + } + + tools[i] = &mcpTool{ + server: s, + name: t.Name, + description: t.Description, + inputSchema: schema, + } + } + + s.cachedTools = tools + s.toolsCached = true + + return tools, nil +} + +// ResetCache clears the cached tool definitions, forcing the next call to +// Tools to re-fetch from the MCP server. +func (s *MCPServer) ResetCache() { + s.mu.Lock() + defer s.mu.Unlock() + + s.cachedTools = nil + s.toolsCached = false +} + +func (t *mcpTool) Name() string { return t.name } + +func (t *mcpTool) Definition() llm.Tool { + return llm.Tool{ + Name: t.name, + Description: t.description, + Parameters: t.inputSchema, + } +} + +func (t *mcpTool) Execute(ctx context.Context, arguments string) (ToolResult, error) { + var args map[string]any + if arguments != "" { + if err := json.Unmarshal([]byte(arguments), &args); err != nil { + return ToolResult{ + Content: fmt.Sprintf("Invalid arguments: %s", err.Error()), + IsError: true, + }, nil + } + } + + result, err := t.server.session.CallTool( + ctx, + &mcp.CallToolParams{ + Name: t.name, + Arguments: args, + }, + ) + if err != nil { + return ToolResult{}, fmt.Errorf("cannot call MCP tool %q: %w", t.name, err) + } + + content := extractMCPContent(result) + + return ToolResult{ + Content: content, + IsError: result.IsError, + }, nil +} + +func extractMCPContent(result *mcp.CallToolResult) string { + if result == nil || len(result.Content) == 0 { + return "" + } + + var parts []string + for _, c := range result.Content { + if tc, ok := c.(*mcp.TextContent); ok { + parts = append(parts, tc.Text) + } + } + + return strings.Join(parts, "\n") +} diff --git a/pkg/agent/mcp_test.go b/pkg/agent/mcp_test.go new file mode 100644 index 000000000..f270209b0 --- /dev/null +++ b/pkg/agent/mcp_test.go @@ -0,0 +1,369 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "encoding/json" + "sync" + "testing" + + "github.com/modelcontextprotocol/go-sdk/mcp" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/llm" +) + +func TestNewMCPServer(t *testing.T) { + t.Parallel() + + s := NewMCPServer("my-server", nil) + + require.NotNil(t, s) + assert.Equal(t, "my-server", s.name) + assert.Nil(t, s.session) + assert.False(t, s.toolsCached) + assert.Nil(t, s.cachedTools) +} + +func TestMCPServer_Name(t *testing.T) { + t.Parallel() + + s := &MCPServer{name: "test-server"} + assert.Equal(t, "test-server", s.Name()) +} + +func TestMCPServer_Tools(t *testing.T) { + t.Parallel() + + t.Run( + "returns cached tools without calling session", + func(t *testing.T) { + t.Parallel() + + tool := &mcpTool{ + name: "cached-tool", + description: "A cached tool", + inputSchema: json.RawMessage(`{"type":"object"}`), + } + + s := &MCPServer{name: "s"} + s.cachedTools = []Tool{tool} + s.toolsCached = true + + tools, err := s.Tools(context.Background()) + require.NoError(t, err) + require.Len(t, tools, 1) + assert.Equal(t, "cached-tool", tools[0].Name()) + }, + ) + + t.Run( + "returns a defensive copy of cached tools", + func(t *testing.T) { + t.Parallel() + + tool := &mcpTool{ + name: "tool-a", + description: "Tool A", + inputSchema: json.RawMessage(`{}`), + } + + s := &MCPServer{name: "s"} + s.cachedTools = []Tool{tool} + s.toolsCached = true + + tools1, err := s.Tools(context.Background()) + require.NoError(t, err) + + tools2, err := s.Tools(context.Background()) + require.NoError(t, err) + + // Mutating one slice must not affect the other. + tools1[0] = nil + assert.NotNil(t, tools2[0]) + + // Underlying cache must be untouched. + s.mu.RLock() + assert.NotNil(t, s.cachedTools[0]) + s.mu.RUnlock() + }, + ) + + t.Run( + "concurrent reads return consistent results", + func(t *testing.T) { + t.Parallel() + + s := &MCPServer{name: "s"} + s.cachedTools = []Tool{ + &mcpTool{ + name: "echo", + description: "Echoes the input", + inputSchema: json.RawMessage(`{"type":"object"}`), + }, + } + s.toolsCached = true + + const goroutines = 50 + var wg sync.WaitGroup + wg.Add(goroutines) + + for range goroutines { + go func() { + defer wg.Done() + + tools, err := s.Tools(context.Background()) + require.NoError(t, err) + assert.Len(t, tools, 1) + assert.Equal(t, "echo", tools[0].Name()) + }() + } + + wg.Wait() + }, + ) +} + +func TestMCPServer_ResetCache(t *testing.T) { + t.Parallel() + + t.Run( + "clears cached tools", + func(t *testing.T) { + t.Parallel() + + s := &MCPServer{name: "s"} + s.cachedTools = []Tool{ + &mcpTool{ + name: "tool1", + description: "A tool", + inputSchema: json.RawMessage(`{"type":"object"}`), + }, + } + s.toolsCached = true + + tools, err := s.Tools(context.Background()) + require.NoError(t, err) + require.Len(t, tools, 1) + + s.ResetCache() + + s.mu.RLock() + assert.False(t, s.toolsCached) + assert.Nil(t, s.cachedTools) + s.mu.RUnlock() + }, + ) + + t.Run( + "is safe to call on fresh server", + func(t *testing.T) { + t.Parallel() + + s := &MCPServer{name: "s"} + s.ResetCache() + + s.mu.RLock() + assert.False(t, s.toolsCached) + assert.Nil(t, s.cachedTools) + s.mu.RUnlock() + }, + ) + + t.Run( + "concurrent resets do not race", + func(t *testing.T) { + t.Parallel() + + s := &MCPServer{name: "s"} + s.cachedTools = []Tool{ + &mcpTool{ + name: "tool1", + description: "A tool", + inputSchema: json.RawMessage(`{"type":"object"}`), + }, + } + s.toolsCached = true + + const goroutines = 50 + var wg sync.WaitGroup + wg.Add(goroutines) + + for range goroutines { + go func() { + defer wg.Done() + s.ResetCache() + }() + } + + wg.Wait() + + s.mu.RLock() + assert.False(t, s.toolsCached) + assert.Nil(t, s.cachedTools) + s.mu.RUnlock() + }, + ) +} + +func TestMCPTool_Name(t *testing.T) { + t.Parallel() + + tool := &mcpTool{name: "get_weather"} + assert.Equal(t, "get_weather", tool.Name()) +} + +func TestMCPTool_Definition(t *testing.T) { + t.Parallel() + + schema := json.RawMessage(`{"type":"object","properties":{"city":{"type":"string"}}}`) + + tool := &mcpTool{ + name: "get_weather", + description: "Returns current weather for a city", + inputSchema: schema, + } + + def := tool.Definition() + + assert.Equal(t, llm.Tool{ + Name: "get_weather", + Description: "Returns current weather for a city", + Parameters: schema, + }, def) +} + +func TestMCPTool_Execute(t *testing.T) { + t.Parallel() + + t.Run( + "invalid JSON arguments returns error result", + func(t *testing.T) { + t.Parallel() + + tool := &mcpTool{name: "my_tool"} + + result, err := tool.Execute(context.Background(), "not-json") + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Invalid arguments") + }, + ) + + t.Run( + "malformed JSON arguments returns error result", + func(t *testing.T) { + t.Parallel() + + tool := &mcpTool{name: "my_tool"} + + result, err := tool.Execute(context.Background(), `{"key": }`) + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Invalid arguments") + }, + ) + + t.Run( + "JSON array arguments returns error result", + func(t *testing.T) { + t.Parallel() + + tool := &mcpTool{name: "my_tool"} + + // Execute expects a JSON object, not an array. + result, err := tool.Execute(context.Background(), `[1,2,3]`) + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Invalid arguments") + }, + ) +} + +func TestExtractMCPContent(t *testing.T) { + t.Parallel() + + t.Run( + "nil result returns empty", + func(t *testing.T) { + t.Parallel() + assert.Equal(t, "", extractMCPContent(nil)) + }, + ) + + t.Run( + "empty content returns empty", + func(t *testing.T) { + t.Parallel() + result := &mcp.CallToolResult{Content: []mcp.Content{}} + assert.Equal(t, "", extractMCPContent(result)) + }, + ) + + t.Run( + "single text content", + func(t *testing.T) { + t.Parallel() + result := &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: "hello world"}, + }, + } + assert.Equal(t, "hello world", extractMCPContent(result)) + }, + ) + + t.Run( + "multiple text contents joined by newline", + func(t *testing.T) { + t.Parallel() + result := &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: "line one"}, + &mcp.TextContent{Text: "line two"}, + }, + } + assert.Equal(t, "line one\nline two", extractMCPContent(result)) + }, + ) + + t.Run( + "non-text content is skipped", + func(t *testing.T) { + t.Parallel() + result := &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.TextContent{Text: "text part"}, + &mcp.ImageContent{Data: []byte("base64data"), MIMEType: "image/png"}, + }, + } + assert.Equal(t, "text part", extractMCPContent(result)) + }, + ) + + t.Run( + "only non-text content returns empty", + func(t *testing.T) { + t.Parallel() + result := &mcp.CallToolResult{ + Content: []mcp.Content{ + &mcp.ImageContent{Data: []byte("base64data"), MIMEType: "image/png"}, + }, + } + assert.Equal(t, "", extractMCPContent(result)) + }, + ) +} diff --git a/pkg/agent/model_settings.go b/pkg/agent/model_settings.go new file mode 100644 index 000000000..3723f8752 --- /dev/null +++ b/pkg/agent/model_settings.go @@ -0,0 +1,28 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import "go.probo.inc/probo/pkg/llm" + +type ModelSettings struct { + Temperature *float64 + TopP *float64 + FrequencyPenalty *float64 + PresencePenalty *float64 + MaxTokens *int + ToolChoice *llm.ToolChoice + ParallelToolCalls *bool +} + diff --git a/pkg/agent/output_type.go b/pkg/agent/output_type.go new file mode 100644 index 000000000..160495682 --- /dev/null +++ b/pkg/agent/output_type.go @@ -0,0 +1,46 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "encoding/json" + + "go.probo.inc/probo/pkg/llm" +) + +// OutputType describes a structured output schema that the agent should +// produce. Build one with NewOutputType[T](). +type OutputType struct { + Name string + Schema json.RawMessage +} + +func NewOutputType[T any](name string) *OutputType { + return &OutputType{ + Name: name, + Schema: jsonSchemaFor[T](), + } +} + +func (o *OutputType) responseFormat() *llm.ResponseFormat { + return &llm.ResponseFormat{ + Type: llm.ResponseFormatJSONSchema, + JSONSchema: &llm.JSONSchema{ + Name: o.Name, + Schema: o.Schema, + Strict: true, + }, + } +} diff --git a/pkg/agent/output_type_test.go b/pkg/agent/output_type_test.go new file mode 100644 index 000000000..d411e96ac --- /dev/null +++ b/pkg/agent/output_type_test.go @@ -0,0 +1,107 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/llm" +) + +func TestNewOutputType_SetsNameAndSchema(t *testing.T) { + t.Parallel() + + type Result struct { + Answer string `json:"answer"` + Score int `json:"score"` + } + + ot := NewOutputType[Result]("test_result") + + assert.Equal(t, "test_result", ot.Name) + require.NotNil(t, ot.Schema) + + var schema map[string]any + require.NoError(t, json.Unmarshal(ot.Schema, &schema)) + assert.Equal(t, "object", schema["type"]) + + props := schema["properties"].(map[string]any) + assert.Contains(t, props, "answer") + assert.Contains(t, props, "score") +} + +func TestNewOutputType_EmptyStruct(t *testing.T) { + t.Parallel() + + type Empty struct{} + + ot := NewOutputType[Empty]("empty") + + assert.Equal(t, "empty", ot.Name) + + var schema map[string]any + require.NoError(t, json.Unmarshal(ot.Schema, &schema)) + assert.Equal(t, "object", schema["type"]) + + props, ok := schema["properties"].(map[string]any) + require.True(t, ok) + assert.Empty(t, props) +} + +func TestOutputType_responseFormat(t *testing.T) { + t.Parallel() + + type Verdict struct { + Approved bool `json:"approved"` + Reason string `json:"reason"` + } + + ot := NewOutputType[Verdict]("verdict") + rf := ot.responseFormat() + + require.NotNil(t, rf) + assert.Equal(t, llm.ResponseFormatJSONSchema, rf.Type) + require.NotNil(t, rf.JSONSchema) + assert.Equal(t, "verdict", rf.JSONSchema.Name) + assert.True(t, rf.JSONSchema.Strict) + assert.JSONEq(t, string(ot.Schema), string(rf.JSONSchema.Schema)) +} + +func TestOutputType_responseFormat_SchemaMatchesOutputType(t *testing.T) { + t.Parallel() + + type Analysis struct { + Summary string `json:"summary"` + Tags []string `json:"tags"` + Priority *int `json:"priority,omitempty"` + } + + ot := NewOutputType[Analysis]("analysis") + rf := ot.responseFormat() + + var schema map[string]any + require.NoError(t, json.Unmarshal(rf.JSONSchema.Schema, &schema)) + + props := schema["properties"].(map[string]any) + assert.Contains(t, props, "summary") + assert.Contains(t, props, "tags") + assert.Contains(t, props, "priority") + + tagsProp := props["tags"].(map[string]any) + assert.Equal(t, "array", tagsProp["type"]) +} diff --git a/pkg/agent/result.go b/pkg/agent/result.go new file mode 100644 index 000000000..b7acf0d3c --- /dev/null +++ b/pkg/agent/result.go @@ -0,0 +1,33 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import "go.probo.inc/probo/pkg/llm" + +type ( + Result struct { + Messages []llm.Message + Usage llm.Usage + Turns int + LastAgent *Agent + } +) + +func (r *Result) FinalMessage() llm.Message { + if len(r.Messages) == 0 { + return llm.Message{} + } + return r.Messages[len(r.Messages)-1] +} diff --git a/pkg/agent/result_test.go b/pkg/agent/result_test.go new file mode 100644 index 000000000..4f8c1987b --- /dev/null +++ b/pkg/agent/result_test.go @@ -0,0 +1,91 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +func TestResult_FinalMessage(t *testing.T) { + t.Parallel() + + t.Run( + "returns zero value when messages is empty", + func(t *testing.T) { + t.Parallel() + + r := &agent.Result{} + msg := r.FinalMessage() + + assert.Equal(t, llm.Message{}, msg) + assert.Equal(t, "", msg.Text()) + }, + ) + + t.Run( + "returns the last message with single message", + func(t *testing.T) { + t.Parallel() + + r := &agent.Result{ + Messages: []llm.Message{ + assistantMessage("Hello!"), + }, + } + + assert.Equal(t, "Hello!", r.FinalMessage().Text()) + }, + ) + + t.Run( + "returns the last message with multiple messages", + func(t *testing.T) { + t.Parallel() + + r := &agent.Result{ + Messages: []llm.Message{ + userMessage("Hi"), + assistantMessage("Hello!"), + userMessage("How are you?"), + assistantMessage("I'm fine."), + }, + } + + assert.Equal(t, "I'm fine.", r.FinalMessage().Text()) + }, + ) + + t.Run( + "returns message regardless of role", + func(t *testing.T) { + t.Parallel() + + r := &agent.Result{ + Messages: []llm.Message{ + assistantMessage("first"), + userMessage("last user message"), + }, + } + + msg := r.FinalMessage() + assert.Equal(t, llm.RoleUser, msg.Role) + assert.Equal(t, "last user message", msg.Text()) + }, + ) +} diff --git a/pkg/agent/run.go b/pkg/agent/run.go new file mode 100644 index 000000000..a0df7cbe9 --- /dev/null +++ b/pkg/agent/run.go @@ -0,0 +1,978 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "errors" + "fmt" + "sync" + + "go.gearno.de/kit/log" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/codes" + "go.opentelemetry.io/otel/trace" + "go.probo.inc/probo/pkg/llm" +) + +const tracerName = "go.probo.inc/probo/pkg/agent" + +type ( + CallLLMFunc func(ctx context.Context, agent *Agent, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) + + runOpts struct { + callLLM CallLLMFunc + onEvent func(ctx context.Context, ev StreamEvent) + skipInputGuardrails bool + skipSessionLoad bool + initialUsage llm.Usage + initialTurns int + } + + loopState struct { + agent *Agent + toolMap map[string]ToolDescriptor + toolDefs []llm.Tool + messages []llm.Message + inputMessages []llm.Message + systemPrompt string + totalUsage llm.Usage + turns int + toolUsedInRun bool + + tracer trace.Tracer + runSpan trace.Span + opts runOpts + logger *log.Logger + } + + parallelToolEntry struct { + result ToolResult + err error + } +) + +func noopEvent(_ context.Context, _ StreamEvent) {} + +func blockingCallLLM(ctx context.Context, agent *Agent, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + return agent.client.ChatCompletion(ctx, req) +} + +func (a *Agent) Run(ctx context.Context, messages []llm.Message) (*Result, error) { + return coreLoop( + ctx, + a, + messages, + runOpts{ + callLLM: blockingCallLLM, + onEvent: noopEvent, + }, + ) +} + +func (s *loopState) resolveAgentTools(ctx context.Context) error { + tools, toolMap, err := s.agent.resolveTools(ctx) + if err != nil { + return err + } + + toolDefs := make([]llm.Tool, len(tools)) + for i, t := range tools { + toolDefs[i] = t.Definition() + } + + s.toolMap = toolMap + s.toolDefs = toolDefs + return nil +} + +func (s *loopState) finishRun(ctx context.Context, result *Result, err error) (*Result, error) { + defer s.runSpan.End() + defer func() { + emitHook(s.agent, func(h RunHooks) { h.OnRunEnd(ctx, s.agent, result, err) }) + }() + + if err != nil { + s.runSpan.RecordError(err) + s.runSpan.SetStatus(codes.Error, err.Error()) + s.opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err}) + + s.logger.ErrorCtx( + ctx, + "agent run failed", + log.String("agent", s.agent.name), + log.Int("turns", s.turns), + log.Error(err), + ) + + return result, err + } + + if s.agent.session != nil { + if saveErr := s.agent.session.Save(ctx, s.agent.sessionID, result.Messages); saveErr != nil { + err = fmt.Errorf("cannot save session: %w", saveErr) + result = nil + + s.runSpan.RecordError(err) + s.runSpan.SetStatus(codes.Error, err.Error()) + s.opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err}) + + s.logger.ErrorCtx( + ctx, + "cannot save session", + log.String("agent", s.agent.name), + log.Error(err), + ) + + return result, err + } + } + + s.runSpan.SetAttributes( + attribute.Int("agent.turns", result.Turns), + attribute.Int("agent.usage.input_tokens", result.Usage.InputTokens), + attribute.Int("agent.usage.output_tokens", result.Usage.OutputTokens), + ) + + s.logger.InfoCtx( + ctx, + "agent run completed", + log.String("agent", s.agent.name), + log.Int("turns", result.Turns), + log.Int("input_tokens", result.Usage.InputTokens), + log.Int("output_tokens", result.Usage.OutputTokens), + ) + + s.opts.onEvent(ctx, StreamEvent{Type: StreamEventComplete, Agent: s.agent, Result: result}) + return result, err +} + +func (s *loopState) applyHandoff(ctx context.Context, handoffTarget *Handoff) error { + emitHook(s.agent, func(h RunHooks) { h.OnHandoff(ctx, s.agent, handoffTarget.Agent) }) + emitAgentHook(handoffTarget.Agent, func(h AgentHooks) { h.OnHandoff(ctx, handoffTarget.Agent, s.agent) }) + + s.opts.onEvent(ctx, StreamEvent{Type: StreamEventHandoff, Agent: handoffTarget.Agent}) + + s.logger.InfoCtx( + ctx, + "agent handoff", + log.String("from", s.agent.name), + log.String("to", handoffTarget.Agent.name), + ) + + if handoffTarget.InputFilter != nil { + s.messages = handoffTarget.InputFilter( + HandoffInputData{ + InputHistory: s.inputMessages, + NewItems: s.messages[len(s.inputMessages):], + }, + ) + } + + s.agent = handoffTarget.Agent + s.logger = s.agent.logger + + if err := s.resolveAgentTools(ctx); err != nil { + return err + } + + emitAgentHook(s.agent, func(h AgentHooks) { h.OnStart(ctx, s.agent) }) + + s.opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentStart, Agent: s.agent}) + + s.systemPrompt = s.agent.buildSystemPrompt(ctx) + s.toolUsedInRun = false + + return nil +} + +func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Message, opts runOpts) (*Result, error) { + s := &loopState{ + agent: startAgent, + inputMessages: inputMessages, + systemPrompt: startAgent.buildSystemPrompt(ctx), + totalUsage: opts.initialUsage, + turns: opts.initialTurns, + tracer: otel.GetTracerProvider().Tracer(tracerName), + opts: opts, + logger: startAgent.logger, + } + + s.messages = make([]llm.Message, 0, len(inputMessages)) + + if !opts.skipSessionLoad && s.agent.session != nil { + prev, err := s.agent.session.Load(ctx, s.agent.sessionID) + if err != nil { + return nil, fmt.Errorf("cannot load session: %w", err) + } + + s.logger.InfoCtx( + ctx, + "session loaded", + log.String("session_id", s.agent.sessionID), + log.Int("message_count", len(prev)), + ) + + s.messages = append(s.messages, prev...) + } + + s.messages = append(s.messages, inputMessages...) + + if !opts.skipInputGuardrails { + if err := runInputGuardrails(ctx, s.agent, s.messages); err != nil { + opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err}) + + s.logger.ErrorCtx( + ctx, + "input guardrail tripped", + log.Error(err), + ) + + return nil, err + } + } + + emitHook(s.agent, func(h RunHooks) { h.OnRunStart(ctx, s.agent, s.messages) }) + emitAgentHook(s.agent, func(h AgentHooks) { h.OnStart(ctx, s.agent) }) + + opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentStart, Agent: s.agent}) + + ctx, s.runSpan = s.tracer.Start( + ctx, + fmt.Sprintf("agent.run %s", s.agent.name), + trace.WithSpanKind(trace.SpanKindInternal), + trace.WithAttributes( + attribute.String("agent.name", s.agent.name), + attribute.String("agent.model", s.agent.model), + ), + ) + + if err := s.resolveAgentTools(ctx); err != nil { + s.runSpan.End() + return nil, err + } + + s.logger.InfoCtx( + ctx, + "agent run started", + log.String("model", s.agent.model), + log.Int("max_turns", s.agent.maxTurns), + log.Int("tool_count", len(s.toolDefs)), + ) + + for { + if err := ctx.Err(); err != nil { + return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: %w", err)) + } + + if s.turns >= s.agent.maxTurns { + return s.finishRun(ctx, nil, &MaxTurnsExceededError{MaxTurns: s.agent.maxTurns}) + } + + fullMessages := buildFullMessages(s.systemPrompt, s.messages) + + responseFormat := s.agent.responseFormat + if responseFormat == nil && s.agent.outputType != nil { + responseFormat = s.agent.outputType.responseFormat() + } + + toolChoice := s.agent.modelSettings.ToolChoice + if s.toolUsedInRun && s.agent.resetToolChoice && toolChoice != nil { + toolChoice = nil + } + + req := &llm.ChatCompletionRequest{ + Model: s.agent.model, + Messages: fullMessages, + Tools: s.toolDefs, + Temperature: s.agent.modelSettings.Temperature, + TopP: s.agent.modelSettings.TopP, + FrequencyPenalty: s.agent.modelSettings.FrequencyPenalty, + PresencePenalty: s.agent.modelSettings.PresencePenalty, + MaxTokens: s.agent.modelSettings.MaxTokens, + ToolChoice: toolChoice, + ParallelToolCalls: s.agent.modelSettings.ParallelToolCalls, + ResponseFormat: responseFormat, + } + + s.logger.InfoCtx( + ctx, + "calling LLM", + log.Int("turn", s.turns+1), + log.Int("message_count", len(fullMessages)), + ) + + resp, err := callLLMWithHooks(ctx, s.agent, req, opts) + if err != nil { + return s.finishRun(ctx, nil, fmt.Errorf("cannot call LLM: %w", err)) + } + + s.totalUsage = s.totalUsage.Add(resp.Usage) + s.turns++ + + s.logger.InfoCtx( + ctx, + "LLM response received", + log.Int("turn", s.turns), + log.String("finish_reason", string(resp.FinishReason)), + log.Int("input_tokens", resp.Usage.InputTokens), + log.Int("output_tokens", resp.Usage.OutputTokens), + ) + + s.messages = append(s.messages, resp.Message) + + switch resp.FinishReason { + case llm.FinishReasonStop, llm.FinishReasonLength: + if err := runOutputGuardrails(ctx, s.agent, resp.Message); err != nil { + return s.finishRun(ctx, nil, err) + } + + result := &Result{ + Messages: s.messages, + Usage: s.totalUsage, + Turns: s.turns, + LastAgent: s.agent, + } + + emitAgentHook(s.agent, func(h AgentHooks) { h.OnEnd(ctx, s.agent, resp.Message.Text()) }) + + opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentEnd, Agent: s.agent}) + return s.finishRun(ctx, result, nil) + + case llm.FinishReasonToolCalls: + s.toolUsedInRun = true + + s.logger.InfoCtx( + ctx, + "dispatching tool calls", + log.Int("count", len(resp.Message.ToolCalls)), + ) + + handoffTarget, toolResults, toolMsgs, err := dispatchToolCalls( + ctx, + s.tracer, + s.agent, + resp.Message.ToolCalls, + s.toolMap, + opts.onEvent, + s.logger, + ) + s.messages = append(s.messages, toolMsgs...) + + if err != nil { + if nae, ok := errors.AsType[*needsApprovalError](err); ok { + s.logger.InfoCtx( + ctx, + "run interrupted, approval required", + log.Int("pending_count", len(nae.pendingApprovals)), + ) + + msgsCopy := make([]llm.Message, len(s.messages)) + copy(msgsCopy, s.messages) + + return s.finishRun( + ctx, + nil, + &InterruptedError{ + ToolCalls: nae.allToolCalls, + PendingApprovals: nae.pendingApprovals, + Agent: s.agent, + Messages: msgsCopy, + Usage: s.totalUsage, + Turns: s.turns, + }, + ) + } + return s.finishRun(ctx, nil, err) + } + + finalOutput, isFinal, behaviorErr := s.agent.toolUseBehavior(ctx, toolResults) + if behaviorErr != nil { + return s.finishRun(ctx, nil, fmt.Errorf("cannot evaluate tool use behavior: %w", behaviorErr)) + } + + if isFinal { + s.messages = append( + s.messages, llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: finalOutput}}, + }, + ) + + result := &Result{ + Messages: s.messages, + Usage: s.totalUsage, + Turns: s.turns, + LastAgent: s.agent, + } + + emitAgentHook(s.agent, func(h AgentHooks) { h.OnEnd(ctx, s.agent, finalOutput) }) + + opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentEnd, Agent: s.agent}) + return s.finishRun(ctx, result, nil) + } + + if handoffTarget != nil { + if handoffErr := s.applyHandoff(ctx, handoffTarget); handoffErr != nil { + return s.finishRun(ctx, nil, handoffErr) + } + } + + case llm.FinishReasonContentFilter: + return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: content was filtered by the provider")) + + default: + return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: unexpected finish reason %q", resp.FinishReason)) + } + } +} + +func buildFullMessages(systemPrompt string, messages []llm.Message) []llm.Message { + fullMessages := make([]llm.Message, 0, len(messages)+1) + if systemPrompt != "" { + fullMessages = append( + fullMessages, + llm.Message{ + Role: llm.RoleSystem, + Parts: []llm.Part{llm.TextPart{Text: systemPrompt}}, + }, + ) + } + + return append(fullMessages, messages...) +} + +func callLLMWithHooks( + ctx context.Context, + agent *Agent, + req *llm.ChatCompletionRequest, + opts runOpts, +) (*llm.ChatCompletionResponse, error) { + emitHook(agent, func(h RunHooks) { h.OnLLMStart(ctx, agent, req.Messages) }) + emitAgentHook(agent, func(h AgentHooks) { h.OnLLMStart(ctx, agent, req.Messages) }) + + resp, err := opts.callLLM(ctx, agent, req) + if err != nil { + emitHook(agent, func(h RunHooks) { h.OnLLMEnd(ctx, agent, nil, err) }) + emitAgentHook(agent, func(h AgentHooks) { h.OnLLMEnd(ctx, agent, nil, err) }) + return nil, err + } + + emitHook(agent, func(h RunHooks) { h.OnLLMEnd(ctx, agent, resp, nil) }) + emitAgentHook(agent, func(h AgentHooks) { h.OnLLMEnd(ctx, agent, resp, nil) }) + + return resp, nil +} + +func dispatchToolCalls( + ctx context.Context, + tracer trace.Tracer, + agent *Agent, + toolCalls []llm.ToolCall, + toolMap map[string]ToolDescriptor, + onEvent func(context.Context, StreamEvent), + logger *log.Logger, +) (*Handoff, []ToolCallResult, []llm.Message, error) { + if err := checkApproval(ctx, agent, toolCalls); err != nil { + return nil, nil, nil, err + } + + return executeToolCalls(ctx, tracer, agent, toolCalls, toolMap, onEvent, logger) +} + +func executeToolCalls( + ctx context.Context, + tracer trace.Tracer, + agent *Agent, + toolCalls []llm.ToolCall, + toolMap map[string]ToolDescriptor, + onEvent func(context.Context, StreamEvent), + logger *log.Logger, +) (*Handoff, []ToolCallResult, []llm.Message, error) { + descriptors := make([]ToolDescriptor, len(toolCalls)) + handoffIdx := -1 + + for i, tc := range toolCalls { + desc, ok := toolMap[tc.Function.Name] + if !ok { + return nil, nil, nil, fmt.Errorf("cannot dispatch tool call: unknown tool %q", tc.Function.Name) + } + descriptors[i] = desc + if _, isHandoff := desc.(*handoffToolAdapter); isHandoff && handoffIdx == -1 { + handoffIdx = i + } + } + + if handoffIdx >= 0 { + return executeWithHandoff(ctx, tracer, agent, toolCalls, descriptors, handoffIdx, onEvent, logger) + } + + tools := make([]Tool, len(descriptors)) + for i, d := range descriptors { + tools[i] = d.(Tool) + } + + results, msgs, err := executeParallel(ctx, tracer, agent, toolCalls, tools, onEvent, logger) + return nil, results, msgs, err +} + +func executeWithHandoff( + ctx context.Context, + tracer trace.Tracer, + agent *Agent, + toolCalls []llm.ToolCall, + descriptors []ToolDescriptor, + handoffIdx int, + onEvent func(context.Context, StreamEvent), + logger *log.Logger, +) (*Handoff, []ToolCallResult, []llm.Message, error) { + var ( + results []ToolCallResult + msgs []llm.Message + ) + + for i := range handoffIdx { + logger.InfoCtx( + ctx, + "executing tool before handoff", + log.String("tool", toolCalls[i].Function.Name), + ) + + tr, err := executeSingleTool(ctx, tracer, agent, toolCalls[i], descriptors[i].(Tool), onEvent, logger) + if err != nil { + return nil, nil, msgs, err + } + + msgs = append( + msgs, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: toolCalls[i].ID, + Parts: []llm.Part{llm.TextPart{Text: tr.Content}}, + }, + ) + + results = append( + results, + ToolCallResult{ + ToolName: toolCalls[i].Function.Name, + Arguments: toolCalls[i].Function.Arguments, + Result: tr, + }, + ) + } + + ht := descriptors[handoffIdx].(*handoffToolAdapter) + if ht.handoff.OnHandoff != nil { + if err := ht.handoff.OnHandoff(ctx); err != nil { + return nil, nil, msgs, fmt.Errorf("cannot execute handoff callback for %q: %w", ht.handoff.Agent.name, err) + } + } + + msgs = append( + msgs, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: toolCalls[handoffIdx].ID, + Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf("Transferred to %s", ht.handoff.Agent.name)}}, + }, + ) + + return ht.handoff, results, msgs, nil +} + +func executeParallel( + ctx context.Context, + tracer trace.Tracer, + agent *Agent, + toolCalls []llm.ToolCall, + tools []Tool, + onEvent func(context.Context, StreamEvent), + logger *log.Logger, +) ([]ToolCallResult, []llm.Message, error) { + entries := make([]parallelToolEntry, len(toolCalls)) + var wg sync.WaitGroup + + wg.Add(len(toolCalls)) + for i := range toolCalls { + go func(idx int, tc llm.ToolCall, tool Tool) { + defer wg.Done() + + tr, err := executeSingleTool(ctx, tracer, agent, tc, tool, onEvent, logger) + if err != nil { + entries[idx] = parallelToolEntry{err: err} + return + } + entries[idx] = parallelToolEntry{result: tr} + }(i, toolCalls[i], tools[i]) + } + + wg.Wait() + + var ( + results []ToolCallResult + msgs []llm.Message + ) + + for i, tc := range toolCalls { + entry := entries[i] + + if entry.err != nil { + msgs = append( + msgs, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: tc.ID, + Parts: []llm.Part{ + llm.TextPart{ + Text: fmt.Sprintf("Error: %s", entry.err.Error()), + }, + }, + }, + ) + results = append( + results, + ToolCallResult{ + ToolName: tc.Function.Name, + Arguments: tc.Function.Arguments, + Result: ToolResult{ + Content: fmt.Sprintf("Error: %s", entry.err.Error()), + IsError: true, + }, + }, + ) + continue + } + + msgs = append( + msgs, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: tc.ID, + Parts: []llm.Part{llm.TextPart{Text: entry.result.Content}}, + }, + ) + + results = append( + results, + ToolCallResult{ + ToolName: tc.Function.Name, + Arguments: tc.Function.Arguments, + Result: entry.result, + }, + ) + } + + return results, msgs, nil +} + +func executeSingleTool( + ctx context.Context, + tracer trace.Tracer, + agent *Agent, + tc llm.ToolCall, + tool Tool, + onEvent func(context.Context, StreamEvent), + logger *log.Logger, +) (ToolResult, error) { + onEvent(ctx, StreamEvent{Type: StreamEventToolStart, Agent: agent, Tool: tool}) + + emitHook(agent, func(h RunHooks) { h.OnToolStart(ctx, agent, tool, tc.Function.Arguments) }) + emitAgentHook(agent, func(h AgentHooks) { h.OnToolStart(ctx, agent, tool) }) + + toolCtx, toolSpan := tracer.Start( + ctx, + fmt.Sprintf("agent.tool %s", tool.Name()), + trace.WithSpanKind(trace.SpanKindInternal), + trace.WithAttributes( + attribute.String("tool.name", tool.Name()), + ), + ) + + logger.InfoCtx( + ctx, + "executing tool", + log.String("tool", tool.Name()), + ) + + result, err := tool.Execute(toolCtx, tc.Function.Arguments) + if err != nil { + toolSpan.RecordError(err) + toolSpan.SetStatus(codes.Error, err.Error()) + toolSpan.End() + + emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, result, err) }) + + logger.ErrorCtx( + ctx, + "tool execution failed", + log.String("tool", tool.Name()), + log.Error(err), + ) + + return ToolResult{}, fmt.Errorf("cannot execute tool %q: %w", tool.Name(), err) + } + + toolSpan.SetAttributes(attribute.Bool("tool.is_error", result.IsError)) + toolSpan.End() + + onEvent(ctx, StreamEvent{Type: StreamEventToolEnd, Agent: agent, Tool: tool, ToolResult: &result}) + + emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, result, nil) }) + emitAgentHook(agent, func(h AgentHooks) { h.OnToolEnd(ctx, agent, tool, result) }) + + logger.InfoCtx( + ctx, + "tool execution completed", + log.String("tool", tool.Name()), + log.Bool("is_error", result.IsError), + ) + + return result, nil +} + +func checkApproval(ctx context.Context, a *Agent, toolCalls []llm.ToolCall) error { + if a.approval == nil { + return nil + } + + var pending []llm.ToolCall + for _, tc := range toolCalls { + if a.approval.requiresApproval(ctx, tc) { + pending = append(pending, tc) + } + } + + if len(pending) > 0 { + return &needsApprovalError{ + allToolCalls: toolCalls, + pendingApprovals: pending, + } + } + + return nil +} + +func runInputGuardrails(ctx context.Context, agent *Agent, messages []llm.Message) error { + for _, g := range agent.inputGuardrails { + result, err := g.Check(ctx, messages) + if err != nil { + return fmt.Errorf("cannot run input guardrail %q: %w", g.Name(), err) + } + + if result != nil && result.Tripwire { + emitHook(agent, func(h RunHooks) { h.OnGuardrailTripped(ctx, agent, g.Name(), result) }) + + return &InputGuardrailTrippedError{ + Guardrail: g.Name(), + Message: result.Message, + } + } + } + + return nil +} + +func runOutputGuardrails(ctx context.Context, agent *Agent, message llm.Message) error { + for _, g := range agent.outputGuardrails { + result, err := g.Check(ctx, message) + if err != nil { + return fmt.Errorf("cannot run output guardrail %q: %w", g.Name(), err) + } + + if result != nil && result.Tripwire { + emitHook(agent, func(h RunHooks) { h.OnGuardrailTripped(ctx, agent, g.Name(), result) }) + + return &OutputGuardrailTrippedError{ + Guardrail: g.Name(), + Message: result.Message, + } + } + } + + return nil +} + +// Resume continues an interrupted run after human approval decisions have +// been collected. It executes or denies each pending tool call according to +// the provided ResumeInput, then re-enters the agent loop. Input guardrails +// are not re-evaluated because the messages were already validated in the +// original Run call. +func Resume(ctx context.Context, interrupted *InterruptedError, input ResumeInput) (*Result, error) { + tracer := otel.GetTracerProvider().Tracer(tracerName) + logger := interrupted.Agent.logger + + agent := interrupted.Agent + messages := interrupted.Messages + + logger.InfoCtx( + ctx, + "resuming interrupted run", + log.String("agent", agent.name), + log.Int("pending_approvals", len(interrupted.PendingApprovals)), + log.Int("total_tool_calls", len(interrupted.ToolCalls)), + ) + + _, toolMap, err := agent.resolveTools(ctx) + if err != nil { + return nil, fmt.Errorf("cannot resolve tools for resume: %w", err) + } + + pendingSet := make(map[string]struct{}, len(interrupted.PendingApprovals)) + for _, tc := range interrupted.PendingApprovals { + pendingSet[tc.ID] = struct{}{} + } + + var handoffTarget *Handoff + + for _, tc := range interrupted.ToolCalls { + if _, needsApproval := pendingSet[tc.ID]; needsApproval { + approval, ok := input.Approvals[tc.ID] + if !ok || !approval.Approved { + reason := "Tool call was denied by human review." + if ok && approval.Message != "" { + reason = approval.Message + } + + logger.InfoCtx( + ctx, + "tool call denied", + log.String("tool", tc.Function.Name), + log.String("tool_call_id", tc.ID), + log.String("reason", reason), + ) + + messages = append( + messages, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: tc.ID, + Parts: []llm.Part{ + llm.TextPart{Text: reason}, + }, + }, + ) + + continue + } + + logger.InfoCtx( + ctx, + "tool call approved", + log.String("tool", tc.Function.Name), + log.String("tool_call_id", tc.ID), + ) + } + + desc, toolOK := toolMap[tc.Function.Name] + if !toolOK { + return nil, fmt.Errorf("unknown tool %q", tc.Function.Name) + } + + if ht, ok := desc.(*handoffToolAdapter); ok { + if ht.handoff.OnHandoff != nil { + if err := ht.handoff.OnHandoff(ctx); err != nil { + return nil, fmt.Errorf("cannot execute handoff callback for %q: %w", ht.handoff.Agent.name, err) + } + } + + messages = append( + messages, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: tc.ID, + Parts: []llm.Part{ + llm.TextPart{ + Text: fmt.Sprintf("Transferred to %s", ht.handoff.Agent.name), + }, + }, + }, + ) + + handoffTarget = ht.handoff + break + } + + tr, execErr := executeSingleTool(ctx, tracer, agent, tc, desc.(Tool), noopEvent, logger) + if execErr != nil { + return nil, execErr + } + + messages = append( + messages, + llm.Message{ + Role: llm.RoleTool, + ToolCallID: tc.ID, + Parts: []llm.Part{ + llm.TextPart{ + Text: tr.Content, + }, + }, + }, + ) + } + + resumeAgent := agent + if handoffTarget != nil { + logger.InfoCtx( + ctx, + "resume handoff", + log.String("from", agent.name), + log.String("to", handoffTarget.Agent.name), + ) + + emitHook(agent, func(h RunHooks) { h.OnHandoff(ctx, agent, handoffTarget.Agent) }) + emitAgentHook(handoffTarget.Agent, func(h AgentHooks) { h.OnHandoff(ctx, handoffTarget.Agent, agent) }) + + if handoffTarget.InputFilter != nil { + filtered := handoffTarget.InputFilter( + HandoffInputData{ + InputHistory: interrupted.Messages, + NewItems: messages[len(interrupted.Messages):], + }, + ) + messages = filtered + } + + resumeAgent = handoffTarget.Agent + } + + return coreLoop( + ctx, + resumeAgent, + messages, + runOpts{ + callLLM: blockingCallLLM, + onEvent: noopEvent, + skipInputGuardrails: true, + skipSessionLoad: true, + initialUsage: interrupted.Usage, + initialTurns: interrupted.Turns, + }, + ) +} + +func emitHook(agent *Agent, fn func(RunHooks)) { + for _, h := range agent.hooks { + fn(h) + } +} + +func emitAgentHook(agent *Agent, fn func(AgentHooks)) { + if agent.agentHooks != nil { + fn(agent.agentHooks) + } +} diff --git a/pkg/agent/schema.go b/pkg/agent/schema.go new file mode 100644 index 000000000..93eabf43c --- /dev/null +++ b/pkg/agent/schema.go @@ -0,0 +1,84 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "encoding/json" + "fmt" + "reflect" + + "github.com/google/jsonschema-go/jsonschema" +) + +func jsonSchemaFor[T any]() json.RawMessage { + t := reflect.TypeFor[T]() + + schema, err := jsonschema.ForType(t, nil) + if err != nil { + panic(fmt.Sprintf("cannot generate schema for %s: %v", t, err)) + } + + stripNullTypes(schema) + + data, err := json.Marshal(schema) + if err != nil { + panic(fmt.Sprintf("cannot marshal schema for %s: %v", t, err)) + } + + return json.RawMessage(data) +} + +// stripNullTypes removes "null" from union types produced by pointer fields +// (e.g. ["null","string"] becomes "string") and clears integer bounds so that +// LLM providers receive a clean schema without Go-specific type constraints. +func stripNullTypes(s *jsonschema.Schema) { + if s == nil { + return + } + + if len(s.Types) > 0 { + filtered := make([]string, 0, len(s.Types)) + for _, t := range s.Types { + if t != "null" { + filtered = append(filtered, t) + } + } + if len(filtered) == 1 { + s.Type = filtered[0] + s.Types = nil + } else if len(filtered) > 1 { + s.Types = filtered + } + } + + s.Minimum = nil + s.Maximum = nil + + if s.Type == "object" && s.Properties == nil { + s.Properties = make(map[string]*jsonschema.Schema) + } + + for _, prop := range s.Properties { + stripNullTypes(prop) + } + + if s.Items != nil { + stripNullTypes(s.Items) + } + + if s.AdditionalProperties != nil { + stripNullTypes(s.AdditionalProperties) + } +} diff --git a/pkg/agent/schema_test.go b/pkg/agent/schema_test.go new file mode 100644 index 000000000..f0b95f832 --- /dev/null +++ b/pkg/agent/schema_test.go @@ -0,0 +1,325 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGenerateSchema_PointerFieldsStripNull(t *testing.T) { + t.Parallel() + + type Params struct { + Name *string `json:"name"` + Count *int `json:"count"` + Score *float64 `json:"score"` + Done *bool `json:"done"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + for _, field := range []struct { + name string + wantType string + }{ + {"name", "string"}, + {"count", "integer"}, + {"score", "number"}, + {"done", "boolean"}, + } { + prop := props[field.name].(map[string]any) + assert.Equal(t, field.wantType, prop["type"], "field %s", field.name) + assert.Nil(t, prop["types"], "field %s should not have union types", field.name) + } +} + +func TestGenerateSchema_IntegerBoundsStripped(t *testing.T) { + t.Parallel() + + type Params struct { + Int8 int8 `json:"int8"` + Int16 int16 `json:"int16"` + Int32 int32 `json:"int32"` + Uint8 uint8 `json:"uint8"` + Uint16 uint16 `json:"uint16"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + for _, name := range []string{"int8", "int16", "int32", "uint8", "uint16"} { + prop := props[name].(map[string]any) + assert.Equal(t, "integer", prop["type"], "field %s", name) + assert.Nil(t, prop["minimum"], "field %s should have no minimum", name) + assert.Nil(t, prop["maximum"], "field %s should have no maximum", name) + } +} + +func TestGenerateSchema_EmptyStruct(t *testing.T) { + t.Parallel() + + type Params struct{} + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + assert.Equal(t, "object", schema["type"]) + + props, ok := schema["properties"].(map[string]any) + require.True(t, ok, "empty struct should have a properties field") + assert.Empty(t, props) +} + +func TestGenerateSchema_MapField(t *testing.T) { + t.Parallel() + + type Params struct { + Metadata map[string]string `json:"metadata"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + metaProp := props["metadata"].(map[string]any) + + assert.Equal(t, "object", metaProp["type"]) + + addlProps, ok := metaProp["additionalProperties"].(map[string]any) + require.True(t, ok, "map field should produce additionalProperties") + assert.Equal(t, "string", addlProps["type"]) +} + +func TestGenerateSchema_NestedPointerStruct(t *testing.T) { + t.Parallel() + + type Inner struct { + Value *string `json:"value"` + } + type Params struct { + Inner *Inner `json:"inner"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + innerProp := props["inner"].(map[string]any) + assert.Equal(t, "object", innerProp["type"]) + assert.Nil(t, innerProp["types"], "pointer to struct should not have union types") + + innerProps := innerProp["properties"].(map[string]any) + valueProp := innerProps["value"].(map[string]any) + assert.Equal(t, "string", valueProp["type"]) + assert.Nil(t, valueProp["types"], "nested pointer field should not have union types") +} + +func TestGenerateSchema_SliceOfPointers(t *testing.T) { + t.Parallel() + + type Params struct { + Names []*string `json:"names"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + namesProp := props["names"].(map[string]any) + assert.Equal(t, "array", namesProp["type"]) + + items := namesProp["items"].(map[string]any) + assert.Equal(t, "string", items["type"]) + assert.Nil(t, items["types"], "array items from pointer should not have union types") +} + +func TestGenerateSchema_MapWithPointerValues(t *testing.T) { + t.Parallel() + + type Params struct { + Scores map[string]*int `json:"scores"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + scoresProp := props["scores"].(map[string]any) + assert.Equal(t, "object", scoresProp["type"]) + + addlProps := scoresProp["additionalProperties"].(map[string]any) + assert.Equal(t, "integer", addlProps["type"]) + assert.Nil(t, addlProps["types"], "map pointer values should not have union types") + assert.Nil(t, addlProps["minimum"]) + assert.Nil(t, addlProps["maximum"]) +} + +func TestGenerateSchema_DescriptionFromJsonschemaTag(t *testing.T) { + t.Parallel() + + type Params struct { + Query string `json:"query" jsonschema:"The search query to execute"` + Limit int `json:"limit" jsonschema:"Maximum number of results to return"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + queryProp := props["query"].(map[string]any) + assert.Equal(t, "The search query to execute", queryProp["description"]) + + limitProp := props["limit"].(map[string]any) + assert.Equal(t, "Maximum number of results to return", limitProp["description"]) +} + +func TestGenerateSchema_RequiredVsOptional(t *testing.T) { + t.Parallel() + + type Params struct { + Required string `json:"required"` + Optional *string `json:"optional,omitempty"` + OmitEmpty string `json:"omit_empty,omitempty"` + AlsoNeeded int `json:"also_needed"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + required := schema["required"].([]any) + assert.Contains(t, required, "required") + assert.Contains(t, required, "also_needed") + assert.NotContains(t, required, "optional") + assert.NotContains(t, required, "omit_empty") +} + +func TestGenerateSchema_DeeplyNestedStructure(t *testing.T) { + t.Parallel() + + type Level3 struct { + Value *int `json:"value"` + } + type Level2 struct { + Items []Level3 `json:"items"` + } + type Level1 struct { + Child *Level2 `json:"child"` + } + type Params struct { + Root Level1 `json:"root"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + rootProp := props["root"].(map[string]any) + assert.Equal(t, "object", rootProp["type"]) + + rootProps := rootProp["properties"].(map[string]any) + childProp := rootProps["child"].(map[string]any) + assert.Equal(t, "object", childProp["type"]) + assert.Nil(t, childProp["types"]) + + childProps := childProp["properties"].(map[string]any) + itemsProp := childProps["items"].(map[string]any) + assert.Equal(t, "array", itemsProp["type"]) + + itemsItems := itemsProp["items"].(map[string]any) + assert.Equal(t, "object", itemsItems["type"]) + + level3Props := itemsItems["properties"].(map[string]any) + valueProp := level3Props["value"].(map[string]any) + assert.Equal(t, "integer", valueProp["type"]) + assert.Nil(t, valueProp["types"]) + assert.Nil(t, valueProp["minimum"]) + assert.Nil(t, valueProp["maximum"]) +} + +func TestGenerateSchema_SliceOfStructs(t *testing.T) { + t.Parallel() + + type Item struct { + Name string `json:"name"` + Count *int `json:"count,omitempty"` + } + type Params struct { + Items []Item `json:"items"` + } + + raw := jsonSchemaFor[Params]() + + var schema map[string]any + require.NoError(t, json.Unmarshal(raw, &schema)) + + props := schema["properties"].(map[string]any) + + itemsProp := props["items"].(map[string]any) + assert.Equal(t, "array", itemsProp["type"]) + + items := itemsProp["items"].(map[string]any) + assert.Equal(t, "object", items["type"]) + + itemProps := items["properties"].(map[string]any) + assert.Contains(t, itemProps, "name") + assert.Contains(t, itemProps, "count") + + countProp := itemProps["count"].(map[string]any) + assert.Equal(t, "integer", countProp["type"]) + assert.Nil(t, countProp["types"]) + assert.Nil(t, countProp["minimum"]) +} + +func TestStripNullTypes_NilSchema(t *testing.T) { + t.Parallel() + + assert.NotPanics(t, func() { + stripNullTypes(nil) + }) +} diff --git a/pkg/agent/session.go b/pkg/agent/session.go new file mode 100644 index 000000000..ebf50a934 --- /dev/null +++ b/pkg/agent/session.go @@ -0,0 +1,26 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + + "go.probo.inc/probo/pkg/llm" +) + +type Session interface { + Load(ctx context.Context, sessionID string) ([]llm.Message, error) + Save(ctx context.Context, sessionID string, messages []llm.Message) error +} diff --git a/pkg/agent/session_memory.go b/pkg/agent/session_memory.go new file mode 100644 index 000000000..ffb08060a --- /dev/null +++ b/pkg/agent/session_memory.go @@ -0,0 +1,84 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "sync" + + "go.probo.inc/probo/pkg/llm" +) + +var _ Session = (*memorySession)(nil) + +type memorySession struct { + mu sync.RWMutex + sessions map[string][]llm.Message +} + +func NewMemorySession() *memorySession { + return &memorySession{ + sessions: make(map[string][]llm.Message), + } +} + +func (s *memorySession) Load(_ context.Context, sessionID string) ([]llm.Message, error) { + s.mu.RLock() + defer s.mu.RUnlock() + + msgs, ok := s.sessions[sessionID] + if !ok { + return nil, nil + } + + cp := make([]llm.Message, len(msgs)) + for i, m := range msgs { + cp[i] = copyMessage(m) + } + + return cp, nil +} + +func (s *memorySession) Save(_ context.Context, sessionID string, messages []llm.Message) error { + s.mu.Lock() + defer s.mu.Unlock() + + cp := make([]llm.Message, len(messages)) + for i, m := range messages { + cp[i] = copyMessage(m) + } + + s.sessions[sessionID] = cp + return nil +} + +func copyMessage(m llm.Message) llm.Message { + cp := llm.Message{ + Role: m.Role, + ToolCallID: m.ToolCallID, + } + + if len(m.Parts) > 0 { + cp.Parts = make([]llm.Part, len(m.Parts)) + copy(cp.Parts, m.Parts) + } + + if len(m.ToolCalls) > 0 { + cp.ToolCalls = make([]llm.ToolCall, len(m.ToolCalls)) + copy(cp.ToolCalls, m.ToolCalls) + } + + return cp +} diff --git a/pkg/agent/session_memory_test.go b/pkg/agent/session_memory_test.go new file mode 100644 index 000000000..8e2b72f87 --- /dev/null +++ b/pkg/agent/session_memory_test.go @@ -0,0 +1,275 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" + "go.probo.inc/probo/pkg/llm" +) + +func TestNewMemorySession(t *testing.T) { + t.Parallel() + + t.Run( + "returns non-nil session", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + assert.NotNil(t, s) + }, + ) +} + +func TestMemorySession_Load(t *testing.T) { + t.Parallel() + + t.Run( + "returns nil for unknown session ID", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + + msgs, err := s.Load(context.Background(), "unknown") + + require.NoError(t, err) + assert.Nil(t, msgs) + }, + ) + + t.Run( + "returns saved messages", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + messages := []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "hello"}}, + }, + { + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: "hi there"}}, + }, + } + + err := s.Save(context.Background(), "sess-1", messages) + require.NoError(t, err) + + loaded, err := s.Load(context.Background(), "sess-1") + + require.NoError(t, err) + require.Len(t, loaded, 2) + assert.Equal(t, llm.RoleUser, loaded[0].Role) + assert.Equal(t, "hello", loaded[0].Text()) + assert.Equal(t, llm.RoleAssistant, loaded[1].Role) + assert.Equal(t, "hi there", loaded[1].Text()) + }, + ) + + t.Run( + "returns a defensive copy of messages", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + messages := []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "original"}}, + }, + } + + err := s.Save(context.Background(), "sess-copy", messages) + require.NoError(t, err) + + loaded, err := s.Load(context.Background(), "sess-copy") + require.NoError(t, err) + + loaded[0].Parts = []llm.Part{llm.TextPart{Text: "mutated"}} + + reloaded, err := s.Load(context.Background(), "sess-copy") + require.NoError(t, err) + + assert.Equal(t, "original", reloaded[0].Text()) + }, + ) + + t.Run( + "different session IDs are independent", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + + err := s.Save( + context.Background(), + "sess-a", + []llm.Message{ + {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "alpha"}}}, + }, + ) + require.NoError(t, err) + + err = s.Save( + context.Background(), + "sess-b", + []llm.Message{ + {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "beta"}}}, + }, + ) + require.NoError(t, err) + + a, err := s.Load(context.Background(), "sess-a") + require.NoError(t, err) + require.Len(t, a, 1) + assert.Equal(t, "alpha", a[0].Text()) + + b, err := s.Load(context.Background(), "sess-b") + require.NoError(t, err) + require.Len(t, b, 1) + assert.Equal(t, "beta", b[0].Text()) + }, + ) +} + +func TestMemorySession_Save(t *testing.T) { + t.Parallel() + + t.Run( + "stores a defensive copy of input messages", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + messages := []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "before"}}, + }, + } + + err := s.Save(context.Background(), "sess-def", messages) + require.NoError(t, err) + + messages[0].Parts = []llm.Part{llm.TextPart{Text: "after"}} + + loaded, err := s.Load(context.Background(), "sess-def") + require.NoError(t, err) + + assert.Equal(t, "before", loaded[0].Text()) + }, + ) + + t.Run( + "overwrites previous messages for same session ID", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + + err := s.Save( + context.Background(), + "sess-ow", + []llm.Message{ + {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "first"}}}, + }, + ) + require.NoError(t, err) + + err = s.Save( + context.Background(), + "sess-ow", + []llm.Message{ + {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: "second"}}}, + {Role: llm.RoleAssistant, Parts: []llm.Part{llm.TextPart{Text: "reply"}}}, + }, + ) + require.NoError(t, err) + + loaded, err := s.Load(context.Background(), "sess-ow") + require.NoError(t, err) + require.Len(t, loaded, 2) + assert.Equal(t, "second", loaded[0].Text()) + assert.Equal(t, "reply", loaded[1].Text()) + }, + ) + + t.Run( + "preserves tool calls and tool call ID", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + messages := []llm.Message{ + { + Role: llm.RoleAssistant, + ToolCalls: []llm.ToolCall{ + { + ID: "call-1", + Function: llm.FunctionCall{ + Name: "get_weather", + Arguments: `{"city":"Paris"}`, + }, + }, + }, + }, + { + Role: llm.RoleTool, + ToolCallID: "call-1", + Parts: []llm.Part{llm.TextPart{Text: "sunny"}}, + }, + } + + err := s.Save(context.Background(), "sess-tc", messages) + require.NoError(t, err) + + loaded, err := s.Load(context.Background(), "sess-tc") + require.NoError(t, err) + require.Len(t, loaded, 2) + + require.Len(t, loaded[0].ToolCalls, 1) + assert.Equal(t, "call-1", loaded[0].ToolCalls[0].ID) + assert.Equal(t, "get_weather", loaded[0].ToolCalls[0].Function.Name) + assert.Equal(t, `{"city":"Paris"}`, loaded[0].ToolCalls[0].Function.Arguments) + + assert.Equal(t, "call-1", loaded[1].ToolCallID) + assert.Equal(t, "sunny", loaded[1].Text()) + }, + ) + + t.Run( + "handles empty message slice", + func(t *testing.T) { + t.Parallel() + + s := agent.NewMemorySession() + + err := s.Save(context.Background(), "sess-empty", []llm.Message{}) + require.NoError(t, err) + + loaded, err := s.Load(context.Background(), "sess-empty") + require.NoError(t, err) + assert.Empty(t, loaded) + }, + ) +} diff --git a/pkg/agent/stream.go b/pkg/agent/stream.go new file mode 100644 index 000000000..2899bb041 --- /dev/null +++ b/pkg/agent/stream.go @@ -0,0 +1,133 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + StreamEventType string + + StreamEvent struct { + Type StreamEventType + Agent *Agent + Delta string + Tool Tool + ToolResult *ToolResult + Result *Result + Err error + } + + StreamedRun struct { + Events <-chan StreamEvent + done chan struct{} + result *Result + err error + } +) + +const ( + StreamEventAgentStart StreamEventType = "agent_start" + StreamEventAgentEnd StreamEventType = "agent_end" + StreamEventLLMDelta StreamEventType = "llm_delta" + StreamEventToolStart StreamEventType = "tool_start" + StreamEventToolEnd StreamEventType = "tool_end" + StreamEventHandoff StreamEventType = "handoff" + StreamEventComplete StreamEventType = "complete" + StreamEventError StreamEventType = "error" +) + +func (sr *StreamedRun) Wait() (*Result, error) { + for range sr.Events { + } + <-sr.done + return sr.result, sr.err +} + +func (a *Agent) RunStreamed(ctx context.Context, messages []llm.Message) *StreamedRun { + events := make(chan StreamEvent, 64) + sr := &StreamedRun{ + Events: events, + done: make(chan struct{}), + } + + go func() { + defer close(sr.done) + defer close(events) + + result, err := coreLoop( + ctx, + a, + messages, + runOpts{ + callLLM: streamingCallLLM(events), + onEvent: func(ctx context.Context, ev StreamEvent) { + trySendEvent(ctx, events, ev) + }, + }, + ) + + sr.result = result + sr.err = err + }() + + return sr +} + +func streamingCallLLM(events chan<- StreamEvent) CallLLMFunc { + return func(ctx context.Context, agent *Agent, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + stream, err := agent.client.ChatCompletionStream(ctx, req) + if err != nil { + return nil, err + } + + acc := llm.NewStreamAccumulator(stream) + for acc.Next() { + ev := acc.Event() + if ev.Delta.Content != "" { + trySendEvent( + ctx, + events, + StreamEvent{ + Type: StreamEventLLMDelta, + Agent: agent, + Delta: ev.Delta.Content, + }, + ) + } + } + + if err := acc.Err(); err != nil { + _ = stream.Close() + return nil, err + } + + if err := stream.Close(); err != nil { + return nil, err + } + + return acc.Response(), nil + } +} + +func trySendEvent(ctx context.Context, events chan<- StreamEvent, ev StreamEvent) { + select { + case events <- ev: + case <-ctx.Done(): + } +} diff --git a/pkg/agent/system_prompt.go b/pkg/agent/system_prompt.go new file mode 100644 index 000000000..d68f92d34 --- /dev/null +++ b/pkg/agent/system_prompt.go @@ -0,0 +1,51 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "bytes" + "text/template" +) + +type ( + systemPromptData struct { + Instructions string + Handoffs []systemPromptHandoff + } + + systemPromptHandoff struct { + Name string + Description string + } +) + +var systemPromptTmpl = template.Must(template.New("system_prompt").Parse( + `{{- .Instructions -}} +{{- if .Handoffs }} + +## Handoffs +You can transfer the conversation to a more specialized agent when appropriate: +{{ range .Handoffs -}} +- {{ .Name }}{{ with .Description }}: {{ . }}{{ end }} +{{ end -}} +{{- end -}} +`)) + +func buildSystemPrompt(data systemPromptData) string { + var buf bytes.Buffer + _ = systemPromptTmpl.Execute(&buf, data) + + return buf.String() +} diff --git a/pkg/agent/system_prompt_test.go b/pkg/agent/system_prompt_test.go new file mode 100644 index 000000000..c7c327634 --- /dev/null +++ b/pkg/agent/system_prompt_test.go @@ -0,0 +1,121 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "strings" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestBuildSystemPrompt(t *testing.T) { + t.Run( + "empty data returns empty string", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{}) + assert.Equal(t, "", got) + }, + ) + + t.Run( + "instructions only", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Instructions: "You are a helpful assistant.", + }) + assert.Equal(t, "You are a helpful assistant.", got) + }, + ) + + t.Run( + "handoffs only", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Handoffs: []systemPromptHandoff{ + {Name: "billing-agent", Description: "Handles billing questions."}, + }, + }) + + assert.Contains(t, got, "## Handoffs") + assert.Contains(t, got, "- billing-agent: Handles billing questions.") + }, + ) + + t.Run( + "instructions with handoffs", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Instructions: "You are a triage agent.", + Handoffs: []systemPromptHandoff{ + {Name: "billing-agent", Description: "Handles billing."}, + {Name: "support-agent", Description: "Handles support."}, + }, + }) + + assert.True(t, strings.HasPrefix(got, "You are a triage agent.")) + assert.Contains(t, got, "## Handoffs") + assert.Contains(t, got, "- billing-agent: Handles billing.") + assert.Contains(t, got, "- support-agent: Handles support.") + }, + ) + + t.Run( + "handoff without description", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Handoffs: []systemPromptHandoff{ + {Name: "silent-agent"}, + }, + }) + + assert.Contains(t, got, "- silent-agent\n") + assert.NotContains(t, got, "- silent-agent:") + }, + ) + + t.Run( + "multiple handoffs preserve order", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Handoffs: []systemPromptHandoff{ + {Name: "alpha"}, + {Name: "beta"}, + {Name: "gamma"}, + }, + }) + + idxAlpha := strings.Index(got, "- alpha") + idxBeta := strings.Index(got, "- beta") + idxGamma := strings.Index(got, "- gamma") + + assert.Greater(t, idxBeta, idxAlpha) + assert.Greater(t, idxGamma, idxBeta) + }, + ) + + t.Run( + "no trailing newline after last handoff", + func(t *testing.T) { + got := buildSystemPrompt(systemPromptData{ + Handoffs: []systemPromptHandoff{ + {Name: "agent-a", Description: "Does A."}, + }, + }) + + assert.False(t, strings.HasSuffix(got, "\n\n"), "should not end with double newline") + }, + ) +} diff --git a/pkg/agent/tool.go b/pkg/agent/tool.go new file mode 100644 index 000000000..d1d637596 --- /dev/null +++ b/pkg/agent/tool.go @@ -0,0 +1,84 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "encoding/json" + "fmt" + + "go.probo.inc/probo/pkg/llm" +) + +type ( + ToolResult struct { + Content string + IsError bool + } + + ToolDescriptor interface { + Name() string + Definition() llm.Tool + } + + Tool interface { + ToolDescriptor + Execute(ctx context.Context, arguments string) (ToolResult, error) + } + + functionTool[P any] struct { + name string + description string + fn func(ctx context.Context, params P) (ToolResult, error) + schema json.RawMessage + } +) + +func FunctionTool[P any]( + name string, + description string, + fn func(ctx context.Context, params P) (ToolResult, error), +) Tool { + schema := jsonSchemaFor[P]() + + return &functionTool[P]{ + name: name, + description: description, + fn: fn, + schema: schema, + } +} + +func (t *functionTool[P]) Name() string { return t.name } + +func (t *functionTool[P]) Definition() llm.Tool { + return llm.Tool{ + Name: t.name, + Description: t.description, + Parameters: t.schema, + } +} + +func (t *functionTool[P]) Execute(ctx context.Context, arguments string) (ToolResult, error) { + var params P + if err := json.Unmarshal([]byte(arguments), ¶ms); err != nil { + return ToolResult{ + Content: fmt.Sprintf("Invalid parameters: %s", err.Error()), + IsError: true, + }, nil + } + + return t.fn(ctx, params) +} diff --git a/pkg/agent/tool_test.go b/pkg/agent/tool_test.go new file mode 100644 index 000000000..f6b5e582f --- /dev/null +++ b/pkg/agent/tool_test.go @@ -0,0 +1,361 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" +) + +func TestFunctionTool_Name(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "my_tool", + "does things", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + assert.Equal(t, "my_tool", tool.Name()) +} + +func TestFunctionTool_Definition(t *testing.T) { + t.Parallel() + + t.Run( + "returns name and description", + func(t *testing.T) { + t.Parallel() + + type Params struct { + Query string `json:"query"` + } + + tool := agent.FunctionTool( + "search", + "Search for items", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + def := tool.Definition() + assert.Equal(t, "search", def.Name) + assert.Equal(t, "Search for items", def.Description) + }, + ) + + t.Run( + "generates valid JSON schema from params type", + func(t *testing.T) { + t.Parallel() + + type Params struct { + Name string `json:"name" jsonschema:"The item name"` + Count int `json:"count"` + } + + tool := agent.FunctionTool( + "create", + "Create items", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + def := tool.Definition() + require.NotNil(t, def.Parameters) + + var schema map[string]any + require.NoError(t, json.Unmarshal(def.Parameters, &schema)) + + assert.Equal(t, "object", schema["type"]) + + props := schema["properties"].(map[string]any) + assert.Contains(t, props, "name") + assert.Contains(t, props, "count") + + nameProp := props["name"].(map[string]any) + assert.Equal(t, "string", nameProp["type"]) + assert.Equal(t, "The item name", nameProp["description"]) + + countProp := props["count"].(map[string]any) + assert.Equal(t, "integer", countProp["type"]) + }, + ) + + t.Run( + "empty struct produces object schema with no properties", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + var schema map[string]any + require.NoError(t, json.Unmarshal(tool.Definition().Parameters, &schema)) + assert.Equal(t, "object", schema["type"]) + }, + ) + + t.Run( + "pointer fields are not nullable in schema", + func(t *testing.T) { + t.Parallel() + + type Params struct { + Title *string `json:"title,omitempty"` + } + + tool := agent.FunctionTool( + "update", + "Update", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + var schema map[string]any + require.NoError(t, json.Unmarshal(tool.Definition().Parameters, &schema)) + + props := schema["properties"].(map[string]any) + titleProp := props["title"].(map[string]any) + assert.Equal(t, "string", titleProp["type"]) + assert.Nil(t, titleProp["types"]) + }, + ) +} + +func TestFunctionTool_Execute(t *testing.T) { + t.Parallel() + + t.Run( + "unmarshals params and calls function", + func(t *testing.T) { + t.Parallel() + + type Params struct { + X int `json:"x"` + Y int `json:"y"` + } + + tool := agent.FunctionTool( + "add", + "Add two numbers", + func(_ context.Context, p Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "42"}, nil + }, + ) + + result, err := tool.Execute(context.Background(), `{"x": 1, "y": 2}`) + require.NoError(t, err) + assert.Equal(t, "42", result.Content) + assert.False(t, result.IsError) + }, + ) + + t.Run( + "passes received params to function", + func(t *testing.T) { + t.Parallel() + + type Params struct { + City string `json:"city"` + } + + var received string + tool := agent.FunctionTool( + "weather", + "Get weather", + func(_ context.Context, p Params) (agent.ToolResult, error) { + received = p.City + return agent.ToolResult{Content: "sunny"}, nil + }, + ) + + _, err := tool.Execute(context.Background(), `{"city":"Paris"}`) + require.NoError(t, err) + assert.Equal(t, "Paris", received) + }, + ) + + t.Run( + "invalid JSON returns tool error not Go error", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "noop", + "No-op", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "ok"}, nil + }, + ) + + result, err := tool.Execute(context.Background(), `{invalid`) + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Contains(t, result.Content, "Invalid parameters") + }, + ) + + t.Run( + "infrastructure error propagated as Go error", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "fail", + "Always fails", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, errors.New("db down") + }, + ) + + _, err := tool.Execute(context.Background(), `{}`) + require.Error(t, err) + assert.Contains(t, err.Error(), "db down") + }, + ) + + t.Run( + "context is forwarded to function", + func(t *testing.T) { + t.Parallel() + + type ctxKey struct{} + type Params struct{} + + tool := agent.FunctionTool( + "ctx_check", + "Check context", + func(ctx context.Context, _ Params) (agent.ToolResult, error) { + val := ctx.Value(ctxKey{}).(string) + return agent.ToolResult{Content: val}, nil + }, + ) + + ctx := context.WithValue(context.Background(), ctxKey{}, "hello") + result, err := tool.Execute(ctx, `{}`) + require.NoError(t, err) + assert.Equal(t, "hello", result.Content) + }, + ) + + t.Run( + "extra JSON fields are ignored", + func(t *testing.T) { + t.Parallel() + + type Params struct { + Name string `json:"name"` + } + + tool := agent.FunctionTool( + "greet", + "Greet", + func(_ context.Context, p Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "hi " + p.Name}, nil + }, + ) + + result, err := tool.Execute(context.Background(), `{"name":"Alice","extra":"ignored"}`) + require.NoError(t, err) + assert.Equal(t, "hi Alice", result.Content) + assert.False(t, result.IsError) + }, + ) + + t.Run( + "empty JSON object works for empty params", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "ping", + "Ping", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "pong"}, nil + }, + ) + + result, err := tool.Execute(context.Background(), `{}`) + require.NoError(t, err) + assert.Equal(t, "pong", result.Content) + }, + ) + + t.Run( + "function can return IsError true", + func(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "validate", + "Validate input", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{Content: "validation failed", IsError: true}, nil + }, + ) + + result, err := tool.Execute(context.Background(), `{}`) + require.NoError(t, err) + assert.True(t, result.IsError) + assert.Equal(t, "validation failed", result.Content) + }, + ) +} + +func TestFunctionTool_InterfaceSatisfaction(t *testing.T) { + t.Parallel() + + type Params struct{} + + tool := agent.FunctionTool( + "test", + "test tool", + func(_ context.Context, _ Params) (agent.ToolResult, error) { + return agent.ToolResult{}, nil + }, + ) + + assert.Implements(t, (*agent.Tool)(nil), tool) + assert.Implements(t, (*agent.ToolDescriptor)(nil), tool) +} diff --git a/pkg/agent/tool_use_behavior.go b/pkg/agent/tool_use_behavior.go new file mode 100644 index 000000000..67ee2c929 --- /dev/null +++ b/pkg/agent/tool_use_behavior.go @@ -0,0 +1,64 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import "context" + +// ToolUseBehavior decides whether tool-call results constitute a final +// output. When isFinal is true the run stops and finalOutput is returned +// to the caller without sending results back to the LLM. +type ToolUseBehavior func(ctx context.Context, results []ToolCallResult) (finalOutput string, isFinal bool, err error) + +type ToolCallResult struct { + ToolName string + Arguments string + Result ToolResult +} + +// RunLLMAgain is the default behavior: tools run, then the LLM receives the +// results and gets to respond. +func RunLLMAgain() ToolUseBehavior { + return func(_ context.Context, _ []ToolCallResult) (string, bool, error) { + return "", false, nil + } +} + +// StopOnFirstTool treats the output from the first tool call as the final +// result, without sending it back to the LLM. +func StopOnFirstTool() ToolUseBehavior { + return func(_ context.Context, results []ToolCallResult) (string, bool, error) { + if len(results) == 0 { + return "", false, nil + } + return results[0].Result.Content, true, nil + } +} + +// StopAtTools stops the agent when any of the listed tool names is called. +// The matching tool's output becomes the final output. +func StopAtTools(names ...string) ToolUseBehavior { + stopSet := make(map[string]struct{}, len(names)) + for _, n := range names { + stopSet[n] = struct{}{} + } + return func(_ context.Context, results []ToolCallResult) (string, bool, error) { + for _, r := range results { + if _, ok := stopSet[r.ToolName]; ok { + return r.Result.Content, true, nil + } + } + return "", false, nil + } +} diff --git a/pkg/agent/tool_use_behavior_test.go b/pkg/agent/tool_use_behavior_test.go new file mode 100644 index 000000000..e4dadfc79 --- /dev/null +++ b/pkg/agent/tool_use_behavior_test.go @@ -0,0 +1,288 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/agent" +) + +func TestRunLLMAgain(t *testing.T) { + t.Parallel() + + t.Run( + "returns not final with no results", + func(t *testing.T) { + t.Parallel() + + behavior := agent.RunLLMAgain() + output, isFinal, err := behavior(context.Background(), nil) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) + + t.Run( + "returns not final with results present", + func(t *testing.T) { + t.Parallel() + + behavior := agent.RunLLMAgain() + results := []agent.ToolCallResult{ + { + ToolName: "search", + Arguments: `{"q":"test"}`, + Result: agent.ToolResult{Content: "found 3 items"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) +} + +func TestStopOnFirstTool(t *testing.T) { + t.Parallel() + + t.Run( + "returns not final when results are empty", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopOnFirstTool() + output, isFinal, err := behavior(context.Background(), nil) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) + + t.Run( + "returns first tool output as final", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopOnFirstTool() + results := []agent.ToolCallResult{ + { + ToolName: "lookup", + Result: agent.ToolResult{Content: "first result"}, + }, + { + ToolName: "search", + Result: agent.ToolResult{Content: "second result"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "first result", output) + }, + ) + + t.Run( + "returns single result as final", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopOnFirstTool() + results := []agent.ToolCallResult{ + { + ToolName: "compute", + Result: agent.ToolResult{Content: "42"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "42", output) + }, + ) + + t.Run( + "returns empty content when first tool has empty output", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopOnFirstTool() + results := []agent.ToolCallResult{ + { + ToolName: "noop", + Result: agent.ToolResult{Content: ""}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "", output) + }, + ) +} + +func TestStopAtTools(t *testing.T) { + t.Parallel() + + t.Run( + "returns not final when results are empty", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools("done") + output, isFinal, err := behavior(context.Background(), nil) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) + + t.Run( + "returns not final when no tool matches", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools("done", "finish") + results := []agent.ToolCallResult{ + { + ToolName: "search", + Result: agent.ToolResult{Content: "results"}, + }, + { + ToolName: "compute", + Result: agent.ToolResult{Content: "42"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) + + t.Run( + "stops on matching tool and returns its output", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools("submit") + results := []agent.ToolCallResult{ + { + ToolName: "search", + Result: agent.ToolResult{Content: "search result"}, + }, + { + ToolName: "submit", + Result: agent.ToolResult{Content: "submitted"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "submitted", output) + }, + ) + + t.Run( + "returns first matching tool when multiple match", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools("done", "finish") + results := []agent.ToolCallResult{ + { + ToolName: "finish", + Result: agent.ToolResult{Content: "finished"}, + }, + { + ToolName: "done", + Result: agent.ToolResult{Content: "done"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "finished", output) + }, + ) + + t.Run( + "matches any of the listed tool names", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools("alpha", "beta", "gamma") + results := []agent.ToolCallResult{ + { + ToolName: "unrelated", + Result: agent.ToolResult{Content: "ignored"}, + }, + { + ToolName: "gamma", + Result: agent.ToolResult{Content: "gamma output"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.True(t, isFinal) + assert.Equal(t, "gamma output", output) + }, + ) + + t.Run( + "no stop names means never final", + func(t *testing.T) { + t.Parallel() + + behavior := agent.StopAtTools() + results := []agent.ToolCallResult{ + { + ToolName: "anything", + Result: agent.ToolResult{Content: "value"}, + }, + } + + output, isFinal, err := behavior(context.Background(), results) + + require.NoError(t, err) + assert.False(t, isFinal) + assert.Equal(t, "", output) + }, + ) +} diff --git a/pkg/agent/typed.go b/pkg/agent/typed.go new file mode 100644 index 000000000..fd3a8fe04 --- /dev/null +++ b/pkg/agent/typed.go @@ -0,0 +1,79 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "encoding/json" + "fmt" + "reflect" + "strings" + + "go.probo.inc/probo/pkg/llm" +) + +type TypedResult[T any] struct { + Result + Output T +} + +func RunTyped[T any]( + ctx context.Context, + a *Agent, + messages []llm.Message, +) (*TypedResult[T], error) { + typed := a.clone() + + typed.responseFormat = &llm.ResponseFormat{ + Type: llm.ResponseFormatJSONSchema, + JSONSchema: &llm.JSONSchema{ + Name: typeName[T](), + Schema: jsonSchemaFor[T](), + Strict: true, + }, + } + + result, err := typed.Run(ctx, messages) + if err != nil { + return nil, err + } + + text := result.FinalMessage().Text() + + var output T + if err := json.Unmarshal([]byte(text), &output); err != nil { + return nil, fmt.Errorf("cannot parse typed output: %w", err) + } + + return &TypedResult[T]{ + Result: *result, + Output: output, + }, nil +} + +func typeName[T any]() string { + var zero T + t := reflect.TypeOf(zero) + if t.Kind() == reflect.Pointer { + t = t.Elem() + } + + name := t.Name() + if name == "" { + name = "output" + } + + return strings.ToLower(name) +} diff --git a/pkg/agent/typed_test.go b/pkg/agent/typed_test.go new file mode 100644 index 000000000..eed59b940 --- /dev/null +++ b/pkg/agent/typed_test.go @@ -0,0 +1,458 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package agent + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/llm" +) + +type typedMockProvider struct { + responses []*llm.ChatCompletionResponse + calls int +} + +func (m *typedMockProvider) ChatCompletion(_ context.Context, _ *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) { + if m.calls >= len(m.responses) { + return nil, errors.New("no more mock responses") + } + resp := m.responses[m.calls] + m.calls++ + return resp, nil +} + +func (m *typedMockProvider) ChatCompletionStream(_ context.Context, _ *llm.ChatCompletionRequest) (llm.ChatCompletionStream, error) { + return nil, errors.New("not implemented") +} + +func typedStopResponse(text string) *llm.ChatCompletionResponse { + return &llm.ChatCompletionResponse{ + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + Parts: []llm.Part{llm.TextPart{Text: text}}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + FinishReason: llm.FinishReasonStop, + } +} + +func TestTypeName(t *testing.T) { + t.Parallel() + + t.Run( + "named struct", + func(t *testing.T) { + t.Parallel() + + type UserInfo struct { + Name string `json:"name"` + } + + assert.Equal(t, "userinfo", typeName[UserInfo]()) + }, + ) + + t.Run( + "pointer to named struct", + func(t *testing.T) { + t.Parallel() + + type Invoice struct { + Amount int `json:"amount"` + } + + assert.Equal(t, "invoice", typeName[*Invoice]()) + }, + ) + + t.Run( + "anonymous struct returns output", + func(t *testing.T) { + t.Parallel() + + assert.Equal(t, "output", typeName[struct{ X int }]()) + }, + ) + + t.Run( + "basic string type", + func(t *testing.T) { + t.Parallel() + + assert.Equal(t, "string", typeName[string]()) + }, + ) + + t.Run( + "basic int type", + func(t *testing.T) { + t.Parallel() + + assert.Equal(t, "int", typeName[int]()) + }, + ) + + t.Run( + "custom type alias", + func(t *testing.T) { + t.Parallel() + + type Status string + + assert.Equal(t, "status", typeName[Status]()) + }, + ) +} + +func TestRunTyped(t *testing.T) { + t.Parallel() + + t.Run( + "successful structured output", + func(t *testing.T) { + t.Parallel() + + type CompanyInfo struct { + Name string `json:"name"` + Country string `json:"country"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`{"name":"Probo","country":"FR"}`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + WithInstructions("Return structured info."), + ) + + result, err := RunTyped[CompanyInfo]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Tell me about Probo"}}, + }}, + ) + + require.NoError(t, err) + assert.Equal(t, "Probo", result.Output.Name) + assert.Equal(t, "FR", result.Output.Country) + assert.Equal(t, 1, result.Turns) + assert.Equal(t, 10, result.Usage.InputTokens) + assert.Equal(t, 5, result.Usage.OutputTokens) + assert.Equal(t, "assistant", result.LastAgent.Name()) + }, + ) + + t.Run( + "nested struct output", + func(t *testing.T) { + t.Parallel() + + type Address struct { + City string `json:"city"` + Country string `json:"country"` + } + + type Person struct { + Name string `json:"name"` + Age int `json:"age"` + Address Address `json:"address"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`{"name":"Alice","age":30,"address":{"city":"Paris","country":"FR"}}`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + ) + + result, err := RunTyped[Person]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "Describe Alice"}}, + }}, + ) + + require.NoError(t, err) + assert.Equal(t, "Alice", result.Output.Name) + assert.Equal(t, 30, result.Output.Age) + assert.Equal(t, "Paris", result.Output.Address.City) + assert.Equal(t, "FR", result.Output.Address.Country) + }, + ) + + t.Run( + "invalid JSON response", + func(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`not valid json`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + ) + + _, err := RunTyped[Info]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "test"}}, + }}, + ) + + require.Error(t, err) + assert.Contains(t, err.Error(), "cannot parse typed output") + }, + ) + + t.Run( + "underlying Run failure propagates", + func(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{}, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + ) + + _, err := RunTyped[Info]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "test"}}, + }}, + ) + + require.Error(t, err) + }, + ) + + t.Run( + "does not mutate original agent", + func(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`{"name":"test"}`), + typedStopResponse("plain text"), + }, + } + + client := llm.NewClient(provider, "test") + + ag := New( + "assistant", + client, + WithModel("test-model"), + ) + + _, err := RunTyped[Info]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "structured"}}, + }}, + ) + require.NoError(t, err) + + assert.Nil(t, ag.responseFormat) + }, + ) + + t.Run( + "slice output", + func(t *testing.T) { + t.Parallel() + + type Items struct { + Names []string `json:"names"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`{"names":["Alice","Bob","Charlie"]}`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + ) + + result, err := RunTyped[Items]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "list names"}}, + }}, + ) + + require.NoError(t, err) + assert.Equal(t, []string{"Alice", "Bob", "Charlie"}, result.Output.Names) + }, + ) + + t.Run( + "context cancellation", + func(t *testing.T) { + t.Parallel() + + type Info struct { + Name string `json:"name"` + } + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + typedStopResponse(`{"name":"test"}`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + ) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := RunTyped[Info]( + ctx, + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "test"}}, + }}, + ) + + require.Error(t, err) + assert.Equal(t, 0, provider.calls) + }, + ) + + t.Run( + "with tool call", + func(t *testing.T) { + t.Parallel() + + type Summary struct { + City string `json:"city"` + Weather string `json:"weather"` + } + + type Params struct { + City string `json:"city"` + } + + weatherTool := FunctionTool[Params]( + "get_weather", + "Get weather for a city", + func(_ context.Context, p Params) (ToolResult, error) { + return ToolResult{Content: "Sunny, 22°C in " + p.City}, nil + }, + ) + + provider := &typedMockProvider{ + responses: []*llm.ChatCompletionResponse{ + { + Model: "test-model", + Message: llm.Message{ + Role: llm.RoleAssistant, + ToolCalls: []llm.ToolCall{{ + ID: "tc_1", + Function: llm.FunctionCall{ + Name: "get_weather", + Arguments: `{"city":"Paris"}`, + }, + }}, + }, + Usage: llm.Usage{InputTokens: 10, OutputTokens: 5}, + FinishReason: llm.FinishReasonToolCalls, + }, + typedStopResponse(`{"city":"Paris","weather":"Sunny, 22°C"}`), + }, + } + + ag := New( + "assistant", + llm.NewClient(provider, "test"), + WithModel("test-model"), + WithTools(weatherTool), + ) + + result, err := RunTyped[Summary]( + context.Background(), + ag, + []llm.Message{{ + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: "What's the weather in Paris?"}}, + }}, + ) + + require.NoError(t, err) + assert.Equal(t, "Paris", result.Output.City) + assert.Equal(t, "Sunny, 22°C", result.Output.Weather) + assert.Equal(t, 2, result.Turns) + }, + ) +} diff --git a/pkg/agents/changelog_generator.go b/pkg/agents/changelog_generator.go index 190073c4b..d2a79cb01 100644 --- a/pkg/agents/changelog_generator.go +++ b/pkg/agents/changelog_generator.go @@ -18,6 +18,7 @@ import ( "context" "fmt" + "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/llm" ) @@ -48,19 +49,30 @@ const ( ) func (a *Agent) GenerateChangelog(ctx context.Context, oldContent string, newContent string) (*string, error) { - resp, err := a.client.ChatCompletion(ctx, &llm.ChatCompletionRequest{ - Model: a.model, - Messages: []llm.Message{ - {Role: llm.RoleSystem, Parts: []llm.Part{llm.TextPart{Text: changelogGeneratorSystemPrompt}}}, - {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf(`Old content: %s`, oldContent)}}}, - {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf(`New content: %s`, newContent)}}}, + ag := agent.New( + "changelog_generator", + a.client, + agent.WithInstructions(changelogGeneratorSystemPrompt), + agent.WithModel(a.model), + agent.WithTemperature(a.temp), + ) + + result, err := ag.Run( + ctx, + []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{ + llm.TextPart{Text: fmt.Sprintf("Old content: %s", oldContent)}, + llm.TextPart{Text: fmt.Sprintf("New content: %s", newContent)}, + }, + }, }, - Temperature: &a.temp, - }) + ) if err != nil { return nil, fmt.Errorf("cannot generate changelog: %w", err) } - text := resp.Message.Text() + text := result.FinalMessage().Text() return &text, nil } diff --git a/pkg/agents/vendor_assessment.go b/pkg/agents/vendor_assessment.go index c89fd1b7e..922fd4fcf 100644 --- a/pkg/agents/vendor_assessment.go +++ b/pkg/agents/vendor_assessment.go @@ -16,9 +16,9 @@ package agents import ( "context" - "encoding/json" "fmt" + "go.probo.inc/probo/pkg/agent" "go.probo.inc/probo/pkg/llm" ) @@ -122,25 +122,27 @@ const ( ) func (a *Agent) AssessVendor(ctx context.Context, websiteURL string) (*vendorInfo, error) { - resp, err := a.client.ChatCompletion( + ag := agent.New( + "vendor_assessor", + a.client, + agent.WithInstructions(assessVendorSystemPrompt), + agent.WithModel(a.model), + agent.WithTemperature(a.temp), + ) + + typedResult, err := agent.RunTyped[vendorInfo]( ctx, - &llm.ChatCompletionRequest{ - Model: a.model, - Messages: []llm.Message{ - {Role: llm.RoleSystem, Parts: []llm.Part{llm.TextPart{Text: assessVendorSystemPrompt}}}, - {Role: llm.RoleUser, Parts: []llm.Part{llm.TextPart{Text: websiteURL}}}, + ag, + []llm.Message{ + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: websiteURL}}, }, - Temperature: &a.temp, }, ) if err != nil { return nil, fmt.Errorf("cannot assess vendor: %w", err) } - var info vendorInfo - if err := json.Unmarshal([]byte(resp.Message.Text()), &info); err != nil { - return nil, fmt.Errorf("cannot parse vendor info: %w", err) - } - - return &info, nil + return &typedResult.Output, nil } diff --git a/pkg/llm/chat.go b/pkg/llm/chat.go index 9b963b809..1eb85c178 100644 --- a/pkg/llm/chat.go +++ b/pkg/llm/chat.go @@ -21,15 +21,18 @@ import ( type ( ChatCompletionRequest struct { - Model string - Messages []Message - MaxTokens *int - Temperature *float64 - TopP *float64 - StopSequences []string - Tools []Tool - ToolChoice *ToolChoice - ResponseFormat *ResponseFormat + Model string + Messages []Message + MaxTokens *int + Temperature *float64 + TopP *float64 + FrequencyPenalty *float64 + PresencePenalty *float64 + StopSequences []string + Tools []Tool + ToolChoice *ToolChoice + ParallelToolCalls *bool + ResponseFormat *ResponseFormat } ToolChoiceType string @@ -50,6 +53,7 @@ type ( Name string Description string Schema json.RawMessage + Strict bool } FinishReason string