Files
2026-09-11 06:14:38 +02:00

109 lines
3.0 KiB
Go

package main
import (
"bufio"
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
)
func TestOpenAIChatNonStreaming(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 12, promptTokens: 7, outputTokens: 3}
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"stream":false}`))
rr := httptest.NewRecorder()
s.openAIChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
var doc struct {
Model string `json:"model"`
Choices []struct {
Message struct {
Content string `json:"content"`
} `json:"message"`
} `json:"choices"`
Usage struct {
Prompt int64 `json:"prompt_tokens"`
Output int64 `json:"completion_tokens"`
} `json:"usage"`
}
if err := json.Unmarshal(rr.Body.Bytes(), &doc); err != nil {
t.Fatal(err)
}
if doc.Model != "mock:latest" || len(doc.Choices) != 1 || len(doc.Choices[0].Message.Content) != 12 {
t.Fatalf("unexpected response: %+v", doc)
}
if doc.Usage.Prompt != 7 || doc.Usage.Output != 3 {
t.Fatalf("unexpected usage: %+v", doc.Usage)
}
}
func TestOpenAIChatStreaming(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 16}
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"stream":true}`))
rr := httptest.NewRecorder()
s.openAIChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
if ct := rr.Header().Get("Content-Type"); !strings.Contains(ct, "text/event-stream") {
t.Fatalf("content-type=%q", ct)
}
if !strings.Contains(rr.Body.String(), "data: [DONE]") {
t.Fatalf("missing done marker: %s", rr.Body.String())
}
}
func TestNativeStreamingEndsWithUsage(t *testing.T) {
s := &server{model: "mock:latest", responseBytes: 8, promptTokens: 11, outputTokens: 5}
req := httptest.NewRequest(http.MethodPost, "/api/chat", strings.NewReader(`{"stream":true}`))
rr := httptest.NewRecorder()
s.nativeChat(rr, req)
if rr.Code != http.StatusOK {
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
}
scanner := bufio.NewScanner(strings.NewReader(rr.Body.String()))
var last map[string]any
for scanner.Scan() {
if err := json.Unmarshal(scanner.Bytes(), &last); err != nil {
t.Fatal(err)
}
}
if err := scanner.Err(); err != nil {
t.Fatal(err)
}
if done, _ := last["done"].(bool); !done {
t.Fatalf("last chunk not done: %#v", last)
}
if last["prompt_eval_count"] != float64(11) || last["eval_count"] != float64(5) {
t.Fatalf("unexpected usage: %#v", last)
}
}
func TestShouldFailConcurrent(t *testing.T) {
const total = 1000
s := &server{failEvery: 5}
var wg sync.WaitGroup
failures := make(chan struct{}, total)
for i := 0; i < total; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if s.shouldFail() {
failures <- struct{}{}
}
}()
}
wg.Wait()
close(failures)
if got, want := len(failures), total/5; got != want {
t.Fatalf("failures=%d want=%d", got, want)
}
if got := s.requests.Load(); got != total {
t.Fatalf("requests=%d want=%d", got, total)
}
}