Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions internal/core/interfaces.go
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,13 @@ type MessagesTokenCounter interface {
CountMessagesTokens(ctx context.Context, model string, body []byte) (int, error)
}

// UnlistedModelAcceptor is implemented by providers that serve model IDs
// their listing omits, such as TypeSafe's versioned Jev IDs (jev-1.13.0), so a
// virtual model can target a provider-qualified name the catalog lacks.
type UnlistedModelAcceptor interface {
AcceptsUnlistedModels() bool
}

// ErrMessagesTokenCountUnsupported reports that the provider owning a model
// has no token counting endpoint.
var ErrMessagesTokenCountUnsupported = errors.New("provider has no token counting endpoint")
Expand Down
3 changes: 3 additions & 0 deletions internal/gateway/failover.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,9 @@ func (o *InferenceOrchestrator) ProviderTypeForSelector(selector core.ModelSelec
if providerType := strings.TrimSpace(o.provider.GetProviderType(selector.QualifiedModel())); providerType != "" {
return providerType
}
if _, providerType := configuredSelectorProvider(o.provider, selector); providerType != "" {
return providerType
}
if provider := strings.TrimSpace(selector.Provider); provider != "" {
return provider
}
Expand Down
25 changes: 25 additions & 0 deletions internal/gateway/inference_orchestrator_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (

"github.com/enterpilot/gomodel/internal/core"
"github.com/enterpilot/gomodel/internal/usage"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -169,6 +170,30 @@ func TestInferenceOrchestratorProviderTypeForSelectorCanonicalizesProviderNameSe
require.Equal(t, "openai", got)
}

// namedProviderStub knows configured provider names and their types, but no
// catalog entry for the models under test.
type namedProviderStub struct {
providerTypeResolverStub
typesByName map[string]string
}

func (p *namedProviderStub) GetProviderTypeForName(name string) string { return p.typesByName[name] }

// A failover target the catalog does not list, such as a pinned jev version
// on a provider named kev, is routed to the provider its selector names, with
// that provider's type, never to the primary's provider.
func TestUnlistedSelectorResolvesToTheProviderItNames(t *testing.T) {
provider := &namedProviderStub{typesByName: map[string]string{"kev": "jev", "jev-down": "jev"}}
orchestrator := NewInferenceOrchestrator(InferenceConfig{Provider: provider})

pinned := core.ModelSelector{Provider: "kev", Model: "kev-4b-2026-09"}
assert.Equal(t, "jev", orchestrator.ProviderTypeForSelector(pinned, "openai"))
assert.Equal(t, "kev", ResolvedProviderName(provider, pinned, "jev-down"))

unknown := core.ModelSelector{Provider: "nope", Model: "x"}
assert.Equal(t, "jev-down", ResolvedProviderName(provider, unknown, "jev-down"))
}

func TestQualifyModelWithProviderPrefixesSlashModelIDs(t *testing.T) {
got := QualifyModelWithProvider("openai/gpt-4o-mini", "openrouter")
require.Equal(t, "openrouter/openai/gpt-4o-mini", got)
Expand Down
21 changes: 21 additions & 0 deletions internal/gateway/request_model_resolution.go
Original file line number Diff line number Diff line change
Expand Up @@ -30,9 +30,30 @@ func ResolvedProviderName(provider core.RoutableProvider, selector core.ModelSel
return providerName
}
}
// A model the catalog does not list (a pinned jev version) still belongs to
// the provider its selector names, not to the fallback's.
if providerName, providerType := configuredSelectorProvider(provider, selector); providerType != "" {
return providerName
}
return fallback
}

// configuredSelectorProvider returns the provider a selector names explicitly
// and that provider's type, or empty strings when the selector names no
// configured provider.
func configuredSelectorProvider(provider core.RoutableProvider, selector core.ModelSelector) (string, string) {
providerName := strings.TrimSpace(selector.Provider)
named, ok := provider.(core.ProviderNameTypeResolver)
if providerName == "" || !ok {
return "", ""
}
providerType := strings.TrimSpace(named.GetProviderTypeForName(providerName))
if providerType == "" {
return "", ""
}
return providerName, providerType
}

// ResolvedWorkflowProviderName returns the configured provider name recorded in a resolution.
func ResolvedWorkflowProviderName(resolution *core.RequestModelResolution) string {
if resolution == nil {
Expand Down
23 changes: 23 additions & 0 deletions internal/pricingoverrides/resolver.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,29 @@ func (s *Service) ResolvePricing(model, providerName string) *core.ModelPricing
return cloneBasePricing(basePricing)
}

// HasModelPricing reports whether pricing is declared for exactly this model:
// catalog pricing, or an override scoped to the model rather than to its
// provider or to every model.
func (s *Service) HasModelPricing(model, providerName string) bool {
if s == nil {
return false
}
providerName = strings.TrimSpace(providerName)
rawModel := strings.TrimSpace(model)
model = modelIDFromSelector(rawModel, providerName)
if model == "" {
return false
}
if s.snapshot().hasModelScopedOverride(providerName, model) {
return true
}
if s.base == nil {
return false
}
return s.base.ResolvePricing(model, providerName) != nil ||
(rawModel != model && s.base.ResolvePricing(rawModel, providerName) != nil)
}

func cloneBasePricing(base *core.ModelPricing) *core.ModelPricing {
if base == nil {
return nil
Expand Down
33 changes: 33 additions & 0 deletions internal/pricingoverrides/service_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"time"

"github.com/enterpilot/gomodel/internal/core"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

Expand Down Expand Up @@ -353,3 +354,35 @@ func TestNormalizedRefreshIntervalClampsBelowRefreshTimeout(t *testing.T) {
})
}
}

func TestServiceHasModelPricing(t *testing.T) {
baseRate := 1.0
service, err := NewService(
newTestStore(
Override{Selector: "/", Pricing: Pricing{InputPerMtok: new(float64(10))}},
Override{Selector: "jev/", Pricing: Pricing{InputPerMtok: new(float64(20))}},
Override{Selector: "jev/jev-1.13.0", Pricing: Pricing{InputPerMtok: new(float64(42))}},
Override{Selector: "kev-4b", Pricing: Pricing{InputPerMtok: new(float64(0))}},
),
testCatalog{providerNames: []string{"jev", "openai"}},
selectivePricingResolver{"openai/gpt-4o": {InputPerMtok: &baseRate}},
)
require.NoError(t, err)
require.NoError(t, service.Refresh(context.Background()))

tests := []struct {
model, provider string
want bool
}{
{model: "jev-1.13.0", provider: "jev", want: true},
{model: "jev/jev-1.13.0", provider: "jev", want: true},
{model: "kev-4b", provider: "kev", want: true},
{model: "gpt-4o", provider: "openai", want: true},
// Only the provider-wide and global overrides match these.
{model: "jev-latest", provider: "jev", want: false},
{model: "gpt-9", provider: "openai", want: false},
}
for _, tt := range tests {
assert.Equal(t, tt.want, service.HasModelPricing(tt.model, tt.provider), "%s/%s", tt.provider, tt.model)
}
}
12 changes: 12 additions & 0 deletions internal/pricingoverrides/snapshot.go
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,18 @@ func (snap snapshot) matchingOverride(providerName, model string) (compiledOverr
return compiledOverride{}, false
}

// hasModelScopedOverride reports whether an override names this model,
// either for one provider or model-wide.
func (snap snapshot) hasModelScopedOverride(providerName, model string) bool {
if key := modelselectors.ExactMatchKey(providerName, model); key != "" {
if _, ok := snap.exact[key]; ok {
return true
}
}
_, ok := snap.modelWide[model]
return ok
}

func snapshotOverrides(snap snapshot) []Override {
result := make([]Override, 0, len(snap.order))
for _, selector := range snap.order {
Expand Down
9 changes: 7 additions & 2 deletions internal/providers/jev/jev.go
Original file line number Diff line number Diff line change
Expand Up @@ -42,8 +42,9 @@ type Provider struct {
}

var (
_ core.Provider = (*Provider)(nil)
_ core.PassthroughProvider = (*Provider)(nil)
_ core.Provider = (*Provider)(nil)
_ core.PassthroughProvider = (*Provider)(nil)
_ core.UnlistedModelAcceptor = (*Provider)(nil)
)

// New creates a Jev provider. The client is rooted at the API origin, which
Expand All @@ -64,6 +65,10 @@ func New(cfg providers.ProviderConfig, opts providers.ProviderOptions) core.Prov
return p
}

// AcceptsUnlistedModels reports that the upstream accepts versioned IDs
// (jev-1.13.0) it does not list: TypeSafe lists only its aliases.
func (p *Provider) AcceptsUnlistedModels() bool { return true }

// SetBaseURL allows configuring a custom base URL for the provider.
func (p *Provider) SetBaseURL(url string) {
p.client.SetBaseURL(baseURL(url))
Expand Down
22 changes: 22 additions & 0 deletions internal/providers/registry_lookup.go
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,28 @@ func (r *ModelRegistry) ModelAvailable(model string) bool {
return !r.providerRuntime[info.ProviderName].inventoryStale
}

// AcceptsUnlistedModel reports whether a provider-qualified model the catalog
// does not list can still be served, because its provider accepts IDs it does
// not list (see core.UnlistedModelAcceptor) and its inventory is fresh. A bare
// name never qualifies: it does not say which provider to use.
func (r *ModelRegistry) AcceptsUnlistedModel(model string) bool {
providerName, _ := splitModelSelector(strings.TrimSpace(model))
if providerName == "" {
return false
}
r.mu.RLock()
defer r.mu.RUnlock()

for _, provider := range r.providers {
if r.providerNames[provider] != providerName {
continue
}
acceptor, ok := provider.(core.UnlistedModelAcceptor)
return ok && acceptor.AcceptsUnlistedModels() && !r.providerRuntime[providerName].inventoryStale
}
return false
}

// GetProviderType returns the provider type string for the given model.
// Returns empty string if the model is not found.
func (r *ModelRegistry) GetProviderType(model string) string {
Expand Down
27 changes: 27 additions & 0 deletions internal/providers/registry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2329,3 +2329,30 @@ func TestSetModelList_ClearsETag(t *testing.T) {
got := registry.currentModelListETag("https://example.test/models.min.json")
require.Empty(t, got)
}

// unlistedAcceptingProvider serves model IDs it does not list, as a jev
// provider serves pinned versions.
type unlistedAcceptingProvider struct {
registryMockProvider
}

func (p *unlistedAcceptingProvider) AcceptsUnlistedModels() bool { return true }

func TestModelRegistryAcceptsUnlistedModel(t *testing.T) {
registry := NewModelRegistry()
registry.RegisterProviderWithNameAndType(&unlistedAcceptingProvider{}, "jev", "jev")
registry.RegisterProviderWithNameAndType(&registryMockProvider{name: "openai"}, "openai", "openai")

tests := []struct {
model string
want bool
}{
{model: "jev/jev-1.13.0", want: true},
{model: "jev-1.13.0", want: false},
{model: "openai/gpt-9", want: false},
{model: "unknown/jev-1.13.0", want: false},
}
for _, tt := range tests {
assert.Equal(t, tt.want, registry.AcceptsUnlistedModel(tt.model), tt.model)
}
}
8 changes: 8 additions & 0 deletions internal/responsecache/responsecache.go
Original file line number Diff line number Diff line change
Expand Up @@ -304,3 +304,11 @@ func NewResponseCacheMiddlewareWithStore(store cache.Store, ttl time.Duration) *
simple: newSimpleCacheMiddleware(store, ttl, nil),
}
}

// NewResponseCacheMiddlewareWithStoreAndUsage creates middleware with a custom
// store that records cache hits in usage (for testing).
func NewResponseCacheMiddlewareWithStoreAndUsage(store cache.Store, ttl time.Duration, usageLogger usage.LoggerInterface, pricingResolver usage.PricingResolver) *ResponseCacheMiddleware {
return &ResponseCacheMiddleware{
simple: newSimpleCacheMiddleware(store, ttl, newUsageHitRecorder(usageLogger, pricingResolver)),
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}
}
20 changes: 16 additions & 4 deletions internal/responsecache/usage_hit.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import (
"log/slog"
"strings"

"github.com/goccy/go-json"

"github.com/enterpilot/gomodel/internal/core"
"github.com/enterpilot/gomodel/internal/usage"
)
Expand Down Expand Up @@ -45,10 +47,8 @@ func newUsageHitRecorder(logger usage.LoggerInterface, pricingResolver usage.Pri
requestID = ex.RequestHeader(core.RequestIDHeader)
}

var pricing *core.ModelPricing
if pricingResolver != nil {
pricing = pricingResolver.ResolvePricing(model, cacheHitPricingProvider(provider, providerName))
}
pricing := usage.ResolveServedModelPricing(pricingResolver, model, cacheHitPricingProvider(provider, providerName),
func() string { return cachedAnsweredModel(body) })

entry := usage.ExtractFromCachedResponseBody(body, requestID, model, provider, endpoint, cacheType, pricing)
if entry == nil {
Expand All @@ -62,6 +62,18 @@ func newUsageHitRecorder(logger usage.LoggerInterface, pricingResolver usage.Pri
}
}

// cachedAnsweredModel returns the model a cached JSON answer names, such as
// jev-1.13.0 for a request routed to jev-latest, or "" for other bodies.
func cachedAnsweredModel(body []byte) string {
var answer struct {
Model string `json:"model"`
}
if err := json.Unmarshal(body, &answer); err != nil {
return ""
}
return answer.Model
}

func cacheHitPricingProvider(provider, providerName string) string {
if name := strings.TrimSpace(providerName); name != "" {
return name
Expand Down
42 changes: 42 additions & 0 deletions internal/server/systemone_dispatch_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -121,6 +121,48 @@ func TestSystemOne_ServesRepeatsFromTheExactCache(t *testing.T) {
assert.Len(t, provider.calls, 1, "the repeat must not reach the provider")
}

// answeredModelPricing prices only the models it names, as an operator who
// declares a price for the versioned model that answers an alias.
type answeredModelPricing map[string]*core.ModelPricing

func (r answeredModelPricing) ResolvePricing(model, _ string) *core.ModelPricing { return r[model] }

// A cache hit is recorded in usage as an exact hit, with the tokens of the
// replayed answer, priced like the live answer: by the model that answered
// when only it carries a price.
func TestSystemOne_RecordsCacheHitsInUsage(t *testing.T) {
provider := newScriptedSystemOneProvider(map[string]string{"kev/kev-latest": "jev"})
store := cache.NewMapStore()
defer store.Close()
rate := 1_000_000.0
pricing := answeredModelPricing{"kev-latest-answered": {InputPerMtok: &rate}}
usageLogger := &collectingUsageLogger{config: usage.Config{Enabled: true}}
mw := responsecache.NewResponseCacheMiddlewareWithStoreAndUsage(store, time.Hour, usageLogger, pricing)
handler := newHandler(provider, nil, usageLogger, pricing, nil, nil, nil, nil)
handler.responseCache = mw

c, first := echotest.Post(t, "/v1/systemone", systemOneRequest("kev-latest"))
require.NoError(t, handler.SystemOne(c))
require.Equal(t, http.StatusOK, first.Code, first.Body.String())
require.NoError(t, mw.Close())

c, second := echotest.Post(t, "/v1/systemone", systemOneRequest("kev-latest"))
require.NoError(t, handler.SystemOne(c))
require.Equal(t, "HIT (exact)", second.Header().Get("X-Cache"))

require.Len(t, usageLogger.entries, 2)
hit := usageLogger.entries[1]
assert.Equal(t, usage.CacheTypeExact, hit.CacheType)
assert.Equal(t, "/v1/systemone", hit.Endpoint)
assert.Equal(t, "jev", hit.Provider)
assert.Equal(t, 10, hit.InputTokens)
assert.Equal(t, 1, hit.OutputTokens)
for _, entry := range usageLogger.entries {
require.NotNil(t, entry.InputCost, "cache type %q", entry.CacheType)
assert.InDelta(t, 10.0, *entry.InputCost, 1e-9)
}
}

// A different state is a different decision: it misses the cache.
func TestSystemOne_CacheKeyCoversTheState(t *testing.T) {
provider := newScriptedSystemOneProvider(map[string]string{"kev/kev-latest": "jev"})
Expand Down
5 changes: 5 additions & 0 deletions internal/usage/extractor.go
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,11 @@ func ExtractFromSSEUsage(
requestID, model, provider, endpoint string,
pricing ...*core.ModelPricing,
) *UsageEntry {
// Anthropic-style usage (and System One answers) report input and output
// tokens without a total.
if totalTokens == 0 {
totalTokens = inputTokens + outputTokens
}
entry := &UsageEntry{
ID: uuid.New().String(),
RequestID: requestID,
Expand Down
Loading
Loading