rony-llm-agent/pkg/agent/integration_test.go

158 lines
4.2 KiB
Go
Raw Permalink Normal View History

package agent_test
import (
"context"
"encoding/json"
"os"
"testing"
"github.com/VictorVargas/rony-llm-agent/pkg/agent"
"github.com/VictorVargas/rony-llm-agent/pkg/llm"
"github.com/VictorVargas/rony-llm-agent/pkg/llm/providers/llamacpp"
"github.com/VictorVargas/rony-llm-agent/pkg/persona"
"github.com/VictorVargas/rony-llm-agent/pkg/tools"
)
// TestIntegration_LlamaCPP_Generate is an integration test that requires llama.cpp running on localhost:8080.
// Run with: go test ./pkg/agent/ -run TestIntegration_LlamaCPP_Generate -tags=integration
func TestIntegration_LlamaCPP_Generate(t *testing.T) {
client, err := llamacpp.New(llamacpp.Config{
BaseURL: "http://localhost:8080/v1",
})
if err != nil {
t.Fatalf("failed to create client: %v", err)
}
resp, err := client.Generate(context.Background(), llm.CompletionRequest{
Messages: []llm.Message{
{Role: llm.RoleUser, Content: "What is 2+2? Answer with just the number."},
},
})
if err != nil {
t.Fatalf("generate failed: %v", err)
}
if resp.Content == "" {
t.Fatal("expected non-empty response")
}
}
// TestIntegration_LlamaCPP_Stream is an integration test that requires llama.cpp running on localhost:8080.
func TestIntegration_LlamaCPP_Stream(t *testing.T) {
client, err := llamacpp.New(llamacpp.Config{
BaseURL: "http://localhost:8080/v1",
})
if err != nil {
t.Fatalf("failed to create client: %v", err)
}
stream := client.Stream(context.Background(), llm.CompletionRequest{
Messages: []llm.Message{
{Role: llm.RoleUser, Content: "Say hello in 5 words."},
},
})
var chunks []llm.StreamChunk
for chunk, err := range stream {
if err != nil {
t.Fatalf("stream error: %v", err)
}
chunks = append(chunks, chunk)
}
if len(chunks) == 0 {
t.Fatal("expected at least one chunk")
}
}
// TestIntegration_AgentLoop_Generate is an integration test for the agent loop with llama.cpp.
func TestIntegration_AgentLoop_Generate(t *testing.T) {
client, err := llamacpp.New(llamacpp.Config{
BaseURL: "http://localhost:8080/v1",
})
if err != nil {
t.Fatalf("failed to create client: %v", err)
}
registry := tools.NewRegistry()
registry.Register(tools.Tool{
Name: "add_numbers",
Description: "Add two numbers together",
InputSchema: json.RawMessage(`{"type": "function", "function": {"name": "add_numbers", "description": "Add two numbers together", "parameters": {"type": "object", "properties": {"a": {"type": "number"}, "b": {"type": "number"}}}}}`),
Handler: func(ctx context.Context, args json.RawMessage) (tools.ToolResult, error) {
var input struct {
A float64 `json:"a"`
B float64 `json:"b"`
}
if err := json.Unmarshal(args, &input); err != nil {
return tools.ToolResult{IsError: true}, err
}
return tools.ToolResult{Content: "42"}, nil
},
Permission: tools.Allow,
})
loop := agent.New(agent.Config{
LLM: client,
Persona: persona.DefaultPersona(),
Tools: registry,
Sandbox: &mockSandbox{},
MaxIters: 3,
})
resp, err := loop.Run(context.Background(), "What is 20+22? Use the add_numbers tool.")
if err != nil {
t.Fatalf("run failed: %v", err)
}
if resp.Content == "" {
t.Fatal("expected non-empty response")
}
}
// TestIntegration_AgentLoop_Stream is an integration test for the agent loop with streaming.
func TestIntegration_AgentLoop_Stream(t *testing.T) {
client, err := llamacpp.New(llamacpp.Config{
BaseURL: "http://localhost:8080/v1",
})
if err != nil {
t.Fatalf("failed to create client: %v", err)
}
loop := agent.New(agent.Config{
LLM: client,
Persona: persona.DefaultPersona(),
Tools: tools.NewRegistry(),
MaxIters: 3,
})
stream := loop.RunStream(context.Background(), "Say something interesting.")
var chunks []llm.StreamChunk
for chunk, err := range stream {
if err != nil {
t.Fatalf("stream error: %v", err)
}
chunks = append(chunks, chunk)
}
if len(chunks) == 0 {
t.Fatal("expected at least one chunk")
}
}
// mockSandbox is a simple sandbox that allows all calls.
type mockSandbox struct{}
func (m *mockSandbox) ValidateToolCall(tool tools.Tool, call llm.ToolCall) error {
return nil
}
func TestMain(m *testing.M) {
// Skip integration tests unless explicitly enabled
if os.Getenv("INTEGRATION_TESTS") != "1" {
return
}
os.Exit(m.Run())
}