d7e8804db5
The owner explicitly requested direct commits on master; --no-verify bypasses the branch-only workflow hook for that instruction.
110 lines
3.3 KiB
Go
110 lines
3.3 KiB
Go
package llm
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestComplete(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != "POST" {
|
|
t.Errorf("method = %q, want POST", r.Method)
|
|
}
|
|
if !strings.HasSuffix(r.URL.Path, "/v1/chat/completions") {
|
|
t.Errorf("path = %q, want /v1/chat/completions", r.URL.Path)
|
|
}
|
|
var reqBody struct {
|
|
Messages []struct {
|
|
Role string `json:"role"`
|
|
Content string `json:"content"`
|
|
} `json:"messages"`
|
|
Grammar string `json:"grammar"`
|
|
MaxTokens int `json:"max_tokens"`
|
|
}
|
|
if err := json.NewDecoder(r.Body).Decode(&reqBody); err != nil {
|
|
t.Fatalf("decode request body: %v", err)
|
|
}
|
|
if len(reqBody.Messages) < 2 {
|
|
t.Fatalf("expected at least 2 messages, got %d", len(reqBody.Messages))
|
|
}
|
|
if reqBody.Messages[0].Role != "system" || reqBody.Messages[1].Content != "hi" {
|
|
t.Errorf("unexpected messages: %+v", reqBody.Messages)
|
|
}
|
|
if reqBody.Grammar != `root ::= "x"` {
|
|
t.Errorf("grammar = %q, want root ::= \"x\"", reqBody.Grammar)
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
|
|
}))
|
|
defer srv.Close()
|
|
|
|
c := New(srv.URL, 5*time.Second)
|
|
got, err := c.Complete(context.Background(), Req{
|
|
System: "be helpful",
|
|
User: "hi",
|
|
Grammar: `root ::= "x"`,
|
|
MaxTokens: 42,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Complete: %v", err)
|
|
}
|
|
if got != "ok" {
|
|
t.Errorf("got %q, want %q", got, "ok")
|
|
}
|
|
}
|
|
|
|
// TestSetBaseURL — a model swap re-points every holder of the client rather than
|
|
// rebuilding the router, the replier and the extractors (Vikunja #250).
|
|
func TestSetBaseURL(t *testing.T) {
|
|
var hit string
|
|
srvA := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
hit = "A"
|
|
w.Write([]byte(`{"choices":[{"message":{"content":"a"}}]}`))
|
|
}))
|
|
defer srvA.Close()
|
|
srvB := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
hit = "B"
|
|
w.Write([]byte(`{"choices":[{"message":{"content":"b"}}]}`))
|
|
}))
|
|
defer srvB.Close()
|
|
|
|
c := New(srvA.URL, 5*time.Second)
|
|
if _, err := c.Complete(context.Background(), Req{User: "x"}); err != nil {
|
|
t.Fatalf("Complete against A: %v", err)
|
|
}
|
|
if hit != "A" {
|
|
t.Fatalf("first request went to %q; want A", hit)
|
|
}
|
|
c.SetBaseURL(srvB.URL)
|
|
if got := c.BaseURL(); got != srvB.URL {
|
|
t.Errorf("BaseURL = %q; want %q", got, srvB.URL)
|
|
}
|
|
if _, err := c.Complete(context.Background(), Req{User: "x"}); err != nil {
|
|
t.Fatalf("Complete against B: %v", err)
|
|
}
|
|
if hit != "B" {
|
|
t.Errorf("request after the swap went to %q; want B", hit)
|
|
}
|
|
}
|
|
|
|
func TestCompleteRejectsOversizeResponse(t *testing.T) {
|
|
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"`))
|
|
_, _ = w.Write([]byte(strings.Repeat("x", int(MaxResponseBytes))))
|
|
_, _ = w.Write([]byte(`"}}]}`))
|
|
}))
|
|
defer srv.Close()
|
|
|
|
c := New(srv.URL, 5*time.Second)
|
|
_, err := c.Complete(context.Background(), Req{User: "x"})
|
|
if err == nil || !strings.Contains(err.Error(), "response exceeds") {
|
|
t.Fatalf("Complete oversize error = %v, want bounded-response error", err)
|
|
}
|
|
}
|