// Package modelcatalog downloads and indexes models.dev reference prices. // The catalog is deliberately separate from the effective local price table: // callers must explicitly import a row before it can affect billing. package modelcatalog import ( "bytes" "context" "crypto/sha256" "encoding/hex" "encoding/json" "errors" "fmt" "io" "math/big" "net/http" "os" "path/filepath" "sort" "strings" "sync" "time" "billing/internal/pricing" ) const ( DefaultSourceURL = "https://models.dev/catalog.json" maxDownloadBytes = 32 << 20 requestTimeout = 30 * time.Second ) // Info describes the currently cached, normalized catalog snapshot. type Info struct { SourceURL string `json:"source_url"` FetchedAt time.Time `json:"fetched_at"` Revision string `json:"revision"` Models int `json:"models"` } // Entry is one provider-specific reference price that can be represented by // billing's deterministic token price model. type Entry struct { ID string `json:"id"` Provider string `json:"provider"` ProviderName string `json:"provider_name"` Model string `json:"model"` ModelName string `json:"model_name"` Base pricing.Rates `json:"base"` LongContext *pricing.LongContext `json:"long_context,omitempty"` } // Policy creates a local price policy from the reference row. Fast pricing is // local business policy and is therefore supplied by the caller, not models.dev. func (e Entry) Policy(localModel string, fastEnabled bool, multiplier pricing.Ratio) pricing.Policy { longContext := cloneLongContext(e.LongContext) return pricing.Policy{ Model: strings.TrimSpace(localModel), Base: e.Base, LongContext: longContext, FastPricingEnabled: fastEnabled, FastMultiplier: multiplier, } } type cacheDocument struct { Info Info `json:"info"` Entries []Entry `json:"entries"` } type index struct { info Info entries []Entry byID map[string]Entry } // Manager owns one immutable in-memory catalog index and a compact cache file. // Refresh downloads and parses outside the read lock, so catalog searches and // CPA request handling never wait on the network. type Manager struct { mu sync.RWMutex refreshMu sync.Mutex cachePath string sourceURL string client *http.Client loaded *index lastError string } func NewManager(cachePath, sourceURL string) *Manager { sourceURL = strings.TrimSpace(sourceURL) if sourceURL == "" { sourceURL = DefaultSourceURL } transport := &http.Transport{Proxy: http.ProxyFromEnvironment, DisableKeepAlives: true} manager := &Manager{ cachePath: strings.TrimSpace(cachePath), sourceURL: sourceURL, client: &http.Client{Timeout: requestTimeout, Transport: transport}, } if err := manager.loadCache(); err != nil && !errors.Is(err, os.ErrNotExist) { manager.lastError = err.Error() } return manager } // Status reports cache availability without attempting network access. func (m *Manager) Status() (Info, string, bool) { m.mu.RLock() defer m.mu.RUnlock() if m.loaded == nil { return Info{SourceURL: m.sourceURL}, m.lastError, false } return m.loaded.info, m.lastError, true } // Refresh downloads and atomically publishes a new last-known-good snapshot. // A failed refresh leaves both the in-memory index and cache file untouched. func (m *Manager) Refresh(ctx context.Context) (Info, error) { m.refreshMu.Lock() defer m.refreshMu.Unlock() request, err := http.NewRequestWithContext(ctx, http.MethodGet, m.sourceURL, nil) if err != nil { return Info{}, m.rememberError(fmt.Errorf("创建 models.dev 请求: %w", err)) } request.Header.Set("Accept", "application/json") request.Header.Set("User-Agent", "billing/models.dev") response, err := m.client.Do(request) if err != nil { return Info{}, m.rememberError(fmt.Errorf("下载 models.dev 价格目录: %w", err)) } defer response.Body.Close() if response.StatusCode != http.StatusOK { return Info{}, m.rememberError(fmt.Errorf("下载 models.dev 价格目录: HTTP %s", response.Status)) } if response.ContentLength > maxDownloadBytes { return Info{}, m.rememberError(fmt.Errorf("models.dev 价格目录超过 %d MiB 限制", maxDownloadBytes>>20)) } raw, err := io.ReadAll(io.LimitReader(response.Body, maxDownloadBytes+1)) if err != nil { return Info{}, m.rememberError(fmt.Errorf("读取 models.dev 价格目录: %w", err)) } if len(raw) > maxDownloadBytes { return Info{}, m.rememberError(fmt.Errorf("models.dev 价格目录超过 %d MiB 限制", maxDownloadBytes>>20)) } loaded, err := parseSource(raw, m.sourceURL, time.Now().UTC()) if err != nil { return Info{}, m.rememberError(fmt.Errorf("解析 models.dev 价格目录: %w", err)) } if err := m.writeCache(loaded); err != nil { return Info{}, m.rememberError(err) } m.mu.Lock() m.loaded = loaded m.lastError = "" m.mu.Unlock() return loaded.info, nil } func (m *Manager) rememberError(err error) error { m.mu.Lock() m.lastError = err.Error() m.mu.Unlock() return err } // Search ranks exact model/ID matches before prefixes and substrings. func (m *Manager) Search(query string, limit int) ([]Entry, Info, bool) { query = strings.ToLower(strings.TrimSpace(query)) if limit <= 0 || limit > 50 { limit = 20 } m.mu.RLock() defer m.mu.RUnlock() if m.loaded == nil { return nil, Info{SourceURL: m.sourceURL}, false } if query == "" { return []Entry{}, m.loaded.info, true } groups := [4][]Entry{} for _, entry := range m.loaded.entries { fields := []string{strings.ToLower(entry.ID), strings.ToLower(entry.Model), strings.ToLower(entry.ModelName), strings.ToLower(entry.Provider), strings.ToLower(entry.ProviderName)} group := -1 switch { case fields[0] == query: group = 0 case fields[1] == query || fields[2] == query || fields[3] == query || fields[4] == query: group = 1 default: for _, field := range fields { if strings.HasPrefix(field, query) { group = 2 break } if group < 0 && strings.Contains(field, query) { group = 3 } } } if group >= 0 { groups[group] = append(groups[group], cloneEntry(entry)) } } result := make([]Entry, 0, limit) for _, group := range groups { for _, entry := range group { result = append(result, entry) if len(result) == limit { return result, m.loaded.info, true } } } return result, m.loaded.info, true } func (m *Manager) Lookup(id string) (Entry, Info, bool) { m.mu.RLock() defer m.mu.RUnlock() if m.loaded == nil { return Entry{}, Info{SourceURL: m.sourceURL}, false } entry, ok := m.loaded.byID[strings.ToLower(strings.TrimSpace(id))] return cloneEntry(entry), m.loaded.info, ok } func (m *Manager) loadCache() error { if m.cachePath == "" { return os.ErrNotExist } raw, err := os.ReadFile(m.cachePath) if err != nil { return err } var document cacheDocument if err := json.Unmarshal(raw, &document); err != nil { return fmt.Errorf("读取 models.dev 本地缓存: %w", err) } if strings.TrimSpace(document.Info.SourceURL) != m.sourceURL { return fmt.Errorf("读取 models.dev 本地缓存: 缓存来源与当前 models_dev_url 不一致") } loaded, err := indexDocument(document) if err != nil { return fmt.Errorf("读取 models.dev 本地缓存: %w", err) } m.loaded = loaded return nil } func (m *Manager) writeCache(loaded *index) error { if m.cachePath == "" { return nil } directory := filepath.Dir(m.cachePath) if err := os.MkdirAll(directory, 0o755); err != nil { return fmt.Errorf("创建 models.dev 缓存目录: %w", err) } raw, err := json.Marshal(cacheDocument{Info: loaded.info, Entries: loaded.entries}) if err != nil { return fmt.Errorf("编码 models.dev 本地缓存: %w", err) } temporary, err := os.CreateTemp(directory, ".billing-models-dev-*.tmp") if err != nil { return fmt.Errorf("创建 models.dev 临时缓存: %w", err) } temporaryPath := temporary.Name() defer os.Remove(temporaryPath) if chmodErr := temporary.Chmod(0o600); chmodErr != nil { _ = temporary.Close() return fmt.Errorf("设置 models.dev 本地缓存权限: %w", chmodErr) } _, err = temporary.Write(raw) if err == nil { err = temporary.Sync() } closeErr := temporary.Close() if err != nil { return fmt.Errorf("写入 models.dev 本地缓存: %w", err) } if closeErr != nil { return fmt.Errorf("关闭 models.dev 本地缓存: %w", closeErr) } if err := os.Rename(temporaryPath, m.cachePath); err != nil { return fmt.Errorf("替换 models.dev 本地缓存: %w", err) } return nil } type sourceCatalog struct { Providers map[string]sourceProvider `json:"providers"` } type sourceProvider struct { ID string `json:"id"` Name string `json:"name"` Models map[string]json.RawMessage `json:"models"` } type sourceModel struct { ID string `json:"id"` Name string `json:"name"` Cost *sourceCost `json:"cost"` } type sourceCost struct { Input *json.Number `json:"input"` Output *json.Number `json:"output"` Reasoning *json.Number `json:"reasoning"` CacheRead *json.Number `json:"cache_read"` CacheWrite *json.Number `json:"cache_write"` InputAudio *json.Number `json:"input_audio"` OutputAudio *json.Number `json:"output_audio"` Tiers []sourceCostTier `json:"tiers"` } type sourceCostTier struct { Input *json.Number `json:"input"` Output *json.Number `json:"output"` Reasoning *json.Number `json:"reasoning"` CacheRead *json.Number `json:"cache_read"` CacheWrite *json.Number `json:"cache_write"` InputAudio *json.Number `json:"input_audio"` OutputAudio *json.Number `json:"output_audio"` Tier struct { Type string `json:"type"` Size int64 `json:"size"` } `json:"tier"` } func parseSource(raw []byte, sourceURL string, fetchedAt time.Time) (*index, error) { var source sourceCatalog decoder := json.NewDecoder(bytes.NewReader(raw)) decoder.UseNumber() if err := decoder.Decode(&source); err != nil { return nil, err } entryByID := make(map[string]Entry) ambiguousIDs := make(map[string]struct{}) for providerKey, provider := range source.Providers { providerID := normalizedID(provider.ID, providerKey) if providerID == "" { continue } providerName := strings.TrimSpace(provider.Name) if providerName == "" { providerName = providerID } for modelKey, rawModel := range provider.Models { var model sourceModel modelDecoder := json.NewDecoder(bytes.NewReader(rawModel)) modelDecoder.UseNumber() if modelDecoder.Decode(&model) != nil || model.Cost == nil { continue } modelID := strings.TrimSpace(model.ID) if modelID == "" { modelID = strings.TrimSpace(modelKey) } entry, ok := entryFromCost(providerID, providerName, modelID, model.Name, model.Cost) if ok { if _, ambiguous := ambiguousIDs[entry.ID]; ambiguous { continue } if existing, duplicate := entryByID[entry.ID]; duplicate { if sameEntryPrice(existing, entry) { continue } delete(entryByID, entry.ID) ambiguousIDs[entry.ID] = struct{}{} continue } entryByID[entry.ID] = entry } } } entries := make([]Entry, 0, len(entryByID)) for _, entry := range entryByID { entries = append(entries, entry) } if len(entries) == 0 { return nil, errors.New("目录中没有可安全表示的非零 Token 价格") } sort.Slice(entries, func(i, j int) bool { return entries[i].ID < entries[j].ID }) digest := sha256.Sum256(raw) document := cacheDocument{ Info: Info{SourceURL: sourceURL, FetchedAt: fetchedAt, Revision: hex.EncodeToString(digest[:]), Models: len(entries)}, Entries: entries, } return indexDocument(document) } func indexDocument(document cacheDocument) (*index, error) { if strings.TrimSpace(document.Info.Revision) == "" || document.Info.FetchedAt.IsZero() || len(document.Entries) == 0 { return nil, errors.New("缓存元数据不完整") } loaded := &index{info: document.Info, entries: make([]Entry, 0, len(document.Entries)), byID: make(map[string]Entry, len(document.Entries))} for _, entry := range document.Entries { entry.ID = strings.ToLower(strings.TrimSpace(entry.ID)) if entry.ID == "" { return nil, errors.New("缓存包含空目录 ID") } if _, duplicate := loaded.byID[entry.ID]; duplicate { return nil, fmt.Errorf("缓存包含重复目录 ID %q", entry.ID) } policy := entry.Policy("catalog-validation", false, pricing.Ratio{Numerator: 1, Denominator: 1}) if err := policy.Validate(); err != nil { return nil, fmt.Errorf("缓存价格 %q 无效: %w", entry.ID, err) } entry = cloneEntry(entry) loaded.entries = append(loaded.entries, entry) loaded.byID[entry.ID] = entry } loaded.info.Models = len(loaded.entries) return loaded, nil } func entryFromCost(providerID, providerName, modelID, modelName string, cost *sourceCost) (Entry, bool) { if strings.TrimSpace(modelID) == "" || cost.Input == nil || cost.Output == nil { return Entry{}, false } input, okInput := numberMicros(cost.Input) output, okOutput := numberMicros(cost.Output) if !okInput || !okOutput || !optionalMatches(cost.Reasoning, output) || !optionalMatches(cost.InputAudio, input) || !optionalMatches(cost.OutputAudio, output) { return Entry{}, false } cacheRead, ok := optionalMicros(cost.CacheRead, input) if !ok { return Entry{}, false } cacheWrite, ok := optionalMicros(cost.CacheWrite, input) if !ok || (input == 0 && output == 0 && cacheRead == 0 && cacheWrite == 0) { return Entry{}, false } entry := Entry{ ID: strings.ToLower(providerID + "/" + strings.TrimSpace(modelID)), Provider: providerID, ProviderName: providerName, Model: strings.TrimSpace(modelID), ModelName: strings.TrimSpace(modelName), Base: pricing.Rates{InputMicrosPer1M: input, CacheReadMicrosPer1M: cacheRead, CacheWriteMicrosPer1M: cacheWrite, OutputMicrosPer1M: output}, } if entry.ModelName == "" { entry.ModelName = entry.Model } if len(cost.Tiers) > 1 { return Entry{}, false } if len(cost.Tiers) == 1 { tier := cost.Tiers[0] if tier.Tier.Type != "context" || tier.Tier.Size <= 0 || tier.Input == nil || tier.Output == nil { return Entry{}, false } tierInput, okInput := numberMicros(tier.Input) tierOutput, okOutput := numberMicros(tier.Output) if !okInput || !okOutput || !optionalMatches(tier.Reasoning, tierOutput) || !optionalMatches(tier.InputAudio, tierInput) || !optionalMatches(tier.OutputAudio, tierOutput) { return Entry{}, false } tierCacheRead, okRead := optionalMicros(tier.CacheRead, tierInput) tierCacheWrite, okWrite := optionalMicros(tier.CacheWrite, tierInput) if !okRead || !okWrite { return Entry{}, false } entry.LongContext = &pricing.LongContext{ ThresholdInputTokens: tier.Tier.Size, Comparison: "gt", Rates: pricing.Rates{InputMicrosPer1M: tierInput, CacheReadMicrosPer1M: tierCacheRead, CacheWriteMicrosPer1M: tierCacheWrite, OutputMicrosPer1M: tierOutput}, } } return entry, true } func normalizedID(id, fallback string) string { id = strings.ToLower(strings.TrimSpace(id)) if id == "" { id = strings.ToLower(strings.TrimSpace(fallback)) } return id } func sameEntryPrice(left, right Entry) bool { if left.Base != right.Base { return false } if left.LongContext == nil || right.LongContext == nil { return left.LongContext == nil && right.LongContext == nil } return *left.LongContext == *right.LongContext } func optionalMicros(number *json.Number, fallback int64) (int64, bool) { if number == nil { return fallback, true } return numberMicros(number) } func optionalMatches(number *json.Number, expected int64) bool { if number == nil { return true } value, ok := numberMicros(number) return ok && value == expected } func numberMicros(number *json.Number) (int64, bool) { if number == nil { return 0, false } rational, ok := new(big.Rat).SetString(number.String()) if !ok || rational.Sign() < 0 { return 0, false } rational.Mul(rational, big.NewRat(1_000_000, 1)) numerator := rational.Num() denominator := rational.Denom() quotient, remainder := new(big.Int), new(big.Int) quotient.QuoRem(numerator, denominator, remainder) if new(big.Int).Lsh(remainder, 1).Cmp(denominator) >= 0 { quotient.Add(quotient, big.NewInt(1)) } return quotient.Int64(), quotient.IsInt64() } func cloneLongContext(value *pricing.LongContext) *pricing.LongContext { if value == nil { return nil } cloned := *value return &cloned } func cloneEntry(value Entry) Entry { value.LongContext = cloneLongContext(value.LongContext) return value }