package llm import ( "context" "encoding/json" "fmt" "net/http" "net/http/httptest" "strings" "testing" ) // fakeServer mimics an OpenAI-compatible server that streams the given chunks. func fakeServer(t *testing.T, chunks []string) *httptest.Server { t.Helper() return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/v1/models": fmt.Fprint(w, `{"data":[{"id":"tiny-model"},{"id":"other"}]}`) case "/v1/chat/completions": if got := r.Header.Get("Authorization"); got != "Bearer secret" { http.Error(w, `{"error":"bad key"}`, http.StatusUnauthorized) return } var req struct { Model string `json:"model"` Stream bool `json:"stream"` Msgs []Message `json:"messages"` } if err := json.NewDecoder(r.Body).Decode(&req); err != nil { t.Errorf("decode request: %v", err) } if req.Model != "tiny-model" || !req.Stream || len(req.Msgs) != 1 { t.Errorf("unexpected request: %+v", req) } w.Header().Set("Content-Type", "text/event-stream") fmt.Fprint(w, ": keep-alive comment\n\n") fmt.Fprint(w, `data: {"choices":[{"delta":{"role":"assistant"}}]}`+"\n\n") for _, c := range chunks { b, _ := json.Marshal(c) fmt.Fprintf(w, `data: {"choices":[{"delta":{"content":%s}}]}`+"\n\n", b) } fmt.Fprint(w, "data: [DONE]\n\n") default: http.NotFound(w, r) } })) } func TestStream(t *testing.T) { srv := fakeServer(t, []string{"Hel", "lo", "\n", "world"}) defer srv.Close() c := &Client{BaseURL: srv.URL + "/v1", APIKey: "secret"} // no model: auto-pick first var got strings.Builder err := c.Stream(context.Background(), []Message{{Role: "user", Content: "hi"}}, func(d string) error { got.WriteString(d) return nil }) if err != nil { t.Fatal(err) } if got.String() != "Hello\nworld" { t.Fatalf("got %q", got.String()) } } func TestStreamHTTPError(t *testing.T) { srv := fakeServer(t, nil) defer srv.Close() c := &Client{BaseURL: srv.URL + "/v1", APIKey: "wrong", Model: "tiny-model"} err := c.Stream(context.Background(), []Message{{Role: "user", Content: "hi"}}, func(string) error { return nil }) if err == nil || !strings.Contains(err.Error(), "401") { t.Fatalf("want 401 error, got %v", err) } } func TestStreamCallbackErrorStops(t *testing.T) { srv := fakeServer(t, []string{"a", "b", "c"}) defer srv.Close() c := &Client{BaseURL: srv.URL + "/v1", APIKey: "secret"} stop := fmt.Errorf("stop") n := 0 err := c.Stream(context.Background(), []Message{{Role: "user", Content: "hi"}}, func(string) error { n++ return stop }) if err != stop || n != 1 { t.Fatalf("err=%v n=%d", err, n) } } func TestResolveModel(t *testing.T) { srv := fakeServer(t, nil) defer srv.Close() c := &Client{BaseURL: srv.URL + "/v1"} if m, err := c.ResolveModel(context.Background()); err != nil || m != "tiny-model" { t.Fatalf("auto: %q %v", m, err) } c.Model = "pinned" if m, _ := c.ResolveModel(context.Background()); m != "pinned" { t.Fatalf("pinned: %q", m) } }