rony-chat-bot/internal/server/server_test.go

420 lines
12 KiB
Go
Raw Permalink Normal View History

package server
import (
"bufio"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"github.com/VictorVargas/rony-llm-agent/pkg/llm"
"github.com/VictorVargas/rony-llm-agent/pkg/llm/mock"
llmpersona "github.com/VictorVargas/rony-llm-agent/pkg/persona"
"github.com/VictorVargas/rony-chat-bot/internal/agent"
"github.com/VictorVargas/rony-chat-bot/internal/config"
"github.com/VictorVargas/rony-chat-bot/internal/portfolio"
)
// newTestServer wires a Server backed by a mock LLM, listens on a random
// port, and returns the base URL plus a teardown. The persona is fixed and
// the RAG store is nil (these tests don't exercise retrieval).
//
// To make /api/health pass, the LLM endpoint is pointed at a tiny stub
// HTTP server that returns 200 on /health. Callers can override this with
// newTestServerWithLLM(t, llmStatus).
func newTestServer(t *testing.T) (string, *config.Config) {
t.Helper()
return newTestServerWithLLM(t, http.StatusOK)
}
func newTestServerWithLLM(t *testing.T, llmStatus int) (string, *config.Config) {
t.Helper()
cli := mock.NewWithStream(mockChunks("hello", "world"))
p := llmpersona.Persona{Name: "TestBot", Tone: "concise", Language: "English"}
// Stub HTTP server that mimics llama-server's /health. If llmStatus != 200
// the health probe will report down.
llmStub := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/health" {
w.WriteHeader(llmStatus)
return
}
w.WriteHeader(http.StatusNotFound)
}))
t.Cleanup(llmStub.Close)
// Real (in-memory) FTS5 store so the store probe reports up. Without
// it the bot is "degraded" (200, but store down) — a confusing default.
dbPath := filepath.Join(t.TempDir(), "t.db")
store, err := portfolio.OpenStore(dbPath)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = store.Close() })
runner := agent.New(cli, p, "test system prompt", store, 5)
cfg := &config.Config{
Server: config.Server{
Host: "127.0.0.1",
Port: 0,
CORSOrigins: []string{"http://localhost:4321"},
RateLimit: config.RateLimit{RequestsPerMinute: 0, Burst: 0}, // disabled
},
Providers: []config.Provider{{
Name: "mock", Type: "llamacpp", Model: "test", Default: true,
Endpoint: llmStub.URL + "/v1",
}},
RAG: config.RAG{Enabled: false, TopK: 5, DataPath: ".", DBPath: dbPath},
Persona: config.Persona{Name: "TestBot", Language: "English"},
}
h := NewHandlers(cfg, runner, store, "test")
srv := New(cfg, h)
ts := httptest.NewUnstartedServer(srv.httpSrv.Handler)
ts.Start()
t.Cleanup(ts.Close)
return ts.URL, cfg
}
// mockChunks converts plain strings into a stream of single-word deltas
// followed by a final usage chunk, matching what a real provider would emit.
func mockChunks(words ...string) []llm.StreamChunk {
out := make([]llm.StreamChunk, 0, len(words)+1)
for _, w := range words {
out = append(out, llm.StreamChunk{Delta: w + " "})
}
out = append(out, llm.StreamChunk{
Delta: "",
FinishReason: "stop",
Usage: llm.TokenUsage{InputTokens: 5, OutputTokens: 3},
})
return out
}
func TestHealth(t *testing.T) {
url, _ := newTestServer(t)
resp, err := http.Get(url + "/api/health")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
var body HealthResponse
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
t.Fatal(err)
}
if body.Status != "healthy" {
t.Errorf("status = %q, want healthy", body.Status)
}
if body.Components["llm"].Status != "up" {
t.Errorf("llm component = %q, want up", body.Components["llm"].Status)
}
if body.Version == "" {
t.Error("version should be set")
}
}
func TestHealthLLMDown(t *testing.T) {
// Stub LLM that returns 500 on /health → bot should be "unhealthy" (503).
url, _ := newTestServerWithLLM(t, http.StatusInternalServerError)
resp, err := http.Get(url + "/api/health")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusServiceUnavailable {
t.Errorf("status = %d, want 503", resp.StatusCode)
}
var body HealthResponse
_ = json.NewDecoder(resp.Body).Decode(&body)
if body.Status != "unhealthy" {
t.Errorf("status = %q, want unhealthy", body.Status)
}
if body.Components["llm"].Status != "down" {
t.Errorf("llm component = %q, want down", body.Components["llm"].Status)
}
}
func TestHealthDeep(t *testing.T) {
url, _ := newTestServer(t)
resp, err := http.Get(url + "/api/health?deep=true")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
var body HealthResponse
_ = json.NewDecoder(resp.Body).Decode(&body)
if body.Components["store"].Status != "up" {
t.Errorf("store component = %q, want up", body.Components["store"].Status)
}
// Deep mode adds a chunks count to the store details.
if _, ok := body.Components["store"].Details["chunks"]; !ok {
t.Errorf("deep mode should include chunks count in details, got: %v", body.Components["store"].Details)
}
}
func TestInfo(t *testing.T) {
url, _ := newTestServer(t)
resp, err := http.Get(url + "/api/info")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
var body map[string]any
_ = json.NewDecoder(resp.Body).Decode(&body)
for _, k := range []string{"name", "version", "provider", "model", "rag", "top_k"} {
if _, ok := body[k]; !ok {
t.Errorf("missing field %q in /api/info: %v", k, body)
}
}
}
func TestChatSSE(t *testing.T) {
url, _ := newTestServer(t)
body := strings.NewReader(`{"messages":[{"role":"user","content":"hi"}],"stream":true}`)
resp, err := http.Post(url+"/api/chat", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if got := resp.Header.Get("Content-Type"); !strings.HasPrefix(got, "text/event-stream") {
t.Errorf("Content-Type = %q, want text/event-stream", got)
}
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
events := parseSSE(t, resp.Body)
// No RAG store in this test → no "sources" event.
wantTypes := []string{"start", "chunk", "chunk", "done"}
if len(events) < len(wantTypes) {
t.Fatalf("got %d events, want >= %d (%v)", len(events), len(wantTypes), events)
}
for i, want := range wantTypes {
if events[i].event != want {
t.Errorf("event[%d] = %q, want %q", i, events[i].event, want)
}
}
if events[1].data == "" {
t.Errorf("first chunk data is empty")
}
}
func TestChatSSEWithRAG(t *testing.T) {
// Build a temp store with one project so the sources event fires.
dir := t.TempDir()
if err := writeFileR(dir+"/src/proj.md", "# Demo\n\n## Tech stack\n- Go\n- SQLite database\n"); err != nil {
t.Fatal(err)
}
store, err := openStoreForTest(dir + "/t.db", dir+"/src")
if err != nil {
t.Fatal(err)
}
defer store.Close()
cli := mock.NewWithStream(mockChunks("answer"))
p := llmpersona.Persona{Name: "TestBot", Tone: "concise", Language: "English"}
runner := agent.New(cli, p, "test system prompt", store, 5)
cfg := &config.Config{
Server: config.Server{Host: "127.0.0.1", Port: 0, CORSOrigins: []string{"*"}},
Providers: []config.Provider{{Name: "mock", Type: "llamacpp", Model: "test", Default: true}},
RAG: config.RAG{Enabled: true, TopK: 5},
Persona: config.Persona{Name: "TestBot", Language: "English"},
}
h := NewHandlers(cfg, runner, store, "test")
srv := New(cfg, h)
ts := httptest.NewUnstartedServer(srv.httpSrv.Handler)
ts.Start()
defer ts.Close()
body := strings.NewReader(`{"messages":[{"role":"user","content":"what database?"}],"stream":true}`)
resp, err := http.Post(ts.URL+"/api/chat", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
events := parseSSE(t, resp.Body)
var sawSources bool
for _, e := range events {
if e.event == "sources" && strings.Contains(e.data, "proj") {
sawSources = true
}
}
if !sawSources {
t.Errorf("no sources event with project id, got events: %v", events)
}
}
func TestChatNoStream(t *testing.T) {
url, _ := newTestServer(t)
body := strings.NewReader(`{"messages":[{"role":"user","content":"hi"}],"stream":false}`)
resp, err := http.Post(url+"/api/chat", "application/json", body)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
if got := resp.Header.Get("Content-Type"); !strings.HasPrefix(got, "application/json") {
t.Errorf("Content-Type = %q, want application/json", got)
}
var out struct {
Content string `json:"content"`
Sources []string `json:"sources"`
Usage struct {
InputTokens int `json:"input_tokens"`
OutputTokens int `json:"output_tokens"`
} `json:"usage"`
}
if err := json.NewDecoder(resp.Body).Decode(&out); err != nil {
t.Fatal(err)
}
if out.Content == "" {
t.Error("content is empty")
}
if out.Usage.InputTokens == 0 && out.Usage.OutputTokens == 0 {
t.Error("usage is all zero")
}
}
func TestChatInvalidBody(t *testing.T) {
url, _ := newTestServer(t)
cases := []struct {
name, body string
}{
{"empty messages", `{"messages":[]}`},
{"empty content", `{"messages":[{"role":"user","content":""}]}`},
{"invalid role", `{"messages":[{"role":"wizard","content":"x"}]}`},
{"bad json", `{not json`},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
resp, err := http.Post(url+"/api/chat", "application/json", strings.NewReader(c.body))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 400 {
t.Errorf("status = %d, want 400", resp.StatusCode)
}
})
}
}
func TestChatWrongMethod(t *testing.T) {
url, _ := newTestServer(t)
resp, err := http.Get(url + "/api/chat")
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 405 {
t.Errorf("status = %d, want 405", resp.StatusCode)
}
}
func TestCORSReject(t *testing.T) {
url, _ := newTestServer(t)
body := strings.NewReader(`{"messages":[{"role":"user","content":"hi"}]}`)
req, _ := http.NewRequest("POST", url+"/api/chat", body)
req.Header.Set("Origin", "https://evil.example")
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 403 {
t.Errorf("status = %d, want 403", resp.StatusCode)
}
if aco := resp.Header.Get("Access-Control-Allow-Origin"); aco != "" {
t.Errorf("ACAO = %q, want empty", aco)
}
}
func TestCORSAllow(t *testing.T) {
url, _ := newTestServer(t)
body := strings.NewReader(`{"messages":[{"role":"user","content":"hi"}]}`)
req, _ := http.NewRequest("POST", url+"/api/chat", body)
req.Header.Set("Origin", "http://localhost:4321")
req.Header.Set("Content-Type", "application/json")
resp, err := http.DefaultClient.Do(req)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
t.Errorf("status = %d, want 200", resp.StatusCode)
}
if aco := resp.Header.Get("Access-Control-Allow-Origin"); aco != "http://localhost:4321" {
t.Errorf("ACAO = %q, want http://localhost:4321", aco)
}
}
// sseEvent is one parsed "event: x\ndata: y\n\n" record.
type sseEvent struct {
event string
data string
}
func parseSSE(t *testing.T, r io.Reader) []sseEvent {
t.Helper()
var out []sseEvent
scanner := bufio.NewScanner(r)
scanner.Buffer(make([]byte, 0, 64*1024), 4*1024*1024)
var cur sseEvent
flush := func() {
if cur.event != "" || cur.data != "" {
out = append(out, cur)
}
cur = sseEvent{}
}
for scanner.Scan() {
line := scanner.Text()
switch {
case line == "":
flush()
case strings.HasPrefix(line, "event: "):
cur.event = strings.TrimPrefix(line, "event: ")
case strings.HasPrefix(line, "data: "):
if cur.data != "" {
cur.data += "\n"
}
cur.data += strings.TrimPrefix(line, "data: ")
}
}
flush()
if err := scanner.Err(); err != nil {
t.Fatal(err)
}
return out
}
// Verify the unused imports are not actually unused (httptest, context, etc.).
var _ = httptest.NewRecorder
// writeFileR + openStoreForTest are tiny shims so the test file doesn't
// need to import os/filepath directly.
func writeFileR(path, content string) error {
return osWriteFile(path, []byte(content), 0o644)
}
func openStoreForTest(dbPath, srcDir string) (ragStore, error) {
return ragOpenStore(dbPath, srcDir)
}