Files
ModelRouter/internal/provider/provider_test.go

301 lines
9.1 KiB
Go

package provider
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"path/filepath"
"sync/atomic"
"testing"
"time"
"llmsproxy/internal/config"
"llmsproxy/internal/lua"
"llmsproxy/internal/types"
)
func newTestProvider(t *testing.T, src config.Source) *Provider {
t.Helper()
vm := lua.NewVM(filepath.Join(t.TempDir(), "adapters"))
if err := vm.Start(); err != nil {
t.Fatalf("vm: %v", err)
}
t.Cleanup(vm.Stop)
return New(src, vm)
}
func src(name, url, adapter string, models ...string) config.Source {
s := config.Source{Name: name, BaseURL: url, Adapter: adapter, MaxConcurrent: 4}
for _, m := range models {
s.Models = append(s.Models, config.Model{ID: m, Priority: 0})
}
return s
}
func TestProviderChat(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var body map[string]interface{}
_ = json.NewDecoder(r.Body).Decode(&body)
if body["model"] != "mock-model" {
t.Errorf("model = %v", body["model"])
}
fmt.Fprint(w, `{"choices":[{"message":{"content":"hi"},"finish_reason":"stop"}]}`)
}))
defer up.Close()
p := newTestProvider(t, src("mock", up.URL, "openai", "mock-model"))
resp, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "mock-model",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hello")}},
})
if err != nil {
t.Fatalf("chat: %v", err)
}
if resp.Content != "hi" {
t.Fatalf("content = %q", resp.Content)
}
}
func TestProviderChatStream(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"a\"}}]}\n\n")
fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"b\"}}]}\n\n")
fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n")
fmt.Fprint(w, "data: [DONE]\n\n")
}))
defer up.Close()
p := newTestProvider(t, src("mock", up.URL, "openai", "m"))
ch, err := p.ChatStream(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if err != nil {
t.Fatalf("stream: %v", err)
}
var text string
var done bool
for ck := range ch {
text += ck.Content
if ck.Done {
done = true
}
}
if text != "ab" {
t.Fatalf("text = %q", text)
}
if !done {
t.Fatal("expected done")
}
}
// TestStreamSuccessClearsBackoff guards P6: a clean streaming end must reset
// a previously cooled (source, model) pair.
func TestStreamSuccessClearsBackoff(t *testing.T) {
var fail atomic.Bool
fail.Store(true)
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if fail.Load() {
w.WriteHeader(500)
return
}
w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(w, "data: {\"choices\":[{\"delta\":{\"content\":\"a\"}}]}\n\n")
fmt.Fprint(w, "data: [DONE]\n\n")
}))
defer up.Close()
p := newTestProvider(t, src("mock", up.URL, "openai", "m"))
if _, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
}); err == nil {
t.Fatal("expected first chat to fail")
}
if p.ModelAvailable("m") {
t.Fatal("m must be cooling after the failed chat")
}
fail.Store(false)
ch, err := p.ChatStream(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if err != nil {
t.Fatalf("stream: %v", err)
}
for range ch {
}
if !p.ModelAvailable("m") {
t.Fatal("clean stream must clear the cooldown")
}
if st := p.state("m"); st.FailCount() != 0 {
t.Fatalf("fail count after clean stream = %d", st.FailCount())
}
}
func TestProviderImage(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
fmt.Fprint(w, `{"created":123,"data":[{"b64_json":"QUJD"}]}`)
}))
defer up.Close()
p := newTestProvider(t, src("img", up.URL, "openai", "flux-1"))
resp, err := p.Image(context.Background(), &types.ImageGenRequest{Model: "flux-1", Prompt: "cat"})
if err != nil {
t.Fatalf("image: %v", err)
}
if len(resp.ImageData) != 1 || resp.ImageData[0].B64JSON != "QUJD" {
t.Fatalf("image data = %+v", resp.ImageData)
}
}
func TestProviderBackoffPerModel(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(500)
fmt.Fprint(w, "boom")
}))
defer up.Close()
// two models on one source: a failure on m1 must not blacklist m2
p := newTestProvider(t, src("mock", up.URL, "openai", "m1", "m2"))
_, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "m1",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if err == nil {
t.Fatal("expected error")
}
if p.ModelAvailable("m1") {
t.Fatal("expected m1 to be cooling down")
}
if !p.ModelAvailable("m2") {
t.Fatal("m2 must stay schedulable (per-model isolation)")
}
if !p.Available() {
t.Fatal("source must stay available while any model is schedulable")
}
if st := p.state("m1"); st.FailCount() != 1 {
t.Fatalf("fail count = %d, want 1", st.FailCount())
}
}
func TestProviderAuthFailureSelfHeals(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(401)
fmt.Fprint(w, `{"error":"API_KEY_DISABLED"}`)
}))
defer up.Close()
p := newTestProvider(t, src("mock2", up.URL, "openai", "m2"))
_, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "m2",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if err == nil {
t.Fatal("expected error")
}
st := p.state("m2")
if st.FailCount() != backoffCapN {
t.Fatalf("auth failure must jump to capped count, got %d", st.FailCount())
}
if until := st.CooldownUntil(); until <= time.Now().Add(25*time.Minute).Unix() {
t.Fatalf("auth failure must cool near the cap (until=%d)", until)
}
if st.Pref() != -2*int64(prefFailStep) {
t.Fatalf("auth failure pref penalty must be doubled, got %d", st.Pref())
}
// not permanent: the reset channel and a later success both restore it
st.reset()
if !p.ModelAvailable("m2") {
t.Fatal("reset must restore schedulability")
}
}
func TestModelStateCooldownAndRecovery(t *testing.T) {
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(500)
}))
defer up.Close()
p := newTestProvider(t, src("mock", up.URL, "openai", "m"))
_, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if err == nil {
t.Fatal("expected error")
}
st := p.state("m")
if st.FailCount() != 1 {
t.Fatalf("fail count = %d", st.FailCount())
}
// one failure -> 5s cooldown from now
until := st.CooldownUntil()
want := time.Now().Add(backoffBase).Unix()
if until < want-2 || until > want+2 {
t.Fatalf("cooldown = %d, want ~%d", until, want)
}
// success resets everything and bumps the preference
st.RecordSuccess()
if !p.ModelAvailable("m") {
t.Fatal("success must clear cooldown")
}
if st.FailCount() != 0 {
t.Fatalf("fail count after success = %d", st.FailCount())
}
if st.Pref() != 1-int64(prefFailStep) {
t.Fatalf("pref after one failure (-5) then success (+1) = %d, want %d", st.Pref(), 1-int64(prefFailStep))
}
}
func TestTryAcquire(t *testing.T) {
p := newTestProvider(t, src("mock", "http://127.0.0.1:1", "openai", "m"))
p.cfg.MaxConcurrent = 1
p.sem = make(chan struct{}, 1)
if err := p.TryAcquire(context.Background()); err != nil {
t.Fatalf("first acquire: %v", err)
}
if err := p.TryAcquire(context.Background()); !errors.Is(err, ErrBusy) {
t.Fatalf("second acquire = %v, want ErrBusy", err)
}
p.Release()
if err := p.TryAcquire(context.Background()); err != nil {
t.Fatalf("acquire after release: %v", err)
}
p.Release()
}
func TestChatBusyFailsFast(t *testing.T) {
release := make(chan struct{})
started := make(chan struct{}, 10)
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
started <- struct{}{}
<-release
fmt.Fprint(w, `{"choices":[{"message":{"content":"ok"}}]}`)
}))
defer up.Close()
p := newTestProvider(t, src("mock", up.URL, "openai", "m"))
p.cfg.MaxConcurrent = 1
p.sem = make(chan struct{}, 1)
done := make(chan error, 1)
go func() {
_, err := p.Chat(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
done <- err
}()
<-started // first request holds the only slot
// second request must fail fast with ErrBusy instead of queueing
_, err2 := p.Chat(context.Background(), &types.ChatRequest{
Model: "m",
Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("x")}},
})
if !errors.Is(err2, ErrBusy) {
t.Fatalf("second chat err = %v, want ErrBusy", err2)
}
close(release)
if err := <-done; err != nil {
t.Fatalf("first chat: %v", err)
}
}