Skip to content

Commit eeff451

Browse files
committed
test(backend): add tests for upstream model tracking and model source filtering
Cover IsValidModelSource/NormalizeModelSource, resolveModelDimensionExpression SQL expressions, invalid model_source 400 responses on both GetModelStats and GetUserBreakdown, upstream_model in scan/insert SQL mock expectations, and updated passthrough/billing test signatures.
1 parent 56fcb20 commit eeff451

7 files changed

Lines changed: 132 additions & 8 deletions

backend/internal/handler/admin/dashboard_handler_request_type_test.go

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -149,6 +149,28 @@ func TestDashboardModelStatsInvalidStream(t *testing.T) {
149149
require.Equal(t, http.StatusBadRequest, rec.Code)
150150
}
151151

152+
func TestDashboardModelStatsInvalidModelSource(t *testing.T) {
153+
repo := &dashboardUsageRepoCapture{}
154+
router := newDashboardRequestTypeTestRouter(repo)
155+
156+
req := httptest.NewRequest(http.MethodGet, "/admin/dashboard/models?model_source=invalid", nil)
157+
rec := httptest.NewRecorder()
158+
router.ServeHTTP(rec, req)
159+
160+
require.Equal(t, http.StatusBadRequest, rec.Code)
161+
}
162+
163+
func TestDashboardModelStatsValidModelSource(t *testing.T) {
164+
repo := &dashboardUsageRepoCapture{}
165+
router := newDashboardRequestTypeTestRouter(repo)
166+
167+
req := httptest.NewRequest(http.MethodGet, "/admin/dashboard/models?model_source=upstream", nil)
168+
rec := httptest.NewRecorder()
169+
router.ServeHTTP(rec, req)
170+
171+
require.Equal(t, http.StatusOK, rec.Code)
172+
}
173+
152174
func TestDashboardUsersRankingLimitAndCache(t *testing.T) {
153175
dashboardUsersRankingCache = newSnapshotCache(5 * time.Minute)
154176
repo := &dashboardUsageRepoCapture{

backend/internal/handler/admin/dashboard_handler_user_breakdown_test.go

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -73,9 +73,35 @@ func TestGetUserBreakdown_ModelFilter(t *testing.T) {
7373

7474
require.Equal(t, http.StatusOK, w.Code)
7575
require.Equal(t, "claude-opus-4-6", repo.capturedDim.Model)
76+
require.Equal(t, usagestats.ModelSourceRequested, repo.capturedDim.ModelType)
7677
require.Equal(t, int64(0), repo.capturedDim.GroupID)
7778
}
7879

80+
func TestGetUserBreakdown_ModelSourceFilter(t *testing.T) {
81+
repo := &userBreakdownRepoCapture{}
82+
router := newUserBreakdownRouter(repo)
83+
84+
req := httptest.NewRequest(http.MethodGet,
85+
"/admin/dashboard/user-breakdown?start_date=2026-03-01&end_date=2026-03-16&model=claude-opus-4-6&model_source=upstream", nil)
86+
w := httptest.NewRecorder()
87+
router.ServeHTTP(w, req)
88+
89+
require.Equal(t, http.StatusOK, w.Code)
90+
require.Equal(t, usagestats.ModelSourceUpstream, repo.capturedDim.ModelType)
91+
}
92+
93+
func TestGetUserBreakdown_InvalidModelSource(t *testing.T) {
94+
repo := &userBreakdownRepoCapture{}
95+
router := newUserBreakdownRouter(repo)
96+
97+
req := httptest.NewRequest(http.MethodGet,
98+
"/admin/dashboard/user-breakdown?start_date=2026-03-01&end_date=2026-03-16&model_source=foobar", nil)
99+
w := httptest.NewRecorder()
100+
router.ServeHTTP(w, req)
101+
102+
require.Equal(t, http.StatusBadRequest, w.Code)
103+
}
104+
79105
func TestGetUserBreakdown_EndpointFilter(t *testing.T) {
80106
repo := &userBreakdownRepoCapture{}
81107
router := newUserBreakdownRouter(repo)
Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,47 @@
1+
package usagestats
2+
3+
import "testing"
4+
5+
func TestIsValidModelSource(t *testing.T) {
6+
tests := []struct {
7+
name string
8+
source string
9+
want bool
10+
}{
11+
{name: "requested", source: ModelSourceRequested, want: true},
12+
{name: "upstream", source: ModelSourceUpstream, want: true},
13+
{name: "mapping", source: ModelSourceMapping, want: true},
14+
{name: "invalid", source: "foobar", want: false},
15+
{name: "empty", source: "", want: false},
16+
}
17+
18+
for _, tc := range tests {
19+
t.Run(tc.name, func(t *testing.T) {
20+
if got := IsValidModelSource(tc.source); got != tc.want {
21+
t.Fatalf("IsValidModelSource(%q)=%v want %v", tc.source, got, tc.want)
22+
}
23+
})
24+
}
25+
}
26+
27+
func TestNormalizeModelSource(t *testing.T) {
28+
tests := []struct {
29+
name string
30+
source string
31+
want string
32+
}{
33+
{name: "requested", source: ModelSourceRequested, want: ModelSourceRequested},
34+
{name: "upstream", source: ModelSourceUpstream, want: ModelSourceUpstream},
35+
{name: "mapping", source: ModelSourceMapping, want: ModelSourceMapping},
36+
{name: "invalid falls back", source: "foobar", want: ModelSourceRequested},
37+
{name: "empty falls back", source: "", want: ModelSourceRequested},
38+
}
39+
40+
for _, tc := range tests {
41+
t.Run(tc.name, func(t *testing.T) {
42+
if got := NormalizeModelSource(tc.source); got != tc.want {
43+
t.Fatalf("NormalizeModelSource(%q)=%q want %q", tc.source, got, tc.want)
44+
}
45+
})
46+
}
47+
}

backend/internal/repository/usage_log_repo_breakdown_test.go

Lines changed: 23 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ package repository
55
import (
66
"testing"
77

8+
"github.com/Wei-Shaw/sub2api/internal/pkg/usagestats"
89
"github.com/stretchr/testify/require"
910
)
1011

@@ -16,8 +17,8 @@ func TestResolveEndpointColumn(t *testing.T) {
1617
{"inbound", "ul.inbound_endpoint"},
1718
{"upstream", "ul.upstream_endpoint"},
1819
{"path", "ul.inbound_endpoint || ' -> ' || ul.upstream_endpoint"},
19-
{"", "ul.inbound_endpoint"}, // default
20-
{"unknown", "ul.inbound_endpoint"}, // fallback
20+
{"", "ul.inbound_endpoint"}, // default
21+
{"unknown", "ul.inbound_endpoint"}, // fallback
2122
}
2223

2324
for _, tc := range tests {
@@ -27,3 +28,23 @@ func TestResolveEndpointColumn(t *testing.T) {
2728
})
2829
}
2930
}
31+
32+
func TestResolveModelDimensionExpression(t *testing.T) {
33+
tests := []struct {
34+
modelType string
35+
want string
36+
}{
37+
{usagestats.ModelSourceRequested, "model"},
38+
{usagestats.ModelSourceUpstream, "COALESCE(NULLIF(TRIM(upstream_model), ''), model)"},
39+
{usagestats.ModelSourceMapping, "(model || ' -> ' || COALESCE(NULLIF(TRIM(upstream_model), ''), model))"},
40+
{"", "model"},
41+
{"invalid", "model"},
42+
}
43+
44+
for _, tc := range tests {
45+
t.Run(tc.modelType, func(t *testing.T) {
46+
got := resolveModelDimensionExpression(tc.modelType)
47+
require.Equal(t, tc.want, got)
48+
})
49+
}
50+
}

backend/internal/repository/usage_log_repo_request_type_test.go

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) {
4444
log.AccountID,
4545
log.RequestID,
4646
log.Model,
47+
sqlmock.AnyArg(), // upstream_model
4748
sqlmock.AnyArg(), // group_id
4849
sqlmock.AnyArg(), // subscription_id
4950
log.InputTokens,
@@ -116,6 +117,7 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) {
116117
log.Model,
117118
sqlmock.AnyArg(),
118119
sqlmock.AnyArg(),
120+
sqlmock.AnyArg(),
119121
log.InputTokens,
120122
log.OutputTokens,
121123
log.CacheCreationTokens,
@@ -353,6 +355,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
353355
int64(30), // account_id
354356
sql.NullString{Valid: true, String: "req-1"},
355357
"gpt-5", // model
358+
sql.NullString{}, // upstream_model
356359
sql.NullInt64{}, // group_id
357360
sql.NullInt64{}, // subscription_id
358361
1, // input_tokens
@@ -404,6 +407,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
404407
int64(31),
405408
sql.NullString{Valid: true, String: "req-2"},
406409
"gpt-5",
410+
sql.NullString{},
407411
sql.NullInt64{},
408412
sql.NullInt64{},
409413
1, 2, 3, 4, 5, 6,
@@ -445,6 +449,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
445449
int64(32),
446450
sql.NullString{Valid: true, String: "req-3"},
447451
"gpt-5.4",
452+
sql.NullString{},
448453
sql.NullInt64{},
449454
sql.NullInt64{},
450455
1, 2, 3, 4, 5, 6,

backend/internal/service/gateway_anthropic_apikey_passthrough_test.go

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -788,7 +788,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardDirect_NonStreamingSuc
788788
rateLimitService: &RateLimitService{},
789789
}
790790

791-
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, newAnthropicAPIKeyAccountForTest(), body, "claude-3-5-sonnet-latest", false, time.Now())
791+
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, newAnthropicAPIKeyAccountForTest(), body, "claude-3-5-sonnet-latest", "claude-3-5-sonnet-latest", false, time.Now())
792792
require.NoError(t, err)
793793
require.NotNil(t, result)
794794
require.Equal(t, 12, result.Usage.InputTokens)
@@ -815,7 +815,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardDirect_InvalidTokenTyp
815815
}
816816
svc := &GatewayService{}
817817

818-
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "claude-3-5-sonnet-latest", false, time.Now())
818+
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{}`), "claude-3-5-sonnet-latest", "claude-3-5-sonnet-latest", false, time.Now())
819819
require.Nil(t, result)
820820
require.Error(t, err)
821821
require.Contains(t, err.Error(), "requires apikey token")
@@ -840,7 +840,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardDirect_UpstreamRequest
840840
}
841841
account := newAnthropicAPIKeyAccountForTest()
842842

843-
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{"model":"x"}`), "x", false, time.Now())
843+
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, account, []byte(`{"model":"x"}`), "x", "x", false, time.Now())
844844
require.Nil(t, result)
845845
require.Error(t, err)
846846
require.Contains(t, err.Error(), "upstream request failed")
@@ -873,7 +873,7 @@ func TestGatewayService_AnthropicAPIKeyPassthrough_ForwardDirect_EmptyResponseBo
873873
httpUpstream: upstream,
874874
}
875875

876-
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, newAnthropicAPIKeyAccountForTest(), []byte(`{"model":"x"}`), "x", false, time.Now())
876+
result, err := svc.forwardAnthropicAPIKeyPassthrough(context.Background(), c, newAnthropicAPIKeyAccountForTest(), []byte(`{"model":"x"}`), "x", "x", false, time.Now())
877877
require.Nil(t, result)
878878
require.Error(t, err)
879879
require.Contains(t, err.Error(), "empty response")

backend/internal/service/openai_gateway_record_usage_test.go

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -846,7 +846,7 @@ func TestExtractOpenAIServiceTierFromBody(t *testing.T) {
846846
require.Nil(t, extractOpenAIServiceTierFromBody(nil))
847847
}
848848

849-
func TestOpenAIGatewayServiceRecordUsage_UsesBillingModelAndMetadataFields(t *testing.T) {
849+
func TestOpenAIGatewayServiceRecordUsage_UsesRequestedModelAndUpstreamModelMetadataFields(t *testing.T) {
850850
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
851851
userRepo := &openAIRecordUsageUserRepoStub{}
852852
subRepo := &openAIRecordUsageSubRepoStub{}
@@ -859,6 +859,7 @@ func TestOpenAIGatewayServiceRecordUsage_UsesBillingModelAndMetadataFields(t *te
859859
RequestID: "resp_billing_model_override",
860860
BillingModel: "gpt-5.1-codex",
861861
Model: "gpt-5.1",
862+
UpstreamModel: "gpt-5.1-codex",
862863
ServiceTier: &serviceTier,
863864
ReasoningEffort: &reasoning,
864865
Usage: OpenAIUsage{
@@ -877,7 +878,9 @@ func TestOpenAIGatewayServiceRecordUsage_UsesBillingModelAndMetadataFields(t *te
877878

878879
require.NoError(t, err)
879880
require.NotNil(t, usageRepo.lastLog)
880-
require.Equal(t, "gpt-5.1-codex", usageRepo.lastLog.Model)
881+
require.Equal(t, "gpt-5.1", usageRepo.lastLog.Model)
882+
require.NotNil(t, usageRepo.lastLog.UpstreamModel)
883+
require.Equal(t, "gpt-5.1-codex", *usageRepo.lastLog.UpstreamModel)
881884
require.NotNil(t, usageRepo.lastLog.ServiceTier)
882885
require.Equal(t, serviceTier, *usageRepo.lastLog.ServiceTier)
883886
require.NotNil(t, usageRepo.lastLog.ReasoningEffort)

0 commit comments

Comments
 (0)