Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2026-03-13 19:58:48 +01:00
parent deac1538e5
commit d01ee914c1
3 changed files with 385 additions and 384 deletions

View File

@@ -246,10 +246,12 @@ type errorChatStream struct {
err error err error
} }
func (s *errorChatStream) Next() bool { return false } func (s *errorChatStream) Next() bool { return false }
func (s *errorChatStream) Event() llm.ChatCompletionStreamEvent { return llm.ChatCompletionStreamEvent{} } func (s *errorChatStream) Event() llm.ChatCompletionStreamEvent {
func (s *errorChatStream) Err() error { return s.err } return llm.ChatCompletionStreamEvent{}
func (s *errorChatStream) Close() error { return nil } }
func (s *errorChatStream) Err() error { return s.err }
func (s *errorChatStream) Close() error { return nil }
type errorStreamProvider struct { type errorStreamProvider struct {
err error err error
@@ -305,9 +307,6 @@ func toolCallResponse(toolCalls ...llm.ToolCall) *llm.ChatCompletionResponse {
} }
} }
func finishReasonPtr(r llm.FinishReason) *llm.FinishReason {
return &r
}
func TestRun(t *testing.T) { func TestRun(t *testing.T) {
t.Parallel() t.Parallel()
@@ -353,14 +352,14 @@ func TestRun(t *testing.T) {
City string `json:"city"` City string `json:"city"`
} }
weatherTool, err := agent.FunctionTool[Params]( weatherTool, err := agent.FunctionTool[Params](
"get_weather", "get_weather",
"Get weather for a city", "Get weather for a city",
func(_ context.Context, p Params) (agent.ToolResult, error) { func(_ context.Context, p Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "Sunny, 22°C in " + p.City}, nil return agent.ToolResult{Content: "Sunny, 22°C in " + p.City}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -412,31 +411,31 @@ func TestRun(t *testing.T) {
}, },
} }
type Params struct{} type Params struct{}
noopTool, err := agent.FunctionTool[Params]( noopTool, err := agent.FunctionTool[Params](
"noop", "noop",
"No-op", "No-op",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
ag := agent.New( ag := agent.New(
"assistant", "assistant",
newTestClient(provider), newTestClient(provider),
agent.WithModel("test-model"), agent.WithModel("test-model"),
agent.WithTools(noopTool), agent.WithTools(noopTool),
agent.WithMaxTurns(2), agent.WithMaxTurns(2),
) )
_, err = ag.Run( _, err = ag.Run(
context.Background(), context.Background(),
[]llm.Message{userMessage("loop")}, []llm.Message{userMessage("loop")},
) )
require.Error(t, err) require.Error(t, err)
var maxTurnsErr *agent.MaxTurnsExceededError var maxTurnsErr *agent.MaxTurnsExceededError
require.ErrorAs(t, err, &maxTurnsErr) require.ErrorAs(t, err, &maxTurnsErr)
assert.Equal(t, 2, maxTurnsErr.MaxTurns) assert.Equal(t, 2, maxTurnsErr.MaxTurns)
}, },
@@ -447,18 +446,18 @@ func TestRun(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
makeTool := func(name string) agent.Tool { makeTool := func(name string) agent.Tool {
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
name, name,
"desc", "desc",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
return tool return tool
} }
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -555,22 +554,22 @@ func TestRun(t *testing.T) {
type Params struct{} type Params struct{}
tool1, err := agent.FunctionTool[Params]( tool1, err := agent.FunctionTool[Params](
"first", "first",
"First tool", "First tool",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "result_1"}, nil return agent.ToolResult{Content: "result_1"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
tool2, err := agent.FunctionTool[Params]( tool2, err := agent.FunctionTool[Params](
"second", "second",
"Second tool", "Second tool",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "result_2"}, nil return agent.ToolResult{Content: "result_2"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -624,22 +623,22 @@ func TestRun(t *testing.T) {
type Params struct{} type Params struct{}
successTool, err := agent.FunctionTool[Params]( successTool, err := agent.FunctionTool[Params](
"succeed", "succeed",
"Always succeeds", "Always succeeds",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "success_result"}, nil return agent.ToolResult{Content: "success_result"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
failTool, err := agent.FunctionTool[Params]( failTool, err := agent.FunctionTool[Params](
"fail", "fail",
"Always fails", "Always fails",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{}, errors.New("tool exploded") return agent.ToolResult{}, errors.New("tool exploded")
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -697,17 +696,17 @@ func TestRun(t *testing.T) {
var capturedTenantID string var capturedTenantID string
type Params struct{} type Params struct{}
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
"check_tenant", "check_tenant",
"Check current tenant", "Check current tenant",
func(ctx context.Context, _ Params) (agent.ToolResult, error) { func(ctx context.Context, _ Params) (agent.ToolResult, error) {
rc := agent.RunContextFrom[*RequestContext](ctx) rc := agent.RunContextFrom[*RequestContext](ctx)
capturedTenantID = rc.TenantID capturedTenantID = rc.TenantID
return agent.ToolResult{Content: "tenant: " + rc.TenantID}, nil return agent.ToolResult{Content: "tenant: " + rc.TenantID}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -1070,27 +1069,27 @@ func TestRun_Hooks(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
noopTool, err := agent.FunctionTool[Params]( noopTool, err := agent.FunctionTool[Params](
"noop", "noop",
"No-op", "No-op",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc_1", ID: "tc_1",
Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, Function: llm.FunctionCall{Name: "noop", Arguments: `{}`},
}), }),
stopResponse("done"), stopResponse("done"),
}, },
} }
hook := &recordingHook{} hook := &recordingHook{}
ag := agent.New( ag := agent.New(
"assistant", "assistant",
@@ -1355,31 +1354,31 @@ func TestRun_ToolUseBehavior(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
"compute", "compute",
"Compute something", "Compute something",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "computed_result"}, nil return agent.ToolResult{Content: "computed_result"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc_1", ID: "tc_1",
Function: llm.FunctionCall{Name: "compute", Arguments: `{}`}, Function: llm.FunctionCall{Name: "compute", Arguments: `{}`},
}), }),
}, },
} }
ag := agent.New( ag := agent.New(
"assistant", "assistant",
newTestClient(provider), newTestClient(provider),
agent.WithModel("test-model"), agent.WithModel("test-model"),
agent.WithTools(tool), agent.WithTools(tool),
agent.WithToolUseBehavior(agent.StopOnFirstTool()), agent.WithToolUseBehavior(agent.StopOnFirstTool()),
) )
result, err := ag.Run( result, err := ag.Run(
@@ -1401,22 +1400,22 @@ func TestRun_ToolUseBehavior(t *testing.T) {
type Params struct{} type Params struct{}
tool1, err := agent.FunctionTool[Params]( tool1, err := agent.FunctionTool[Params](
"search", "search",
"Search", "Search",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "search_result"}, nil return agent.ToolResult{Content: "search_result"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
tool2, err := agent.FunctionTool[Params]( tool2, err := agent.FunctionTool[Params](
"submit", "submit",
"Submit", "Submit",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "submitted"}, nil return agent.ToolResult{Content: "submitted"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -1450,23 +1449,23 @@ func TestRun_ToolUseBehavior(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
"noop", "noop",
"No-op", "No-op",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc_1", ID: "tc_1",
Function: llm.FunctionCall{Name: "noop", Arguments: `{}`}, Function: llm.FunctionCall{Name: "noop", Arguments: `{}`},
}), }),
stopResponse("Final answer."), stopResponse("Final answer."),
}, },
} }
@@ -1493,42 +1492,42 @@ func TestRun_ToolUseBehavior(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
"compute", "compute",
"Compute something", "Compute something",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "result"}, nil return agent.ToolResult{Content: "result"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc_1", ID: "tc_1",
Function: llm.FunctionCall{Name: "compute", Arguments: `{}`}, Function: llm.FunctionCall{Name: "compute", Arguments: `{}`},
}), }),
}, },
} }
ag := agent.New( ag := agent.New(
"assistant", "assistant",
newTestClient(provider), newTestClient(provider),
agent.WithModel("test-model"), agent.WithModel("test-model"),
agent.WithTools(tool), agent.WithTools(tool),
agent.WithToolUseBehavior(agent.ToolUseBehavior(func(_ context.Context, _ []agent.ToolCallResult) (string, bool, error) { agent.WithToolUseBehavior(agent.ToolUseBehavior(func(_ context.Context, _ []agent.ToolCallResult) (string, bool, error) {
return "", false, errors.New("custom behavior failed") return "", false, errors.New("custom behavior failed")
})), })),
) )
_, err = ag.Run( _, err = ag.Run(
context.Background(), context.Background(),
[]llm.Message{userMessage("compute")}, []llm.Message{userMessage("compute")},
) )
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "custom behavior failed") assert.Contains(t, err.Error(), "custom behavior failed")
}, },
) )
} }
@@ -1583,38 +1582,38 @@ func TestRun_Approval(t *testing.T) {
}, },
} }
deleteTool, err := agent.FunctionTool[struct{}]( deleteTool, err := agent.FunctionTool[struct{}](
"delete_account", "delete_account",
"Deletes the user account", "Deletes the user account",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
return agent.ToolResult{Content: "deleted"}, nil return agent.ToolResult{Content: "deleted"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
ag := agent.New( ag := agent.New(
"assistant", "assistant",
newTestClient(provider), newTestClient(provider),
agent.WithModel("test-model"), agent.WithModel("test-model"),
agent.WithTools(deleteTool), agent.WithTools(deleteTool),
agent.WithApproval(agent.ApprovalConfig{ agent.WithApproval(agent.ApprovalConfig{
ToolNames: []string{"delete_account"}, ToolNames: []string{"delete_account"},
}), }),
) )
_, err = ag.Run( _, err = ag.Run(
context.Background(), context.Background(),
[]llm.Message{userMessage("Delete my account")}, []llm.Message{userMessage("Delete my account")},
) )
require.Error(t, err) require.Error(t, err)
var interrupted *agent.InterruptedError var interrupted *agent.InterruptedError
require.ErrorAs(t, err, &interrupted) require.ErrorAs(t, err, &interrupted)
assert.Len(t, interrupted.ToolCalls, 1) assert.Len(t, interrupted.ToolCalls, 1)
assert.Equal(t, "delete_account", interrupted.ToolCalls[0].Function.Name) assert.Equal(t, "delete_account", interrupted.ToolCalls[0].Function.Name)
assert.Len(t, interrupted.PendingApprovals, 1) assert.Len(t, interrupted.PendingApprovals, 1)
assert.Equal(t, "delete_account", interrupted.PendingApprovals[0].Function.Name) assert.Equal(t, "delete_account", interrupted.PendingApprovals[0].Function.Name)
assert.Equal(t, 1, interrupted.Turns) assert.Equal(t, 1, interrupted.Turns)
}, },
) )
@@ -1623,19 +1622,19 @@ func TestRun_Approval(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
var toolExecuted bool var toolExecuted bool
deleteTool, err := agent.FunctionTool[struct{}]( deleteTool, err := agent.FunctionTool[struct{}](
"delete_account", "delete_account",
"Deletes the user account", "Deletes the user account",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
toolExecuted = true toolExecuted = true
return agent.ToolResult{Content: "account deleted"}, nil return agent.ToolResult{Content: "account deleted"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc1", ID: "tc1",
@@ -1655,14 +1654,14 @@ func TestRun_Approval(t *testing.T) {
}), }),
) )
_, err = ag.Run( _, err = ag.Run(
context.Background(), context.Background(),
[]llm.Message{userMessage("Delete my account")}, []llm.Message{userMessage("Delete my account")},
) )
var interrupted *agent.InterruptedError var interrupted *agent.InterruptedError
require.ErrorAs(t, err, &interrupted) require.ErrorAs(t, err, &interrupted)
assert.False(t, toolExecuted) assert.False(t, toolExecuted)
result, err := agent.Resume( result, err := agent.Resume(
context.Background(), context.Background(),
@@ -1685,23 +1684,23 @@ func TestRun_Approval(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
deleteTool, err := agent.FunctionTool[struct{}]( deleteTool, err := agent.FunctionTool[struct{}](
"delete_account", "delete_account",
"Deletes the user account", "Deletes the user account",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
t.Fatal("tool should not be executed") t.Fatal("tool should not be executed")
return agent.ToolResult{}, nil return agent.ToolResult{}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc1", ID: "tc1",
Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`},
}), }),
stopResponse("OK, I won't delete your account."), stopResponse("OK, I won't delete your account."),
}, },
} }
@@ -1753,22 +1752,22 @@ func TestRun_Approval(t *testing.T) {
}, },
} }
safeTool, err := agent.FunctionTool[struct{}]( safeTool, err := agent.FunctionTool[struct{}](
"safe_tool", "safe_tool",
"A safe tool", "A safe tool",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
return agent.ToolResult{Content: "safe result"}, nil return agent.ToolResult{Content: "safe result"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
ag := agent.New( ag := agent.New(
"assistant", "assistant",
newTestClient(provider), newTestClient(provider),
agent.WithModel("test-model"), agent.WithModel("test-model"),
agent.WithTools(safeTool), agent.WithTools(safeTool),
agent.WithApproval(agent.ApprovalConfig{ agent.WithApproval(agent.ApprovalConfig{
ShouldApprove: func(_ context.Context, tc llm.ToolCall) bool { ShouldApprove: func(_ context.Context, tc llm.ToolCall) bool {
return tc.Function.Name == "dangerous_tool" return tc.Function.Name == "dangerous_tool"
}, },
}), }),
@@ -1791,25 +1790,25 @@ func TestRun_Approval(t *testing.T) {
var safeExecuted, dangerExecuted bool var safeExecuted, dangerExecuted bool
safeTool, err := agent.FunctionTool[struct{}]( safeTool, err := agent.FunctionTool[struct{}](
"safe_action", "safe_action",
"A safe action", "A safe action",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
safeExecuted = true safeExecuted = true
return agent.ToolResult{Content: "safe done"}, nil return agent.ToolResult{Content: "safe done"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
dangerTool, err := agent.FunctionTool[struct{}]( dangerTool, err := agent.FunctionTool[struct{}](
"danger_action", "danger_action",
"A dangerous action", "A dangerous action",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
dangerExecuted = true dangerExecuted = true
return agent.ToolResult{Content: "danger done"}, nil return agent.ToolResult{Content: "danger done"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -1945,22 +1944,22 @@ func TestResume(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
deleteTool, err := agent.FunctionTool[struct{}]( deleteTool, err := agent.FunctionTool[struct{}](
"delete_account", "delete_account",
"Deletes the user account", "Deletes the user account",
func(_ context.Context, _ struct{}) (agent.ToolResult, error) { func(_ context.Context, _ struct{}) (agent.ToolResult, error) {
return agent.ToolResult{Content: "deleted"}, nil return agent.ToolResult{Content: "deleted"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
toolCallResponse(llm.ToolCall{ toolCallResponse(llm.ToolCall{
ID: "tc1", ID: "tc1",
Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`}, Function: llm.FunctionCall{Name: "delete_account", Arguments: `{}`},
}), }),
stopResponse("Account deleted."), stopResponse("Account deleted."),
}, },
} }
@@ -2107,7 +2106,7 @@ func TestRunStreamed(t *testing.T) {
{ {
Delta: llm.MessageDelta{Content: "!"}, Delta: llm.MessageDelta{Content: "!"},
Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3}, Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3},
FinishReason: finishReasonPtr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
}, },
}, },
} }
@@ -2153,17 +2152,17 @@ func TestRunStreamed(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool, err := agent.FunctionTool[Params]( tool, err := agent.FunctionTool[Params](
"noop", "noop",
"No-op", "No-op",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
stream1 := &mockChatStream{ stream1 := &mockChatStream{
events: []llm.ChatCompletionStreamEvent{ events: []llm.ChatCompletionStreamEvent{
{ {
Delta: llm.MessageDelta{ Delta: llm.MessageDelta{
@@ -2172,7 +2171,7 @@ func TestRunStreamed(t *testing.T) {
}, },
}, },
Usage: &llm.Usage{InputTokens: 10, OutputTokens: 5}, Usage: &llm.Usage{InputTokens: 10, OutputTokens: 5},
FinishReason: finishReasonPtr(llm.FinishReasonToolCalls), FinishReason: new(llm.FinishReasonToolCalls),
}, },
}, },
} }
@@ -2183,7 +2182,7 @@ func TestRunStreamed(t *testing.T) {
{ {
Delta: llm.MessageDelta{}, Delta: llm.MessageDelta{},
Usage: &llm.Usage{InputTokens: 15, OutputTokens: 3}, Usage: &llm.Usage{InputTokens: 15, OutputTokens: 3},
FinishReason: finishReasonPtr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
}, },
}, },
} }
@@ -2239,7 +2238,7 @@ func TestRunStreamed(t *testing.T) {
{ {
Delta: llm.MessageDelta{Content: "!"}, Delta: llm.MessageDelta{Content: "!"},
Usage: &llm.Usage{InputTokens: 10, OutputTokens: 2}, Usage: &llm.Usage{InputTokens: 10, OutputTokens: 2},
FinishReason: finishReasonPtr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
}, },
}, },
} }
@@ -2288,7 +2287,7 @@ func TestRunStreamed(t *testing.T) {
{ {
Delta: llm.MessageDelta{Content: "!"}, Delta: llm.MessageDelta{Content: "!"},
Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3}, Usage: &llm.Usage{InputTokens: 10, OutputTokens: 3},
FinishReason: finishReasonPtr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
}, },
}, },
} }
@@ -2379,23 +2378,23 @@ func TestClone(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool1, err := agent.FunctionTool[Params]( tool1, err := agent.FunctionTool[Params](
"t1", "t1",
"desc", "desc",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
tool2, err := agent.FunctionTool[Params]( tool2, err := agent.FunctionTool[Params](
"t2", "t2",
"desc", "desc",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "ok"}, nil return agent.ToolResult{Content: "ok"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -2796,16 +2795,16 @@ func TestRun_HandoffWithPreHandoffTools(t *testing.T) {
var executionOrder []string var executionOrder []string
type Params struct{} type Params struct{}
tool1, err := agent.FunctionTool[Params]( tool1, err := agent.FunctionTool[Params](
"prepare", "prepare",
"Prepare data", "Prepare data",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
executionOrder = append(executionOrder, "prepare") executionOrder = append(executionOrder, "prepare")
return agent.ToolResult{Content: "prepared"}, nil return agent.ToolResult{Content: "prepared"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -2856,24 +2855,24 @@ func TestRun_HandoffWithPreHandoffTools(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
tool1, err := agent.FunctionTool[Params]( tool1, err := agent.FunctionTool[Params](
"prepare", "prepare",
"Prepare data", "Prepare data",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{Content: "prepared"}, nil return agent.ToolResult{Content: "prepared"}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
tool2, err := agent.FunctionTool[Params]( tool2, err := agent.FunctionTool[Params](
"finalize", "finalize",
"Finalize data", "Finalize data",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
t.Fatal("tool after handoff should not be executed") t.Fatal("tool after handoff should not be executed")
return agent.ToolResult{}, nil return agent.ToolResult{}, nil
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{
@@ -2942,15 +2941,15 @@ func TestRun_HandoffWithPreHandoffTools(t *testing.T) {
func(t *testing.T) { func(t *testing.T) {
t.Parallel() t.Parallel()
type Params struct{} type Params struct{}
failingTool, err := agent.FunctionTool[Params]( failingTool, err := agent.FunctionTool[Params](
"prepare", "prepare",
"Prepare data", "Prepare data",
func(_ context.Context, _ Params) (agent.ToolResult, error) { func(_ context.Context, _ Params) (agent.ToolResult, error) {
return agent.ToolResult{}, errors.New("preparation failed") return agent.ToolResult{}, errors.New("preparation failed")
}, },
) )
require.NoError(t, err) require.NoError(t, err)
provider := &mockProvider{ provider := &mockProvider{
responses: []*llm.ChatCompletionResponse{ responses: []*llm.ChatCompletionResponse{

View File

@@ -91,7 +91,7 @@ func spanAttrMap(recorder *tracetest.SpanRecorder) map[string]any {
return m return m
} }
func ptr[T any](v T) *T { return &v } //go:fix inline
// --------------------------------------------------------------------------- // ---------------------------------------------------------------------------
// Message.Text // Message.Text
@@ -357,7 +357,7 @@ func TestChatCompletionStream(t *testing.T) {
{Delta: llm.MessageDelta{Content: "Hello"}}, {Delta: llm.MessageDelta{Content: "Hello"}},
{Delta: llm.MessageDelta{Content: " world"}}, {Delta: llm.MessageDelta{Content: " world"}},
{ {
FinishReason: ptr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
Usage: &llm.Usage{InputTokens: 8, OutputTokens: 4}, Usage: &llm.Usage{InputTokens: 8, OutputTokens: 4},
}, },
} }
@@ -493,7 +493,7 @@ func TestChatCompletionStream(t *testing.T) {
events := []llm.ChatCompletionStreamEvent{ events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "done"}}, {Delta: llm.MessageDelta{Content: "done"}},
{ {
FinishReason: ptr(llm.FinishReasonLength), FinishReason: new(llm.FinishReasonLength),
Usage: &llm.Usage{InputTokens: 100, OutputTokens: 50}, Usage: &llm.Usage{InputTokens: 100, OutputTokens: 50},
}, },
} }
@@ -554,7 +554,7 @@ func TestStreamAccumulator(t *testing.T) {
}, },
}}, }},
{ {
FinishReason: ptr(llm.FinishReasonToolCalls), FinishReason: new(llm.FinishReasonToolCalls),
Usage: &llm.Usage{InputTokens: 20, OutputTokens: 15}, Usage: &llm.Usage{InputTokens: 20, OutputTokens: 15},
}, },
} }
@@ -604,7 +604,7 @@ func TestStreamAccumulator(t *testing.T) {
}, },
}}, }},
{ {
FinishReason: ptr(llm.FinishReasonToolCalls), FinishReason: new(llm.FinishReasonToolCalls),
Usage: &llm.Usage{InputTokens: 30, OutputTokens: 10}, Usage: &llm.Usage{InputTokens: 30, OutputTokens: 10},
}, },
} }
@@ -632,7 +632,7 @@ func TestStreamAccumulator(t *testing.T) {
events := []llm.ChatCompletionStreamEvent{ events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "Just text."}}, {Delta: llm.MessageDelta{Content: "Just text."}},
{ {
FinishReason: ptr(llm.FinishReasonStop), FinishReason: new(llm.FinishReasonStop),
Usage: &llm.Usage{InputTokens: 5, OutputTokens: 3}, Usage: &llm.Usage{InputTokens: 5, OutputTokens: 3},
}, },
} }
@@ -654,7 +654,7 @@ func TestStreamAccumulator(t *testing.T) {
events := []llm.ChatCompletionStreamEvent{ events := []llm.ChatCompletionStreamEvent{
{Delta: llm.MessageDelta{Content: "a"}}, {Delta: llm.MessageDelta{Content: "a"}},
{Delta: llm.MessageDelta{Content: "b"}}, {Delta: llm.MessageDelta{Content: "b"}},
{FinishReason: ptr(llm.FinishReasonStop)}, {FinishReason: new(llm.FinishReasonStop)},
} }
acc := llm.NewStreamAccumulator(&mockStream{events: events}) acc := llm.NewStreamAccumulator(&mockStream{events: events})

View File

@@ -14,6 +14,8 @@
package llm package llm
import "strings"
import "encoding/json" import "encoding/json"
type ( type (
@@ -42,11 +44,11 @@ type (
) )
func (m Message) Text() string { func (m Message) Text() string {
var s string var s strings.Builder
for _, p := range m.Parts { for _, p := range m.Parts {
if tp, ok := p.(TextPart); ok { if tp, ok := p.(TextPart); ok {
s += tp.Text s.WriteString(tp.Text)
} }
} }
return s return s.String()
} }