package proxy import ( "encoding/json" "io" "net/http" "net/http/httptest" "sort" "strings" "testing" "time" "github.com/Riotpiaole/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 }