112 lines
3.3 KiB
Go
112 lines
3.3 KiB
Go
|
|
package agent
|
||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"iter"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
|
||
|
|
"github.com/VictorVargas/rony-llm-agent/pkg/llm"
|
||
|
|
llmpersona "github.com/VictorVargas/rony-llm-agent/pkg/persona"
|
||
|
|
|
||
|
|
"github.com/VictorVargas/rony-chat-bot/internal/portfolio"
|
||
|
|
)
|
||
|
|
|
||
|
|
// stubClient is a minimal llm.LLMClient that echoes the system prompt's
|
||
|
|
// last "###" block as a single chunk, then a usage chunk.
|
||
|
|
type stubClient struct {
|
||
|
|
gotMessages []llm.Message
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *stubClient) Generate(_ context.Context, _ llm.CompletionRequest) (llm.CompletionResponse, error) {
|
||
|
|
return llm.CompletionResponse{Content: "ok", Usage: llm.TokenUsage{InputTokens: 1, OutputTokens: 1}}, nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *stubClient) Stream(_ context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] {
|
||
|
|
s.gotMessages = req.Messages
|
||
|
|
return func(yield func(llm.StreamChunk, error) bool) {
|
||
|
|
yield(llm.StreamChunk{Delta: "echo: " + req.Messages[len(req.Messages)-1].Content}, nil)
|
||
|
|
yield(llm.StreamChunk{Delta: " [done]", FinishReason: "stop", Usage: llm.TokenUsage{InputTokens: 7, OutputTokens: 3}}, nil)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func (s *stubClient) Name() string { return "stub" }
|
||
|
|
func (s *stubClient) Capabilities() llm.ProviderCapabilities {
|
||
|
|
return llm.ProviderCapabilities{MaxContextWindow: 4096}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRunnerStreamNoRAG(t *testing.T) {
|
||
|
|
cli := &stubClient{}
|
||
|
|
p := llmpersona.Persona{Name: "Tester", Tone: "concise", Language: "English"}
|
||
|
|
r := New(cli, p, "test system prompt", nil, 5)
|
||
|
|
|
||
|
|
var got strings.Builder
|
||
|
|
for chunk, err := range r.Stream(context.Background(), []llm.Message{{Role: RoleUser, Content: "hi"}}) {
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
got.WriteString(chunk.Delta)
|
||
|
|
}
|
||
|
|
want := "echo: hi [done]"
|
||
|
|
if got.String() != want {
|
||
|
|
t.Errorf("got %q, want %q", got.String(), want)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRunnerStreamWithRAG(t *testing.T) {
|
||
|
|
cli := &stubClient{}
|
||
|
|
p := llmpersona.Persona{Name: "Tester", Tone: "concise", Language: "English"}
|
||
|
|
|
||
|
|
// Build a tiny on-disk store with one project.
|
||
|
|
dir := t.TempDir()
|
||
|
|
dbPath := dir + "/t.db"
|
||
|
|
srcDir := dir + "/src"
|
||
|
|
if err := writeFile(srcDir+"/proj.md", "# Demo\n\n## Tech stack\n- Go\n- SQLite database\n"); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
store, err := portfolio.OpenStore(dbPath)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
defer store.Close()
|
||
|
|
if _, _, err := store.Reindex(context.Background(), srcDir, portfolio.DefaultChunkerConfig()); err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
|
||
|
|
r := New(cli, p, "test system prompt", store, 5)
|
||
|
|
for chunk, err := range r.Stream(context.Background(), []llm.Message{{Role: RoleUser, Content: "What database?"}}) {
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
_ = chunk
|
||
|
|
}
|
||
|
|
// The system message the LLM saw should include the RAG context.
|
||
|
|
if len(cli.gotMessages) == 0 {
|
||
|
|
t.Fatal("LLM never received messages")
|
||
|
|
}
|
||
|
|
sys := cli.gotMessages[0].Content
|
||
|
|
if !strings.Contains(sys, "Relevant context") {
|
||
|
|
t.Errorf("system prompt missing RAG context block:\n%s", sys)
|
||
|
|
}
|
||
|
|
if !strings.Contains(sys, "SQLite") {
|
||
|
|
t.Errorf("RAG context missing the SQLite content (got: %s)", sys)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func writeFile(path, content string) error {
|
||
|
|
if err := mkdirAll(path); err != nil {
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
return writeFileRaw(path, content)
|
||
|
|
}
|
||
|
|
|
||
|
|
// small helpers to avoid importing os/filepath in tests for a 2-call need
|
||
|
|
func mkdirAll(path string) error { return osMkdirAll(dirOf(path)) }
|
||
|
|
func dirOf(p string) string {
|
||
|
|
for i := len(p) - 1; i >= 0; i-- {
|
||
|
|
if p[i] == '/' {
|
||
|
|
return p[:i]
|
||
|
|
}
|
||
|
|
}
|
||
|
|
return "."
|
||
|
|
}
|