79 lines
2.9 KiB
Go
79 lines
2.9 KiB
Go
package ollama
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func poolTestServer(t *testing.T, calls *atomic.Int64, fail bool) *httptest.Server {
|
|
t.Helper()
|
|
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
switch r.URL.Path {
|
|
case "/api/tags":
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"models": []map[string]any{{"name": "qwen3:8b", "digest": "chat"}, {"name": "embeddinggemma:latest", "digest": "embed"}}})
|
|
case "/api/chat":
|
|
calls.Add(1)
|
|
if fail {
|
|
http.Error(w, "busy", http.StatusServiceUnavailable)
|
|
return
|
|
}
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"message": map[string]any{"content": `{"ok":true}`}})
|
|
case "/api/embed":
|
|
calls.Add(1)
|
|
_ = json.NewEncoder(w).Encode(map[string]any{"embeddings": [][]float64{{0.1, 0.2}}})
|
|
default:
|
|
http.NotFound(w, r)
|
|
}
|
|
}))
|
|
}
|
|
|
|
func TestPoolLeastInflightBalancesSerialRequests(t *testing.T) {
|
|
var aCalls, bCalls atomic.Int64
|
|
a := poolTestServer(t, &aCalls, false)
|
|
defer a.Close()
|
|
b := poolTestServer(t, &bCalls, false)
|
|
defer b.Close()
|
|
c := NewPool(PoolConfig{Nodes: []NodeConfig{{Name: "a", URL: a.URL}, {Name: "b", URL: b.URL}}, RoutingMode: "least_inflight", NodeMaxInflight: 1, HealthInterval: time.Minute, FailureCooldown: time.Second, RequestTimeout: time.Second, FailoverEnabled: true, RequireSameModelDigest: true, RequireEmbeddingModel: true}, "qwen3:8b", "embeddinggemma")
|
|
if err := c.Ping(context.Background()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for i := 0; i < 4; i++ {
|
|
var out struct {
|
|
OK bool `json:"ok"`
|
|
}
|
|
if err := c.ChatJSON(context.Background(), "system", "user", map[string]any{"type": "object"}, &out); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
}
|
|
if aCalls.Load() != 2 || bCalls.Load() != 2 {
|
|
t.Fatalf("distribution a=%d b=%d", aCalls.Load(), bCalls.Load())
|
|
}
|
|
}
|
|
|
|
func TestPoolFailsOverOnRetryableError(t *testing.T) {
|
|
var badCalls, goodCalls atomic.Int64
|
|
bad := poolTestServer(t, &badCalls, true)
|
|
defer bad.Close()
|
|
good := poolTestServer(t, &goodCalls, false)
|
|
defer good.Close()
|
|
c := NewPool(PoolConfig{Nodes: []NodeConfig{{Name: "a-bad", URL: bad.URL}, {Name: "b-good", URL: good.URL}}, RoutingMode: "least_inflight", NodeMaxInflight: 1, HealthInterval: time.Minute, FailureCooldown: time.Second, RequestTimeout: time.Second, FailoverEnabled: true, FailoverAttempts: 2, RequireSameModelDigest: true, RequireEmbeddingModel: true}, "qwen3:8b", "embeddinggemma")
|
|
if err := c.Ping(context.Background()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var out struct {
|
|
OK bool `json:"ok"`
|
|
}
|
|
if err := c.ChatJSON(context.Background(), "system", "user", map[string]any{"type": "object"}, &out); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if badCalls.Load() != 1 || goodCalls.Load() != 1 {
|
|
t.Fatalf("failover bad=%d good=%d", badCalls.Load(), goodCalls.Load())
|
|
}
|
|
}
|