Files
homelab-frontend/internal/proxy/models_endpoint_test.go
T

471 lines
11 KiB
Go
Raw Normal View History

package proxy
import (
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"sort"
"strings"
"testing"
"time"
"forgejo.riotpiao.com/rock/homelab-frontend/internal/config"
)
// ModelListResponse represents the response shape for GET /v1/models
type ModelListResponse struct {
Object string `json:"object"`
Data []ModelEntry `json:"data"`
}
// ModelEntry represents a single model in the list
type ModelEntry struct {
ID string `json:"id"`
Object string `json:"object"`
OwnedBy string `json:"owned_by"`
Created int64 `json:"created"`
}
// TestModelsEndpointReturns200 verifies GET /v1/models returns 200
func TestModelsEndpointReturns200(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Errorf("expected 200, got %d", resp.StatusCode)
}
}
// TestModelsEndpointContentType verifies correct content type
func TestModelsEndpointContentType(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
ct := resp.Header.Get("Content-Type")
if !strings.Contains(ct, "application/json") {
t.Errorf("expected content-type application/json, got %s", ct)
}
}
// TestModelsEndpointResponseShape verifies correct JSON structure
func TestModelsEndpointResponseShape(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
var result ModelListResponse
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
t.Fatalf("failed to decode response: %v", err)
}
if result.Object != "list" {
t.Errorf("expected object='list', got %q", result.Object)
}
if len(result.Data) != 1 {
t.Errorf("expected 1 model, got %d", len(result.Data))
}
model := result.Data[0]
if model.ID != "reasoning" {
t.Errorf("expected id='reasoning', got %q", model.ID)
}
if model.Object != "model" {
t.Errorf("expected object='model', got %q", model.Object)
}
}
// TestModelsEndpointEnumeratesAllModels verifies all models are listed
func TestModelsEndpointEnumeratesAllModels(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
"ornith:35b": {
Name: "ornith:35b",
Address: "localhost:9000",
},
"qwen2.5:3b-instruct": {
Name: "qwen2.5:3b-instruct",
Address: "localhost:9000",
},
"nomic-ai/nomic-embed-text-v2-moe": {
Name: "nomic-ai/nomic-embed-text-v2-moe",
Address: "localhost:9000",
},
"BAAI/bge-reranker-base": {
Name: "BAAI/bge-reranker-base",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
var result ModelListResponse
json.NewDecoder(resp.Body).Decode(&result)
if len(result.Data) != 5 {
t.Errorf("expected 5 models, got %d", len(result.Data))
}
// Collect actual model IDs
modelIDs := make(map[string]bool)
for _, model := range result.Data {
modelIDs[model.ID] = true
}
// Verify all expected models are present
expectedModels := []string{
"reasoning",
"ornith:35b",
"qwen2.5:3b-instruct",
"nomic-ai/nomic-embed-text-v2-moe",
"BAAI/bge-reranker-base",
}
for _, expected := range expectedModels {
if !modelIDs[expected] {
t.Errorf("expected model %q in response", expected)
}
}
}
// TestModelsEndpointHasRequiredFields verifies all required fields are present
func TestModelsEndpointHasRequiredFields(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
var result ModelListResponse
json.NewDecoder(resp.Body).Decode(&result)
model := result.Data[0]
if model.ID == "" {
t.Errorf("expected id field")
}
if model.Object == "" {
t.Errorf("expected object field")
}
if model.OwnedBy == "" {
t.Errorf("expected owned_by field")
}
if model.Created == 0 {
t.Errorf("expected created field (unix timestamp)")
}
}
// TestModelsEndpointNoUpstreamContact verifies endpoint doesn't contact upstream
func TestModelsEndpointNoUpstreamContact(t *testing.T) {
upstreamCalled := false
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalled = true
w.WriteHeader(http.StatusOK)
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: upstreamAddr,
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
_, _ = http.Get(server.URL + "/v1/models")
if upstreamCalled {
t.Errorf("upstream should not be called for /v1/models endpoint")
}
}
// TestModelsEndpointDerivedFromConfig verifies models come from config, not hardcoded
func TestModelsEndpointDerivedFromConfig(t *testing.T) {
// Create config with specific models
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"custom-model-1": {
Name: "custom-model-1",
Address: "localhost:9000",
},
"custom-model-2": {
Name: "custom-model-2",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer resp.Body.Close()
var result ModelListResponse
json.NewDecoder(resp.Body).Decode(&result)
// Verify only the configured models are returned
if len(result.Data) != 2 {
t.Errorf("expected 2 models from config, got %d", len(result.Data))
}
modelIDs := make([]string, len(result.Data))
for i, model := range result.Data {
modelIDs[i] = model.ID
}
sort.Strings(modelIDs)
expected := []string{"custom-model-1", "custom-model-2"}
if !equal(modelIDs, expected) {
t.Errorf("expected models %v, got %v", expected, modelIDs)
}
}
// TestModelsEndpointConsistentWithDispatch verifies advertised models can dispatch
func TestModelsEndpointConsistentWithDispatch(t *testing.T) {
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: upstreamAddr,
},
"ornith:35b": {
Name: "ornith:35b",
Address: upstreamAddr,
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
// Get list of models
resp, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
var result ModelListResponse
json.NewDecoder(resp.Body).Decode(&result)
resp.Body.Close()
// Try to dispatch to each advertised model
for _, model := range result.Data {
dispatchResp, err := http.Post(
server.URL+"/v1/chat/completions",
"application/json",
strings.NewReader(`{"model":"`+model.ID+`","messages":[]}`),
)
if err != nil {
t.Fatalf("request failed: %v", err)
}
defer dispatchResp.Body.Close()
// Should not return 400 (unknown model error)
if dispatchResp.StatusCode == http.StatusBadRequest {
body, _ := io.ReadAll(dispatchResp.Body)
if strings.Contains(string(body), "unknown model") {
t.Errorf("model %q advertised in /v1/models but not accepted for dispatch", model.ID)
}
}
}
}
// TestModelsEndpointResponseIsConsistent verifies response is consistent across calls
func TestModelsEndpointResponseIsConsistent(t *testing.T) {
cfg := &config.Config{
Models: map[string]*config.ModelUpstream{
"reasoning": {
Name: "reasoning",
Address: "localhost:9000",
},
"ornith:35b": {
Name: "ornith:35b",
Address: "localhost:9000",
},
},
Routes: make(map[string]*config.Route),
}
handler := New(cfg)
defer handler.Close()
server := httptest.NewServer(handler)
defer server.Close()
// Call endpoint twice
resp1, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
var result1 ModelListResponse
json.NewDecoder(resp1.Body).Decode(&result1)
resp1.Body.Close()
time.Sleep(10 * time.Millisecond)
resp2, err := http.Get(server.URL + "/v1/models")
if err != nil {
t.Fatalf("request failed: %v", err)
}
var result2 ModelListResponse
json.NewDecoder(resp2.Body).Decode(&result2)
resp2.Body.Close()
// Verify both responses have same models
if len(result1.Data) != len(result2.Data) {
t.Errorf("response length inconsistent: %d vs %d", len(result1.Data), len(result2.Data))
}
ids1 := make([]string, len(result1.Data))
ids2 := make([]string, len(result2.Data))
for i, m := range result1.Data {
ids1[i] = m.ID
}
for i, m := range result2.Data {
ids2[i] = m.ID
}
sort.Strings(ids1)
sort.Strings(ids2)
if !equal(ids1, ids2) {
t.Errorf("responses differ: %v vs %v", ids1, ids2)
}
}
// Helper function to compare string slices
func equal(a, b []string) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}