Files
2026-08-15 22:31:12 +08:00

141 lines
6.0 KiB
Go

package modelcatalog
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"path/filepath"
"strings"
"testing"
"time"
)
const testCatalog = `{
"providers": {
"openai": {"id":"openai","name":"OpenAI","models":{
"gpt-5.6-sol":{"id":"gpt-5.6-sol","name":"GPT-5.6 Sol","cost":{"input":5,"output":30,"cache_read":0.5,"cache_write":6.25,"tiers":[{"input":10,"output":45,"cache_read":1,"cache_write":12.5,"tier":{"type":"context","size":272000}}]}},
"gpt-sol-alias":{"id":"gpt-5.6-sol","name":"GPT Sol Alias","cost":{"input":5,"output":30,"cache_read":0.5,"cache_write":6.25,"tiers":[{"input":10,"output":45,"cache_read":1,"cache_write":12.5,"tier":{"type":"context","size":272000}}]}},
"free-placeholder":{"id":"free-placeholder","cost":{"input":0,"output":0}},
"audio-special":{"id":"audio-special","cost":{"input":1,"output":2,"input_audio":3}}
}},
"google": {"id":"google","name":"Google","models":{
"gemini-flash":{"id":"gemini-flash","name":"Gemini Flash","cost":{"input":0.1,"output":0.4}}
}}
}
}`
func TestParseCatalogKeepsRepresentableNonZeroProviderPrices(t *testing.T) {
loaded, err := parseSource([]byte(testCatalog), DefaultSourceURL, time.Date(2026, 8, 15, 0, 0, 0, 0, time.UTC))
if err != nil {
t.Fatal(err)
}
if loaded.info.Models != 2 {
t.Fatalf("models = %d, want 2", loaded.info.Models)
}
entry, ok := loaded.byID["openai/gpt-5.6-sol"]
if !ok || entry.Base.InputMicrosPer1M != 5_000_000 || entry.Base.CacheReadMicrosPer1M != 500_000 || entry.Base.CacheWriteMicrosPer1M != 6_250_000 || entry.Base.OutputMicrosPer1M != 30_000_000 {
t.Fatalf("openai entry = %+v", entry)
}
if entry.LongContext == nil || entry.LongContext.ThresholdInputTokens != 272_000 || entry.LongContext.Rates.OutputMicrosPer1M != 45_000_000 {
t.Fatalf("long context = %+v", entry.LongContext)
}
gemini := loaded.byID["google/gemini-flash"]
if gemini.Base.CacheReadMicrosPer1M != gemini.Base.InputMicrosPer1M || gemini.Base.CacheWriteMicrosPer1M != gemini.Base.InputMicrosPer1M {
t.Fatalf("missing cache prices did not fall back to input: %+v", gemini.Base)
}
if _, exists := loaded.byID["openai/free-placeholder"]; exists {
t.Fatal("zero/zero placeholder was treated as a usable free price")
}
if _, exists := loaded.byID["openai/audio-special"]; exists {
t.Fatal("unrepresentable audio price was imported")
}
}
func TestParseCatalogDropsOnlyConflictingDuplicateID(t *testing.T) {
raw := []byte(`{"providers":{"provider":{"id":"provider","models":{"alias-a":{"id":"same","cost":{"input":1,"output":2}},"alias-b":{"id":"same","cost":{"input":3,"output":4}},"safe":{"id":"safe","cost":{"input":5,"output":6}}}}}}`)
loaded, err := parseSource(raw, DefaultSourceURL, time.Now().UTC())
if err != nil {
t.Fatal(err)
}
if loaded.info.Models != 1 {
t.Fatalf("models = %d, want only safe row", loaded.info.Models)
}
if _, exists := loaded.byID["provider/same"]; exists {
t.Fatal("conflicting duplicate ID was retained")
}
if _, exists := loaded.byID["provider/safe"]; !exists {
t.Fatal("unrelated safe row was dropped")
}
}
func TestManagerRefreshSearchCacheAndLastKnownGood(t *testing.T) {
body := testCatalog
status := http.StatusOK
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
response.WriteHeader(status)
_, _ = response.Write([]byte(body))
}))
defer server.Close()
cachePath := filepath.Join(t.TempDir(), "models-dev.json")
manager := NewManager(cachePath, server.URL)
if _, _, loaded := manager.Status(); loaded {
t.Fatal("new manager unexpectedly had a catalog")
}
first, err := manager.Refresh(context.Background())
if err != nil || first.Models != 2 || first.Revision == "" {
t.Fatalf("refresh info=%+v err=%v", first, err)
}
results, info, loaded := manager.Search("gpt-5.6", 20)
if !loaded || info.Revision != first.Revision || len(results) != 1 || results[0].ID != "openai/gpt-5.6-sol" {
t.Fatalf("search results=%+v info=%+v loaded=%v", results, info, loaded)
}
results, _, _ = manager.Search("openai/gpt-5.6-sol", 20)
if len(results) == 0 || results[0].ID != "openai/gpt-5.6-sol" {
t.Fatalf("exact catalog ID was not ranked first: %+v", results)
}
reopened := NewManager(cachePath, server.URL)
results, cached, loaded := reopened.Search("OpenAI", 20)
if !loaded || cached.Revision != first.Revision || len(results) != 1 {
t.Fatalf("cached search results=%+v info=%+v loaded=%v", results, cached, loaded)
}
wrongSource := NewManager(cachePath, "https://unreachable.invalid/catalog.json")
if _, lastError, loaded := wrongSource.Status(); loaded || !strings.Contains(lastError, "来源") {
t.Fatalf("wrong-source cache loaded=%v error=%q", loaded, lastError)
}
status = http.StatusBadGateway
if _, err := manager.Refresh(context.Background()); err == nil || !strings.Contains(err.Error(), "502") {
t.Fatalf("failed refresh error = %v", err)
}
entry, afterFailure, ok := manager.Lookup("openai/gpt-5.6-sol")
if !ok || entry.Model != "gpt-5.6-sol" || afterFailure.Revision != first.Revision {
t.Fatalf("last-known-good entry=%+v info=%+v ok=%v", entry, afterFailure, ok)
}
body = strings.ReplaceAll(testCatalog, `"input":5`, `"input":6`)
status = http.StatusOK
second, err := manager.Refresh(context.Background())
if err != nil || second.Revision == first.Revision {
t.Fatalf("second refresh info=%+v err=%v", second, err)
}
updated, _, _ := manager.Lookup("openai/gpt-5.6-sol")
if updated.Base.InputMicrosPer1M != 6_000_000 {
t.Fatalf("updated input = %d", updated.Base.InputMicrosPer1M)
}
}
func TestManagerRejectsOversizedCatalogWithoutReadingPastLimit(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, _ *http.Request) {
response.Header().Set("Content-Length", fmt.Sprint(maxDownloadBytes+1))
_, _ = response.Write([]byte("x"))
}))
defer server.Close()
manager := NewManager(filepath.Join(t.TempDir(), "catalog.json"), server.URL)
if _, err := manager.Refresh(context.Background()); err == nil || !strings.Contains(err.Error(), "超过") {
t.Fatalf("oversized refresh error = %v", err)
}
}