diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 816e7e53a..cb625f777 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -33,10 +33,9 @@ repos: language: system pass_filenames: false files: '^(internal/(auditlog|core|providers|server|usage)/.*\.go|tests/perf/.*\.go|Makefile|\.github/workflows/test\.yml)$' - - # 3. High-Performance Go Linter (golangci-lint) - - repo: https://github.com/golangci/golangci-lint - rev: v2.7.2 # Latest stable version (Dec 2025) - hooks: - - id: golangci-lint - args: [--fix=false, --timeout=5m] # Optional: increase timeout if project is large + - id: go-lint + name: make lint + entry: make lint + language: system + pass_filenames: false + files: '^(\.golangci\.yml|Makefile|\.github/workflows/test\.yml|.*\.go)$' diff --git a/Makefile b/Makefile index e44468828..5dfc64d84 100644 --- a/Makefile +++ b/Makefile @@ -83,11 +83,11 @@ swagger: # Run linter lint: - golangci-lint run ./... + golangci-lint run ./cmd/... ./config/... ./internal/... golangci-lint run --build-tags=e2e ./tests/e2e/... golangci-lint run --build-tags=integration ./tests/integration/... golangci-lint run --build-tags=contract ./tests/contract/... # Run linter with auto-fix lint-fix: - golangci-lint run --fix ./... + golangci-lint run --fix ./cmd/... ./config/... ./internal/... diff --git a/internal/admin/dashboard/dashboard.go b/internal/admin/dashboard/dashboard.go index 11826e752..79d0ada31 100644 --- a/internal/admin/dashboard/dashboard.go +++ b/internal/admin/dashboard/dashboard.go @@ -3,10 +3,13 @@ package dashboard import ( "bytes" + "crypto/sha256" "embed" + "encoding/hex" "html/template" "io/fs" "net/http" + "strings" "github.com/labstack/echo/v5" ) @@ -22,7 +25,16 @@ type Handler struct { // New creates a new dashboard handler with parsed templates and static file server. func New() (*Handler, error) { - tmpl, err := template.ParseFS(content, "templates/*.html") + assetVersions, err := buildAssetVersions("css/dashboard.css") + if err != nil { + return nil, err + } + + tmpl, err := template.New("layout").Funcs(template.FuncMap{ + "assetURL": func(path string) string { + return assetURL(path, assetVersions) + }, + }).ParseFS(content, "templates/*.html") if err != nil { return nil, err } @@ -55,3 +67,32 @@ func (h *Handler) Static(c *echo.Context) error { h.staticFS.ServeHTTP(c.Response(), c.Request()) return nil } + +func buildAssetVersions(paths ...string) (map[string]string, error) { + versions := make(map[string]string, len(paths)) + for _, path := range paths { + normalizedPath := strings.TrimLeft(strings.TrimSpace(path), "/") + if normalizedPath == "" { + continue + } + data, err := content.ReadFile("static/" + normalizedPath) + if err != nil { + return nil, err + } + sum := sha256.Sum256(data) + versions[normalizedPath] = hex.EncodeToString(sum[:6]) + } + return versions, nil +} + +func assetURL(path string, versions map[string]string) string { + normalizedPath := strings.TrimLeft(strings.TrimSpace(path), "/") + if normalizedPath == "" { + return "/admin/static/" + } + urlPath := "/admin/static/" + normalizedPath + if version := versions[normalizedPath]; version != "" { + return urlPath + "?v=" + version + } + return urlPath +} diff --git a/internal/admin/dashboard/dashboard_test.go b/internal/admin/dashboard/dashboard_test.go index c8ccfa0c6..c8c36f98b 100644 --- a/internal/admin/dashboard/dashboard_test.go +++ b/internal/admin/dashboard/dashboard_test.go @@ -3,6 +3,7 @@ package dashboard import ( "net/http" "net/http/httptest" + "regexp" "strings" "testing" @@ -59,6 +60,9 @@ func TestIndex_ReturnsHTML(t *testing.T) { if strings.Contains(body, `x-init="init()"`) { t.Errorf("expected dashboard HTML not to call init() explicitly") } + if !regexp.MustCompile(`/admin/static/css/dashboard\.css\?v=[0-9a-f]+`).MatchString(rec.Body.String()) { + t.Errorf("expected versioned dashboard CSS link in page HTML") + } } func TestStatic_ServesCSS(t *testing.T) { diff --git a/internal/admin/dashboard/static/css/dashboard.css b/internal/admin/dashboard/static/css/dashboard.css index 7ae708030..e90102eb5 100644 --- a/internal/admin/dashboard/static/css/dashboard.css +++ b/internal/admin/dashboard/static/css/dashboard.css @@ -835,6 +835,10 @@ body { border-color: var(--accent); } +.models-filter-input { + max-width: 840px; +} + .table-wrapper { background: var(--bg-surface); border: 1px solid var(--border); @@ -875,6 +879,9 @@ body { .mono { font-family: 'SF Mono', Menlo, Consolas, monospace; +} + +.font-size-md { font-size: 13px; } @@ -942,6 +949,21 @@ td.col-price { border-color: color-mix(in srgb, var(--danger) 50%, var(--border)); } +.table-icon-btn { + width: 32px; + min-width: 32px; + height: 32px; + padding: 0; + gap: 0; + border-radius: 999px; +} + +.table-icon-svg { + width: 14px; + height: 14px; + flex-shrink: 0; +} + .model-alias-editor { background: var(--bg-surface); border: 1px solid var(--border); @@ -1487,27 +1509,44 @@ td.col-price { gap: 10px; } -.audit-filter-row .filter-input:first-child { - grid-column: span 4; +.audit-filter-select { + grid-column: span 2; + min-width: 0; } -.audit-filter-input { - grid-column: span 2; +.audit-filter-row-search .filter-input { + grid-column: 1 / -1; max-width: none; } -.audit-filter-select { +.audit-filter-row-controls .audit-filter-select { grid-column: span 2; - min-width: 0; } -.audit-filter-row .pagination-btn { - grid-column: span 1; - min-width: 84px; +.audit-filter-row-controls .pagination-btn { + grid-column: 11 / -1; + justify-self: end; + min-width: 108px; } -.audit-filter-row-selects .audit-filter-select { - grid-column: span 2; +.audit-clear-btn { + display: inline-flex; + align-items: center; + justify-content: center; + gap: 8px; + background: #fff; + border-color: color-mix(in srgb, #fff 80%, var(--border)); + color: #111110; + font-weight: 600; +} + +.audit-clear-btn:hover:not(:disabled) { + background: #f5f5f5; +} + +.audit-clear-btn .table-icon-svg { + width: 12px; + height: 12px; } .audit-log-summary { @@ -2661,6 +2700,10 @@ body.conversation-drawer-open { .aliases-editor-header .table-action-btn { width: 100%; } + .alias-actions-cell .table-icon-btn, + .aliases-editor-header .table-icon-btn { + width: 36px; + } /* Audit page mobile */ .audit-log-toolbar { @@ -2669,8 +2712,7 @@ body.conversation-drawer-open { .audit-filter-row { grid-template-columns: 1fr; } - .audit-filter-row .filter-input:first-child, - .audit-filter-input, + .audit-filter-row .filter-input, .audit-filter-select, .audit-filter-row .pagination-btn { grid-column: auto; @@ -2763,27 +2805,132 @@ body.conversation-drawer-open { position: absolute; top: 12px; right: 14px; - display: flex; - align-items: center; - gap: 6px; - flex-wrap: wrap; - justify-content: flex-end; - max-width: calc(100% - 28px); -} - -.exec-pipeline-meta-chip { display: inline-flex; align-items: center; - padding: 3px 8px; - border-radius: 999px; + gap: 0; + min-width: 0; + max-width: calc(100% - 28px); + padding: 2px 10px; + border-radius: 12px; border: 1px solid var(--border); background: color-mix(in srgb, var(--bg-surface) 86%, transparent); color: var(--text-muted); - font-size: 10px; + font-size: 12px; + font-weight: 500; line-height: 1.2; white-space: nowrap; } +.exec-pipeline-meta-copy { + appearance: none; + cursor: pointer; + text-align: left; + overflow: hidden; + transition: background-color 0.15s, border-color 0.15s, color 0.15s, box-shadow 0.15s; +} + +.exec-pipeline-meta-copy:hover, +.exec-pipeline-meta-copy:focus-visible { + border-color: color-mix(in srgb, var(--accent) 40%, var(--border)); + background: color-mix(in srgb, var(--accent) 8%, var(--bg-surface)); + color: color-mix(in srgb, var(--accent) 74%, var(--text)); +} + +.exec-pipeline-meta-copy:focus-visible { + outline: none; + box-shadow: 0 0 0 2px color-mix(in srgb, var(--accent) 18%, transparent); +} + +.exec-pipeline-meta-label { + flex: 0 0 auto; + font-weight: 700; +} + +.exec-pipeline-meta-placeholder { + flex: 0 0 auto; + max-width: 3ch; + margin-left: 4px; + overflow: hidden; + opacity: 1; + transition: max-width 0.18s ease, margin-left 0.18s ease, opacity 0.15s ease; +} + +.exec-pipeline-meta-value { + flex: 0 1 auto; + max-width: 0; + margin-left: 0; + overflow: hidden; + opacity: 0; + text-overflow: clip; + transition: max-width 0.22s ease, margin-left 0.18s ease, opacity 0.15s ease; +} + +.exec-pipeline-meta-copy:hover .exec-pipeline-meta-placeholder, +.exec-pipeline-meta-copy:focus-visible .exec-pipeline-meta-placeholder, +.exec-pipeline-meta-copied .exec-pipeline-meta-placeholder, +.exec-pipeline-meta-error .exec-pipeline-meta-placeholder { + max-width: 0; + margin-left: 0; + opacity: 0; +} + +.exec-pipeline-meta-copy:hover .exec-pipeline-meta-value, +.exec-pipeline-meta-copy:focus-visible .exec-pipeline-meta-value, +.exec-pipeline-meta-copied .exec-pipeline-meta-value, +.exec-pipeline-meta-error .exec-pipeline-meta-value { + max-width: 42ch; + margin-left: 4px; + opacity: 1; +} + +.exec-pipeline-meta-icon { + display: inline-flex; + align-items: center; + justify-content: center; + flex: 0 0 auto; + width: 0; + height: 14px; + margin-left: 0; + overflow: hidden; + opacity: 0; + line-height: 0; + transform: translateX(4px) translateY(1px) scale(0.84); + transition: width 0.18s ease, margin-left 0.18s ease, opacity 0.15s ease, transform 0.18s ease; +} + +.exec-pipeline-meta-icon svg { + width: 14px; + height: 14px; + stroke: currentcolor; + fill: none; + stroke-width: 2; + stroke-linecap: round; + stroke-linejoin: round; +} + +.exec-pipeline-meta-copied, +.exec-pipeline-meta-copied:hover, +.exec-pipeline-meta-copied:focus-visible { + background: color-mix(in srgb, var(--success) 12%, var(--bg)); + border-color: color-mix(in srgb, var(--success) 40%, var(--border)); + color: var(--success); +} + +.exec-pipeline-meta-copied .exec-pipeline-meta-icon { + width: 14px; + margin-left: 6px; + opacity: 1; + transform: translateY(1px); +} + +.exec-pipeline-meta-error, +.exec-pipeline-meta-error:hover, +.exec-pipeline-meta-error:focus-visible { + background: color-mix(in srgb, var(--danger) 10%, var(--bg)); + border-color: color-mix(in srgb, var(--danger) 34%, var(--border)); + color: var(--danger); +} + /* ─── Main pipeline row ─── */ .exec-pipeline-row { diff --git a/internal/admin/dashboard/static/js/dashboard.js b/internal/admin/dashboard/static/js/dashboard.js index 35da1bd82..2d69711c3 100644 --- a/internal/admin/dashboard/static/js/dashboard.js +++ b/internal/admin/dashboard/static/js/dashboard.js @@ -83,11 +83,7 @@ function dashboard() { // Audit page state auditLog: { entries: [], total: 0, limit: 25, offset: 0 }, auditSearch: '', - auditModel: '', - auditProvider: '', auditMethod: '', - auditPath: '', - auditUserPath: '', auditStatusCode: '', auditStream: '', auditFetchToken: 0, @@ -407,13 +403,43 @@ function dashboard() { const f = this.modelFilter.toLowerCase(); return this.models.filter((m) => (m.model?.id ?? '').toLowerCase().includes(f) || + (m.provider_name ?? '').toLowerCase().includes(f) || (m.provider_type ?? '').toLowerCase().includes(f) || + (m.selector ?? '').toLowerCase().includes(f) || (m.model?.owned_by ?? '').toLowerCase().includes(f) || (m.model?.metadata?.modes ?? []).join(',').toLowerCase().includes(f) || (m.model?.metadata?.categories ?? []).join(',').toLowerCase().includes(f) ); }, + providerTypeValue(value) { + return String(value && value.provider || '').trim(); + }, + + providerDisplayValue(value) { + const providerName = String(value && value.provider_name || '').trim(); + if (providerName) return providerName; + return this.providerTypeValue(value); + }, + + qualifiedModelDisplay(value) { + const model = String(value && value.model || '').trim(); + if (!model) return '-'; + if (model.includes('/')) return model; + const provider = this.providerDisplayValue(value); + if (!provider) return model; + return provider + '/' + model; + }, + + qualifiedResolvedModelDisplay(value) { + const model = String(value && value.resolved_model || '').trim(); + if (!model) return '-'; + if (model.includes('/')) return model; + const provider = this.providerDisplayValue(value); + if (!provider) return model; + return provider + '/' + model; + }, + formatNumber(n) { if (n == null || n === undefined) return '-'; return n.toLocaleString(); diff --git a/internal/admin/dashboard/static/js/modules/audit-list.js b/internal/admin/dashboard/static/js/modules/audit-list.js index 2ce2a475e..e0e9b7cc4 100644 --- a/internal/admin/dashboard/static/js/modules/audit-list.js +++ b/internal/admin/dashboard/static/js/modules/audit-list.js @@ -23,11 +23,7 @@ let qs = this._auditQueryStr(); qs += '&limit=' + this.auditLog.limit + '&offset=' + this.auditLog.offset; if (this.auditSearch) qs += '&search=' + encodeURIComponent(this.auditSearch); - if (this.auditModel) qs += '&model=' + encodeURIComponent(this.auditModel); - if (this.auditProvider) qs += '&provider=' + encodeURIComponent(this.auditProvider); if (this.auditMethod) qs += '&method=' + encodeURIComponent(this.auditMethod); - if (this.auditPath) qs += '&path=' + encodeURIComponent(this.auditPath); - if (this.auditUserPath) qs += '&user_path=' + encodeURIComponent(this.auditUserPath); if (this.auditStatusCode) qs += '&status_code=' + encodeURIComponent(this.auditStatusCode); if (this.auditStream) qs += '&stream=' + encodeURIComponent(this.auditStream); @@ -57,11 +53,7 @@ clearAuditFilters() { this.auditSearch = ''; - this.auditModel = ''; - this.auditProvider = ''; this.auditMethod = ''; - this.auditPath = ''; - this.auditUserPath = ''; this.auditStatusCode = ''; this.auditStream = ''; this.fetchAuditLog(true); diff --git a/internal/admin/dashboard/static/js/modules/audit-list.test.js b/internal/admin/dashboard/static/js/modules/audit-list.test.js index ca0bd2b84..4f5af6634 100644 --- a/internal/admin/dashboard/static/js/modules/audit-list.test.js +++ b/internal/admin/dashboard/static/js/modules/audit-list.test.js @@ -125,6 +125,62 @@ test('fetchAuditLog preserves a successful payload when workflow prefetch fails' assert.match(String(loggedErrors[0][0]), /Failed to prefetch audit workflows:/); }); +test('fetchAuditLog sends the consolidated audit search and select filters only', async () => { + const requests = []; + const module = createAuditListModule({ + fetch(url) { + requests.push(url); + return Promise.resolve({ + ok: true, + json: async () => ({ + entries: [], + total: 0, + limit: 25, + offset: 0 + }) + }); + } + }); + module.auditFetchToken = 0; + module.auditLog = { entries: [], total: 0, limit: 25, offset: 0 }; + module.days = 30; + module.auditSearch = 'team/alpha'; + module.auditMethod = 'POST'; + module.auditStatusCode = '500'; + module.auditStream = 'true'; + module.headers = () => ({}); + module.handleFetchResponse = () => true; + + await module.fetchAuditLog(true); + + assert.equal(requests.length, 1); + assert.match(requests[0], /search=team%2Falpha/); + assert.match(requests[0], /method=POST/); + assert.match(requests[0], /status_code=500/); + assert.match(requests[0], /stream=true/); + assert.doesNotMatch(requests[0], /[?&](model|provider|path|user_path)=/); +}); + +test('clearAuditFilters resets the consolidated audit controls', () => { + const module = createAuditListModule(); + let fetchCalled = false; + module.auditSearch = 'req_123'; + module.auditMethod = 'DELETE'; + module.auditStatusCode = '404'; + module.auditStream = 'false'; + module.fetchAuditLog = (resetOffset) => { + fetchCalled = resetOffset === true; + }; + + module.clearAuditFilters(); + + assert.equal(module.auditSearch, ''); + assert.equal(module.auditMethod, ''); + assert.equal(module.auditStatusCode, ''); + assert.equal(module.auditStream, ''); + assert.equal(fetchCalled, true); +}); + test('auditPaneState copies the formatted body and resets success feedback', async () => { let resetCallback = null; const writes = []; diff --git a/internal/admin/dashboard/static/js/modules/charts.js b/internal/admin/dashboard/static/js/modules/charts.js index 8ef2bee50..d56290d3a 100644 --- a/internal/admin/dashboard/static/js/modules/charts.js +++ b/internal/admin/dashboard/static/js/modules/charts.js @@ -245,7 +245,9 @@ const top = sorted.slice(0, 10); const rest = sorted.slice(10); - const labels = top.map((m) => m.model); + const labels = top.map((m) => typeof this.qualifiedModelDisplay === 'function' + ? this.qualifiedModelDisplay(m) + : m.model); const values = top.map((m) => { if (this.usageMode === 'costs') return m.total_cost || 0; return m.input_tokens + m.output_tokens; diff --git a/internal/admin/dashboard/static/js/modules/charts.test.js b/internal/admin/dashboard/static/js/modules/charts.test.js index b3343fe8d..59faf59a1 100644 --- a/internal/admin/dashboard/static/js/modules/charts.test.js +++ b/internal/admin/dashboard/static/js/modules/charts.test.js @@ -67,6 +67,7 @@ function createChartsContext() { module.interval = 'weekly'; module.page = 'overview'; module.formatTokensShort = (value) => String(value); + module.qualifiedModelDisplay = (value) => value && value.model ? value.model : '-'; return { module, canvas }; } @@ -125,3 +126,20 @@ test('renderBarChart recreates the usage bar chart instance on refresh', () => { assert.equal(JSON.stringify(module.usageBarChart.data.datasets[0].data), JSON.stringify([55])); assert.equal(JSON.stringify(firstChart.updateCalls), JSON.stringify([])); }); + +test('renderBarChart prefers provider_name in model labels when available', () => { + FakeChart.instances = []; + const { module } = createChartsContext(); + module.page = 'usage'; + module.usageMode = 'tokens'; + module.qualifiedModelDisplay = (value) => value.provider_name + ? value.provider_name + '/' + value.model + : value.model; + module.modelUsage = [ + { model: 'gpt-4o', provider_name: 'primary-openai', input_tokens: 5, output_tokens: 7, total_cost: 0.01 } + ]; + + module.renderBarChart(); + + assert.equal(JSON.stringify(module.usageBarChart.data.labels), JSON.stringify(['primary-openai/gpt-4o'])); +}); diff --git a/internal/admin/dashboard/static/js/modules/dashboard-layout.test.js b/internal/admin/dashboard/static/js/modules/dashboard-layout.test.js index 1e5df2910..72da9d0df 100644 --- a/internal/admin/dashboard/static/js/modules/dashboard-layout.test.js +++ b/internal/admin/dashboard/static/js/modules/dashboard-layout.test.js @@ -56,9 +56,24 @@ test('sidebar and main content share the flex layout without manual content offs assert.match(collapsedSidebarRule, /flex-basis:\s*60px/); }); +test('mono utility only sets the font family and font-size-md carries the 13px size', () => { + const css = readFixture('../../css/dashboard.css'); + + const monoRule = readCSSRule(css, '.mono'); + assert.match(monoRule, /font-family:\s*'SF Mono', Menlo, Consolas, monospace/); + assert.doesNotMatch(monoRule, /font-size:/); + + const fontSizeMdRule = readCSSRule(css, '.font-size-md'); + assert.match(fontSizeMdRule, /font-size:\s*13px/); +}); + test('dashboard layout pins Chart.js to 4.5.0', () => { const template = readFixture('../../../templates/layout.html'); + assert.match( + template, + // + ); assert.match( template, / - +
diff --git a/internal/admin/dashboard/templates/x-icon.html b/internal/admin/dashboard/templates/x-icon.html new file mode 100644 index 000000000..db1fd85b5 --- /dev/null +++ b/internal/admin/dashboard/templates/x-icon.html @@ -0,0 +1,6 @@ +{{define "x-icon"}} + +{{end}} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index a6e103dab..fc6e16dce 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -377,9 +377,23 @@ func (h *Handler) DailyUsage(c *echo.Context) error { // @Failure 401 {object} core.GatewayError // @Router /admin/api/v1/usage/models [get] func (h *Handler) UsageByModel(c *echo.Context) error { - return usageSliceResponse(c, h.usageReader, func(ctx context.Context, params usage.UsageQueryParams) ([]usage.ModelUsage, error) { - return h.usageReader.GetUsageByModel(ctx, params) - }) + if h.usageReader == nil { + return c.JSON(http.StatusOK, []usage.ModelUsage{}) + } + + params, err := parseUsageParams(c) + if err != nil { + return handleError(c, err) + } + + values, err := h.usageReader.GetUsageByModel(c.Request().Context(), params) + if err != nil { + return handleError(c, err) + } + if values == nil { + values = []usage.ModelUsage{} + } + return c.JSON(http.StatusOK, values) } // UsageLog handles GET /admin/api/v1/usage/log @@ -392,7 +406,7 @@ func (h *Handler) UsageByModel(c *echo.Context) error { // @Param start_date query string false "Start date (YYYY-MM-DD)" // @Param end_date query string false "End date (YYYY-MM-DD)" // @Param model query string false "Filter by model name" -// @Param provider query string false "Filter by provider" +// @Param provider query string false "Filter by provider name or provider type" // @Param user_path query string false "Filter by tracked user path subtree" // @Param cache_mode query string false "Cache mode filter: uncached, cached, all (default uncached)" // @Param search query string false "Search across model, provider, request_id, provider_id" @@ -500,15 +514,15 @@ func (h *Handler) CacheOverview(c *echo.Context) error { // @Param days query int false "Number of days (default 30)" // @Param start_date query string false "Start date (YYYY-MM-DD)" // @Param end_date query string false "End date (YYYY-MM-DD)" -// @Param model query string false "Filter by model name" -// @Param provider query string false "Filter by provider" +// @Param requested_model query string false "Filter by requested model selector" +// @Param provider query string false "Filter by provider name or provider type" // @Param method query string false "Filter by HTTP method" // @Param path query string false "Filter by request path" // @Param user_path query string false "Filter by tracked user path subtree" // @Param error_type query string false "Filter by error type" // @Param status_code query int false "Filter by status code" // @Param stream query bool false "Filter by stream mode (true/false)" -// @Param search query string false "Search across request_id/model/provider/method/path/error_type" +// @Param search query string false "Search across request_id/requested_model/provider/method/path/error_type" // @Param limit query int false "Page size (default 25, max 100)" // @Param offset query int false "Offset for pagination" // @Success 200 {object} auditlog.LogListResult @@ -531,18 +545,23 @@ func (h *Handler) AuditLog(c *echo.Context) error { return handleError(c, err) } + requestedModel := c.QueryParam("requested_model") + if requestedModel == "" { + requestedModel = c.QueryParam("model") + } + params := auditlog.LogQueryParams{ QueryParams: auditlog.QueryParams{ StartDate: dateRange.StartDate, EndDate: dateRange.EndDate, }, - Model: c.QueryParam("model"), - Provider: c.QueryParam("provider"), - Method: strings.ToUpper(c.QueryParam("method")), - Path: c.QueryParam("path"), - UserPath: userPath, - ErrorType: c.QueryParam("error_type"), - Search: c.QueryParam("search"), + RequestedModel: requestedModel, + Provider: c.QueryParam("provider"), + Method: strings.ToUpper(c.QueryParam("method")), + Path: c.QueryParam("path"), + UserPath: userPath, + ErrorType: c.QueryParam("error_type"), + Search: c.QueryParam("search"), } if sc := c.QueryParam("status_code"); sc != "" { diff --git a/internal/admin/handler_test.go b/internal/admin/handler_test.go index 63909ac92..7316da3b1 100644 --- a/internal/admin/handler_test.go +++ b/internal/admin/handler_test.go @@ -441,6 +441,37 @@ func TestUsageByModel_Success(t *testing.T) { } } +func TestUsageByModel_PreservesProviderName(t *testing.T) { + reader := &mockUsageReader{ + modelUsage: []usage.ModelUsage{ + {Model: "gpt-4o", Provider: "openai", ProviderName: "primary-openai", InputTokens: 100, OutputTokens: 25}, + }, + } + h := NewHandler(reader, nil) + c, rec := newHandlerContext("/admin/api/v1/usage/models?days=30") + + if err := h.UsageByModel(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", rec.Code) + } + + var models []usage.ModelUsage + if err := json.Unmarshal(rec.Body.Bytes(), &models); err != nil { + t.Fatalf("failed to unmarshal: %v", err) + } + if len(models) != 1 { + t.Fatalf("expected 1 entry, got %d", len(models)) + } + if models[0].ProviderName != "primary-openai" { + t.Fatalf("ProviderName = %q, want %q", models[0].ProviderName, "primary-openai") + } + if models[0].Provider != "openai" { + t.Fatalf("Provider = %q, want %q", models[0].Provider, "openai") + } +} + func TestUsageByModel_Error(t *testing.T) { reader := &mockUsageReader{ modelUsageErr: errors.New("db failure"), @@ -525,6 +556,51 @@ func TestUsageLog_Success(t *testing.T) { } } +func TestUsageLog_PreservesProviderName(t *testing.T) { + now := time.Now().UTC() + reader := &mockUsageReader{ + usageLog: &usage.UsageLogResult{ + Entries: []usage.UsageLogEntry{ + { + ID: "1", + RequestID: "req-1", + Model: "gpt-4o", + Provider: "openai", + ProviderName: "primary-openai", + Timestamp: now, + InputTokens: 100, + TotalTokens: 100, + }, + }, + Total: 1, + Limit: 50, + }, + } + h := NewHandler(reader, nil) + c, rec := newHandlerContext("/admin/api/v1/usage/log?days=30") + + if err := h.UsageLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", rec.Code) + } + + var result usage.UsageLogResult + if err := json.Unmarshal(rec.Body.Bytes(), &result); err != nil { + t.Fatalf("failed to unmarshal: %v", err) + } + if len(result.Entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(result.Entries)) + } + if result.Entries[0].ProviderName != "primary-openai" { + t.Fatalf("ProviderName = %q, want %q", result.Entries[0].ProviderName, "primary-openai") + } + if result.Entries[0].Provider != "openai" { + t.Fatalf("Provider = %q, want %q", result.Entries[0].Provider, "openai") + } +} + func TestUsageLog_Error(t *testing.T) { reader := &mockUsageReader{ usageLogErr: core.NewProviderError("test", http.StatusBadGateway, "upstream failed", nil), @@ -591,15 +667,15 @@ func TestAuditLog_Success(t *testing.T) { logResult: &auditlog.LogListResult{ Entries: []auditlog.LogEntry{ { - ID: "log-1", - Timestamp: now, - DurationNs: 12_000_000, - Model: "gpt-4o", - Provider: "openai", - StatusCode: 200, - RequestID: "req-1", - Method: http.MethodPost, - Path: "/v1/chat/completions", + ID: "log-1", + Timestamp: now, + DurationNs: 12_000_000, + RequestedModel: "gpt-4o", + Provider: "openai", + StatusCode: 200, + RequestID: "req-1", + Method: http.MethodPost, + Path: "/v1/chat/completions", Data: &auditlog.LogData{ RequestBody: map[string]any{ "model": "gpt-4o", @@ -644,6 +720,50 @@ func TestAuditLog_Success(t *testing.T) { } } +func TestAuditLog_PreservesProviderName(t *testing.T) { + now := time.Now().UTC() + reader := &mockAuditReader{ + logResult: &auditlog.LogListResult{ + Entries: []auditlog.LogEntry{ + { + ID: "log-1", + Timestamp: now, + RequestedModel: "smart", + ResolvedModel: "primary-openai/gpt-4o", + Provider: "openai", + ProviderName: "primary-openai", + StatusCode: 200, + }, + }, + Total: 1, + Limit: 25, + }, + } + h := NewHandler(nil, nil, WithAuditReader(reader)) + c, rec := newHandlerContext("/admin/api/v1/audit/log?days=7") + + if err := h.AuditLog(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", rec.Code) + } + + var result auditlog.LogListResult + if err := json.Unmarshal(rec.Body.Bytes(), &result); err != nil { + t.Fatalf("failed to unmarshal: %v", err) + } + if len(result.Entries) != 1 { + t.Fatalf("expected 1 entry, got %d", len(result.Entries)) + } + if result.Entries[0].ProviderName != "primary-openai" { + t.Fatalf("ProviderName = %q, want %q", result.Entries[0].ProviderName, "primary-openai") + } + if result.Entries[0].Provider != "openai" { + t.Fatalf("Provider = %q, want %q", result.Entries[0].Provider, "openai") + } +} + func TestAuditLog_WithFilters(t *testing.T) { reader := &mockAuditReader{ logResult: &auditlog.LogListResult{ @@ -664,8 +784,8 @@ func TestAuditLog_WithFilters(t *testing.T) { t.Errorf("expected 200, got %d", rec.Code) } - if reader.lastQuery.Model != "gpt-4" { - t.Errorf("expected model filter gpt-4, got %q", reader.lastQuery.Model) + if reader.lastQuery.RequestedModel != "gpt-4" { + t.Errorf("expected requested model filter gpt-4, got %q", reader.lastQuery.RequestedModel) } if reader.lastQuery.Provider != "openai" { t.Errorf("expected provider filter openai, got %q", reader.lastQuery.Provider) @@ -795,6 +915,42 @@ func TestAuditConversation_Success(t *testing.T) { } } +func TestAuditConversation_PreservesProviderName(t *testing.T) { + now := time.Now().UTC() + reader := &mockAuditReader{ + conversationResult: &auditlog.ConversationResult{ + AnchorID: "log-2", + Entries: []auditlog.LogEntry{ + {ID: "log-1", Timestamp: now.Add(-time.Minute), ResolvedModel: "primary-openai/gpt-4o", Provider: "openai", ProviderName: "primary-openai", Path: "/v1/responses"}, + {ID: "log-2", Timestamp: now, ResolvedModel: "primary-openai/gpt-4o", Provider: "openai", ProviderName: "primary-openai", Path: "/v1/responses"}, + }, + }, + } + h := NewHandler(nil, nil, WithAuditReader(reader)) + c, rec := newHandlerContext("/admin/api/v1/audit/conversation?log_id=log-2") + + if err := h.AuditConversation(c); err != nil { + t.Fatalf("unexpected error: %v", err) + } + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d", rec.Code) + } + + var result auditlog.ConversationResult + if err := json.Unmarshal(rec.Body.Bytes(), &result); err != nil { + t.Fatalf("failed to unmarshal: %v", err) + } + if len(result.Entries) != 2 { + t.Fatalf("expected 2 entries, got %d", len(result.Entries)) + } + if result.Entries[0].ProviderName != "primary-openai" { + t.Fatalf("ProviderName = %q, want %q", result.Entries[0].ProviderName, "primary-openai") + } + if result.Entries[0].Provider != "openai" { + t.Fatalf("Provider = %q, want %q", result.Entries[0].Provider, "openai") + } +} + func TestAuditConversation_MissingLogID(t *testing.T) { reader := &mockAuditReader{} h := NewHandler(nil, nil, WithAuditReader(reader)) diff --git a/internal/aliases/provider.go b/internal/aliases/provider.go index dc290b83d..fce30274f 100644 --- a/internal/aliases/provider.go +++ b/internal/aliases/provider.go @@ -156,6 +156,13 @@ func (p *Provider) GetProviderType(model string) string { return p.inner.GetProviderType(model) } +func (p *Provider) GetProviderName(model string) string { + if named, ok := p.inner.(core.ProviderNameResolver); ok { + return strings.TrimSpace(named.GetProviderName(model)) + } + return "" +} + func (p *Provider) ModelCount() int { if counted, ok := p.inner.(interface{ ModelCount() int }); ok { return counted.ModelCount() diff --git a/internal/auditlog/auditlog.go b/internal/auditlog/auditlog.go index 7c442ce05..b352a7825 100644 --- a/internal/auditlog/auditlog.go +++ b/internal/auditlog/auditlog.go @@ -47,9 +47,10 @@ type LogEntry struct { DurationNs int64 `json:"duration_ns" bson:"duration_ns"` // Core fields (indexed for queries) - Model string `json:"model" bson:"model"` + RequestedModel string `json:"requested_model" bson:"requested_model,omitempty"` ResolvedModel string `json:"resolved_model,omitempty" bson:"resolved_model,omitempty"` - Provider string `json:"provider" bson:"provider"` + Provider string `json:"provider" bson:"provider"` // canonical provider type used for routing and filters + ProviderName string `json:"provider_name,omitempty" bson:"provider_name,omitempty"` AliasUsed bool `json:"alias_used,omitempty" bson:"alias_used,omitempty"` ExecutionPlanVersionID string `json:"execution_plan_version_id,omitempty" bson:"execution_plan_version_id,omitempty"` CacheType string `json:"cache_type,omitempty" bson:"cache_type,omitempty"` @@ -143,6 +144,13 @@ func normalizeCacheType(value string) string { } } +func displayAuditProviderName(providerName, provider string) string { + if trimmed := strings.TrimSpace(providerName); trimmed != "" { + return trimmed + } + return strings.TrimSpace(provider) +} + // RedactedHeaders contains headers that should be automatically redacted. // Values are replaced with "[REDACTED]" to prevent leaking secrets. var RedactedHeaders = []string{ diff --git a/internal/auditlog/auditlog_test.go b/internal/auditlog/auditlog_test.go index e2081c994..b9d083827 100644 --- a/internal/auditlog/auditlog_test.go +++ b/internal/auditlog/auditlog_test.go @@ -119,20 +119,20 @@ func TestRedactHeaders(t *testing.T) { func TestLogEntryJSON(t *testing.T) { entry := &LogEntry{ - ID: "test-id-123", - Timestamp: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC), - DurationNs: 1500000, - Model: "friendly-alias", - ResolvedModel: "openai/gpt-4", - Provider: "openai", - AliasUsed: true, - CacheType: CacheTypeExact, - StatusCode: 200, - RequestID: "req-123", - ClientIP: "192.168.1.1", - Method: "POST", - Path: "/v1/chat/completions", - Stream: false, + ID: "test-id-123", + Timestamp: time.Date(2024, 1, 15, 10, 30, 0, 0, time.UTC), + DurationNs: 1500000, + RequestedModel: "friendly-alias", + ResolvedModel: "openai/gpt-4", + Provider: "openai", + AliasUsed: true, + CacheType: CacheTypeExact, + StatusCode: 200, + RequestID: "req-123", + ClientIP: "192.168.1.1", + Method: "POST", + Path: "/v1/chat/completions", + Stream: false, Data: &LogData{ UserAgent: "test-agent", }, @@ -154,8 +154,8 @@ func TestLogEntryJSON(t *testing.T) { if decoded.ID != entry.ID { t.Errorf("ID mismatch: expected %q, got %q", entry.ID, decoded.ID) } - if decoded.Model != entry.Model { - t.Errorf("Model mismatch: expected %q, got %q", entry.Model, decoded.Model) + if decoded.RequestedModel != entry.RequestedModel { + t.Errorf("RequestedModel mismatch: expected %q, got %q", entry.RequestedModel, decoded.RequestedModel) } if decoded.Provider != entry.Provider { t.Errorf("Provider mismatch: expected %q, got %q", entry.Provider, decoded.Provider) @@ -308,9 +308,9 @@ func TestLogger(t *testing.T) { // Write some entries for i := range 5 { logger.Write(&LogEntry{ - ID: fmt.Sprintf("entry-%d", i), - Timestamp: time.Now(), - Model: "test-model", + ID: fmt.Sprintf("entry-%d", i), + Timestamp: time.Now(), + RequestedModel: "test-model", }) } @@ -498,8 +498,8 @@ func TestMiddleware_PrefersExecutionPlanOverLegacyResolution(t *testing.T) { } entry := logger.entries[0] - if entry.Model != "anthropic/claude-opus-4-6" { - t.Fatalf("Model = %q, want requested alias", entry.Model) + if entry.RequestedModel != "anthropic/claude-opus-4-6" { + t.Fatalf("RequestedModel = %q, want requested alias", entry.RequestedModel) } if entry.ResolvedModel != "openai/gpt-5-nano" { t.Fatalf("ResolvedModel = %q, want openai/gpt-5-nano", entry.ResolvedModel) @@ -573,8 +573,8 @@ func TestMiddleware_DoesNotApplyModelMetadataWithoutExecutionPlan(t *testing.T) } entry := logger.entries[0] - if entry.Model != "" { - t.Fatalf("Model = %q, want empty", entry.Model) + if entry.RequestedModel != "" { + t.Fatalf("RequestedModel = %q, want empty", entry.RequestedModel) } if entry.ResolvedModel != "" { t.Fatalf("ResolvedModel = %q, want empty", entry.ResolvedModel) @@ -619,8 +619,8 @@ func TestMiddleware_PassthroughExecutionPlanUsesPassthroughModel(t *testing.T) { } entry := logger.entries[0] - if entry.Model != "gpt-4.1-nano" { - t.Fatalf("Model = %q, want gpt-4.1-nano", entry.Model) + if entry.RequestedModel != "gpt-4.1-nano" { + t.Fatalf("RequestedModel = %q, want gpt-4.1-nano", entry.RequestedModel) } if entry.Provider != "openai" { t.Fatalf("Provider = %q, want openai", entry.Provider) @@ -890,10 +890,10 @@ data: [DONE] logger := NewLogger(store, cfg) entry := &LogEntry{ - ID: "test-entry", - Timestamp: time.Now(), - Model: "gpt-4", - Data: &LogData{}, + ID: "test-entry", + Timestamp: time.Now(), + RequestedModel: "gpt-4", + Data: &LogData{}, } observedStream := streaming.NewObservedSSEStream( @@ -943,7 +943,7 @@ func TestCreateStreamEntry(t *testing.T) { ID: "test-id", Timestamp: time.Now(), DurationNs: 1000, - Model: "claude-opus-4-6", + RequestedModel: "claude-opus-4-6", ResolvedModel: "openai/gpt-5-nano", Provider: "openai", AliasUsed: true, @@ -983,8 +983,8 @@ func TestCreateStreamEntry(t *testing.T) { if streamEntry.ID != baseEntry.ID { t.Errorf("ID mismatch") } - if streamEntry.Model != baseEntry.Model { - t.Errorf("Model mismatch") + if streamEntry.RequestedModel != baseEntry.RequestedModel { + t.Errorf("RequestedModel mismatch") } if streamEntry.ResolvedModel != baseEntry.ResolvedModel { t.Errorf("ResolvedModel mismatch") diff --git a/internal/auditlog/logger.go b/internal/auditlog/logger.go index de50a9f8e..8be33369b 100644 --- a/internal/auditlog/logger.go +++ b/internal/auditlog/logger.go @@ -79,7 +79,7 @@ func (l *Logger) Write(entry *LogEntry) { } slog.Warn("audit log buffer full, dropping entry", "request_id", requestID, - "model", entry.Model, + "requested_model", entry.RequestedModel, ) } } diff --git a/internal/auditlog/middleware.go b/internal/auditlog/middleware.go index 465eedc7b..3b06be0f6 100644 --- a/internal/auditlog/middleware.go +++ b/internal/auditlog/middleware.go @@ -190,14 +190,14 @@ func enrichEntryWithExecutionPlan(entry *LogEntry, plan *core.ExecutionPlan) { entry.RequestID = requestID } if requestedModel := plan.RequestedQualifiedModel(); requestedModel != "" { - entry.Model = requestedModel + entry.RequestedModel = requestedModel } - if resolvedModel := plan.ResolvedQualifiedModel(); resolvedModel != "" { + if resolvedModel := resolvedModelForAuditLog(plan); resolvedModel != "" { entry.ResolvedModel = resolvedModel } if plan.Mode == core.ExecutionModePassthrough && plan.Passthrough != nil { if model := strings.TrimSpace(plan.Passthrough.Model); model != "" { - entry.Model = model + entry.RequestedModel = model } } if providerType := strings.TrimSpace(plan.ProviderType); providerType != "" { @@ -206,6 +206,9 @@ func enrichEntryWithExecutionPlan(entry *LogEntry, plan *core.ExecutionPlan) { entry.Provider = strings.TrimSpace(plan.Resolution.ProviderType) } if plan.Resolution != nil { + if providerName := strings.TrimSpace(plan.Resolution.ProviderName); providerName != "" { + entry.ProviderName = providerName + } entry.AliasUsed = plan.Resolution.AliasApplied } if versionID := strings.TrimSpace(plan.ExecutionPlanVersionID()); versionID != "" { @@ -222,6 +225,23 @@ func enrichEntryWithExecutionPlan(entry *LogEntry, plan *core.ExecutionPlan) { } } +func resolvedModelForAuditLog(plan *core.ExecutionPlan) string { + if plan == nil || plan.Resolution == nil { + return "" + } + model := strings.TrimSpace(plan.Resolution.ResolvedSelector.Model) + if model == "" { + return "" + } + if providerName := strings.TrimSpace(plan.Resolution.ProviderName); providerName != "" { + return providerName + "/" + model + } + if provider := strings.TrimSpace(plan.Resolution.ResolvedSelector.Provider); provider != "" { + return provider + "/" + model + } + return model +} + func captureLoggedRequestBody(entry *LogEntry, bodyBytes []byte) { entry.Data.RequestBody = captureLoggedBody(bodyBytes) } @@ -364,7 +384,7 @@ func EnrichEntry(c *echo.Context, model, provider string) { return } - entry.Model = model + entry.RequestedModel = model entry.Provider = provider } @@ -392,6 +412,43 @@ func EnrichLogEntryWithExecutionPlan(entry *LogEntry, plan *core.ExecutionPlan) enrichEntryWithExecutionPlan(entry, plan) } +// EnrichEntryWithResolvedRoute attaches the final executed route to the live +// audit entry after execution resolved to a concrete provider/model. +func EnrichEntryWithResolvedRoute(c *echo.Context, resolvedModel, providerType, providerName string) { + entryVal := c.Get(string(LogEntryKey)) + if entryVal == nil { + return + } + + entry, ok := entryVal.(*LogEntry) + if !ok || entry == nil { + return + } + + enrichEntryWithResolvedRoute(entry, resolvedModel, providerType, providerName) +} + +// EnrichLogEntryWithResolvedRoute attaches the final executed route directly to +// an existing audit log entry. +func EnrichLogEntryWithResolvedRoute(entry *LogEntry, resolvedModel, providerType, providerName string) { + enrichEntryWithResolvedRoute(entry, resolvedModel, providerType, providerName) +} + +func enrichEntryWithResolvedRoute(entry *LogEntry, resolvedModel, providerType, providerName string) { + if entry == nil { + return + } + if resolvedModel = strings.TrimSpace(resolvedModel); resolvedModel != "" { + entry.ResolvedModel = resolvedModel + } + if providerType = strings.TrimSpace(providerType); providerType != "" { + entry.Provider = providerType + } + if providerName = strings.TrimSpace(providerName); providerName != "" { + entry.ProviderName = providerName + } +} + // EnrichEntryWithCacheType attaches cache-hit metadata to the live audit entry. // The value is intentionally sourced directly from the cache middleware, not // inferred from response headers after the fact. diff --git a/internal/auditlog/middleware_test.go b/internal/auditlog/middleware_test.go new file mode 100644 index 000000000..fc25a7436 --- /dev/null +++ b/internal/auditlog/middleware_test.go @@ -0,0 +1,42 @@ +package auditlog + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/labstack/echo/v5" + + "gomodel/internal/core" +) + +func TestEnrichEntryWithExecutionPlan_PrefersProviderNameForResolvedModel(t *testing.T) { + e := echo.New() + req := httptest.NewRequest(http.MethodGet, "/", nil) + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + entry := &LogEntry{ID: "provider-name-prefill"} + c.Set(string(LogEntryKey), entry) + + EnrichEntryWithExecutionPlan(c, &core.ExecutionPlan{ + ProviderType: "openai", + Resolution: &core.RequestModelResolution{ + ResolvedSelector: core.ModelSelector{ + Provider: "openai", + Model: "gpt-5-nano", + }, + ProviderName: "openai_test", + }, + }) + + if got := entry.Provider; got != "openai" { + t.Fatalf("Provider = %q, want %q", got, "openai") + } + if got := entry.ProviderName; got != "openai_test" { + t.Fatalf("ProviderName = %q, want %q", got, "openai_test") + } + if got := entry.ResolvedModel; got != "openai_test/gpt-5-nano" { + t.Fatalf("ResolvedModel = %q, want %q", got, "openai_test/gpt-5-nano") + } +} diff --git a/internal/auditlog/reader.go b/internal/auditlog/reader.go index 31d1bab53..0d0bac72a 100644 --- a/internal/auditlog/reader.go +++ b/internal/auditlog/reader.go @@ -14,17 +14,17 @@ type QueryParams struct { // LogQueryParams specifies query parameters for paginated audit log retrieval. type LogQueryParams struct { QueryParams - Model string - Provider string - Method string - Path string - UserPath string - ErrorType string - Search string - StatusCode *int - Stream *bool - Limit int - Offset int + RequestedModel string + Provider string // filter by provider name or provider type + Method string + Path string + UserPath string + ErrorType string + Search string + StatusCode *int + Stream *bool + Limit int + Offset int } // LogListResult holds a paginated list of audit log entries. diff --git a/internal/auditlog/reader_mongodb.go b/internal/auditlog/reader_mongodb.go index edaec0ecc..0f3704533 100644 --- a/internal/auditlog/reader_mongodb.go +++ b/internal/auditlog/reader_mongodb.go @@ -22,9 +22,11 @@ type mongoLogRow struct { ID string `bson:"_id"` Timestamp time.Time `bson:"timestamp"` DurationNs int64 `bson:"duration_ns"` - Model string `bson:"model"` + RequestedModel string `bson:"requested_model"` + LegacyModel string `bson:"model"` ResolvedModel string `bson:"resolved_model"` Provider string `bson:"provider"` + ProviderName string `bson:"provider_name"` AliasUsed bool `bson:"alias_used"` ExecutionPlanVersionID string `bson:"execution_plan_version_id"` CacheType string `bson:"cache_type"` @@ -46,9 +48,10 @@ func (r mongoLogRow) toLogEntry() *LogEntry { ID: r.ID, Timestamp: r.Timestamp, DurationNs: r.DurationNs, - Model: r.Model, + RequestedModel: firstNonEmpty(r.RequestedModel, r.LegacyModel), ResolvedModel: r.ResolvedModel, Provider: r.Provider, + ProviderName: displayAuditProviderName(r.ProviderName, r.Provider), AliasUsed: r.AliasUsed, ExecutionPlanVersionID: r.ExecutionPlanVersionID, CacheType: normalizeCacheType(r.CacheType), @@ -113,23 +116,30 @@ func (r *MongoDBReader) GetLogs(ctx context.Context, params LogQueryParams) (*Lo if tsFilter := mongoDateRangeFilter(params.QueryParams); tsFilter != nil { matchFilters = append(matchFilters, bson.E{Key: "timestamp", Value: tsFilter}) } - if params.Model != "" { + if params.RequestedModel != "" { matchFilters = append(matchFilters, bson.E{ - Key: "model", - Value: bson.D{ - {Key: "$regex", Value: regexp.QuoteMeta(params.Model)}, - {Key: "$options", Value: "i"}, + Key: "$or", + Value: bson.A{ + bson.D{{Key: "requested_model", Value: bson.D{ + {Key: "$regex", Value: regexp.QuoteMeta(params.RequestedModel)}, + {Key: "$options", Value: "i"}, + }}}, + bson.D{{Key: "model", Value: bson.D{ + {Key: "$regex", Value: regexp.QuoteMeta(params.RequestedModel)}, + {Key: "$options", Value: "i"}, + }}}, }, }) } if params.Provider != "" { - matchFilters = append(matchFilters, bson.E{ - Key: "provider", - Value: bson.D{ - {Key: "$regex", Value: regexp.QuoteMeta(params.Provider)}, - {Key: "$options", Value: "i"}, - }, - }) + regex := bson.D{ + {Key: "$regex", Value: regexp.QuoteMeta(params.Provider)}, + {Key: "$options", Value: "i"}, + } + matchFilters = append(matchFilters, bson.E{Key: "$or", Value: bson.A{ + bson.D{{Key: "provider", Value: regex}}, + bson.D{{Key: "provider_name", Value: regex}}, + }}) } if params.Method != "" { matchFilters = append(matchFilters, bson.E{Key: "method", Value: params.Method}) @@ -169,10 +179,13 @@ func (r *MongoDBReader) GetLogs(ctx context.Context, params LogQueryParams) (*Lo matchFilters = append(matchFilters, bson.E{Key: "$or", Value: bson.A{ bson.D{{Key: "request_id", Value: regex}}, bson.D{{Key: "auth_key_id", Value: regex}}, + bson.D{{Key: "requested_model", Value: regex}}, bson.D{{Key: "model", Value: regex}}, bson.D{{Key: "provider", Value: regex}}, + bson.D{{Key: "provider_name", Value: regex}}, bson.D{{Key: "method", Value: regex}}, bson.D{{Key: "path", Value: regex}}, + bson.D{{Key: "user_path", Value: regex}}, bson.D{{Key: "error_type", Value: regex}}, bson.D{{Key: "data.error_message", Value: regex}}, }}) @@ -238,6 +251,15 @@ func (r *MongoDBReader) GetLogs(ctx context.Context, params LogQueryParams) (*Lo }, nil } +func firstNonEmpty(values ...string) string { + for _, value := range values { + if value != "" { + return value + } + } + return "" +} + // GetLogByID returns a single audit log entry by ID. func (r *MongoDBReader) GetLogByID(ctx context.Context, id string) (*LogEntry, error) { var row mongoLogRow diff --git a/internal/auditlog/reader_mongodb_test.go b/internal/auditlog/reader_mongodb_test.go index 96dcc50af..bbd94fb0a 100644 --- a/internal/auditlog/reader_mongodb_test.go +++ b/internal/auditlog/reader_mongodb_test.go @@ -59,10 +59,10 @@ func TestSanitizeLogDataNilSafe(t *testing.T) { func TestMongoLogRowToLogEntryPreservesCacheType(t *testing.T) { row := mongoLogRow{ - ID: "log-1", - Model: "gpt-4", - Provider: "openai", - CacheType: CacheTypeSemantic, + ID: "log-1", + RequestedModel: "gpt-4", + Provider: "openai", + CacheType: CacheTypeSemantic, } entry := row.toLogEntry() diff --git a/internal/auditlog/reader_postgresql.go b/internal/auditlog/reader_postgresql.go index 938c22af2..cb21d0b25 100644 --- a/internal/auditlog/reader_postgresql.go +++ b/internal/auditlog/reader_postgresql.go @@ -32,15 +32,15 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( return nil, err } - if params.Model != "" { - conditions = append(conditions, fmt.Sprintf("model ILIKE $%d ESCAPE '\\'", argIdx)) - args = append(args, "%"+escapeLikeWildcards(params.Model)+"%") + if params.RequestedModel != "" { + conditions = append(conditions, fmt.Sprintf("requested_model ILIKE $%d ESCAPE '\\'", argIdx)) + args = append(args, "%"+escapeLikeWildcards(params.RequestedModel)+"%") argIdx++ } if params.Provider != "" { - conditions = append(conditions, fmt.Sprintf("provider ILIKE $%d ESCAPE '\\'", argIdx)) - args = append(args, "%"+escapeLikeWildcards(params.Provider)+"%") - argIdx++ + conditions = append(conditions, fmt.Sprintf("(provider ILIKE $%d ESCAPE '\\' OR provider_name ILIKE $%d ESCAPE '\\')", argIdx, argIdx+1)) + args = append(args, "%"+escapeLikeWildcards(params.Provider)+"%", "%"+escapeLikeWildcards(params.Provider)+"%") + argIdx += 2 } if params.Method != "" { conditions = append(conditions, fmt.Sprintf("method = $%d", argIdx)) @@ -78,7 +78,7 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( } if params.Search != "" { s := "%" + escapeLikeWildcards(params.Search) + "%" - conditions = append(conditions, fmt.Sprintf("(request_id ILIKE $%d ESCAPE '\\' OR auth_key_id ILIKE $%d ESCAPE '\\' OR model ILIKE $%d ESCAPE '\\' OR provider ILIKE $%d ESCAPE '\\' OR method ILIKE $%d ESCAPE '\\' OR path ILIKE $%d ESCAPE '\\' OR error_type ILIKE $%d ESCAPE '\\')", argIdx, argIdx, argIdx, argIdx, argIdx, argIdx, argIdx)) + conditions = append(conditions, fmt.Sprintf("(request_id ILIKE $%d ESCAPE '\\' OR auth_key_id ILIKE $%d ESCAPE '\\' OR requested_model ILIKE $%d ESCAPE '\\' OR provider ILIKE $%d ESCAPE '\\' OR provider_name ILIKE $%d ESCAPE '\\' OR method ILIKE $%d ESCAPE '\\' OR path ILIKE $%d ESCAPE '\\' OR user_path ILIKE $%d ESCAPE '\\' OR error_type ILIKE $%d ESCAPE '\\')", argIdx, argIdx, argIdx, argIdx, argIdx, argIdx, argIdx, argIdx, argIdx)) args = append(args, s) argIdx++ } @@ -91,7 +91,7 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( return nil, fmt.Errorf("failed to count audit log entries: %w", err) } - dataQuery := fmt.Sprintf(`SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + dataQuery := fmt.Sprintf(`SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs%s ORDER BY timestamp DESC LIMIT $%d OFFSET $%d`, where, argIdx, argIdx+1) dataArgs := append(append([]any(nil), args...), limit, offset) @@ -106,13 +106,14 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( for rows.Next() { var e LogEntry var dataJSON *string + var providerName *string var executionPlanVersionID *string var cacheType *string var authKeyID *string var authMethod *string var userPath *string - if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.Model, &e.ResolvedModel, &e.Provider, &e.AliasUsed, &executionPlanVersionID, &cacheType, &e.StatusCode, + if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &e.AliasUsed, &executionPlanVersionID, &cacheType, &e.StatusCode, &e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &e.ErrorType, &dataJSON); err != nil { return nil, fmt.Errorf("failed to scan audit log row: %w", err) } @@ -128,6 +129,11 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( if cacheType != nil { e.CacheType = normalizeCacheType(*cacheType) } + if providerName != nil { + e.ProviderName = displayAuditProviderName(*providerName, e.Provider) + } else { + e.ProviderName = displayAuditProviderName("", e.Provider) + } if userPath != nil { e.UserPath = *userPath } @@ -158,7 +164,7 @@ func (r *PostgreSQLReader) GetLogs(ctx context.Context, params LogQueryParams) ( // GetLogByID returns a single audit log entry by ID. func (r *PostgreSQLReader) GetLogByID(ctx context.Context, id string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE id::text = $1 LIMIT 1` @@ -200,7 +206,7 @@ func pgDateRangeConditions(params QueryParams, argIdx int) (conditions []string, } func (r *PostgreSQLReader) findByResponseID(ctx context.Context, responseID string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE data->'response_body'->>'id' = $1 @@ -219,7 +225,7 @@ func (r *PostgreSQLReader) findByResponseID(ctx context.Context, responseID stri } func (r *PostgreSQLReader) findByPreviousResponseID(ctx context.Context, previousResponseID string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE data->'request_body'->>'previous_response_id' = $1 @@ -242,13 +248,14 @@ func scanPostgreSQLLogEntry(rows interface { }) (*LogEntry, error) { var e LogEntry var dataJSON *string + var providerName *string var executionPlanVersionID *string var cacheType *string var authKeyID *string var authMethod *string var userPath *string - if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.Model, &e.ResolvedModel, &e.Provider, &e.AliasUsed, &executionPlanVersionID, &cacheType, &e.StatusCode, + if err := rows.Scan(&e.ID, &e.Timestamp, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &e.AliasUsed, &executionPlanVersionID, &cacheType, &e.StatusCode, &e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &e.Stream, &e.ErrorType, &dataJSON); err != nil { return nil, fmt.Errorf("failed to scan audit log row: %w", err) } @@ -264,6 +271,11 @@ func scanPostgreSQLLogEntry(rows interface { if cacheType != nil { e.CacheType = normalizeCacheType(*cacheType) } + if providerName != nil { + e.ProviderName = displayAuditProviderName(*providerName, e.Provider) + } else { + e.ProviderName = displayAuditProviderName("", e.Provider) + } if userPath != nil { e.UserPath = *userPath } diff --git a/internal/auditlog/reader_sqlite.go b/internal/auditlog/reader_sqlite.go index 4f7162fa5..a18fec257 100644 --- a/internal/auditlog/reader_sqlite.go +++ b/internal/auditlog/reader_sqlite.go @@ -35,13 +35,13 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log return nil, err } - if params.Model != "" { - conditions = append(conditions, "model LIKE ? ESCAPE '\\'") - args = append(args, "%"+escapeLikeWildcards(params.Model)+"%") + if params.RequestedModel != "" { + conditions = append(conditions, "requested_model LIKE ? ESCAPE '\\'") + args = append(args, "%"+escapeLikeWildcards(params.RequestedModel)+"%") } if params.Provider != "" { - conditions = append(conditions, "provider LIKE ? ESCAPE '\\'") - args = append(args, "%"+escapeLikeWildcards(params.Provider)+"%") + conditions = append(conditions, "(provider LIKE ? ESCAPE '\\' OR provider_name LIKE ? ESCAPE '\\')") + args = append(args, "%"+escapeLikeWildcards(params.Provider)+"%", "%"+escapeLikeWildcards(params.Provider)+"%") } if params.Method != "" { conditions = append(conditions, "method = ?") @@ -73,8 +73,8 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log } if params.Search != "" { s := "%" + escapeLikeWildcards(params.Search) + "%" - conditions = append(conditions, `(request_id LIKE ? ESCAPE '\' OR auth_key_id LIKE ? ESCAPE '\' OR model LIKE ? ESCAPE '\' OR provider LIKE ? ESCAPE '\' OR method LIKE ? ESCAPE '\' OR path LIKE ? ESCAPE '\' OR error_type LIKE ? ESCAPE '\')`) - args = append(args, s, s, s, s, s, s, s) + conditions = append(conditions, `(request_id LIKE ? ESCAPE '\' OR auth_key_id LIKE ? ESCAPE '\' OR requested_model LIKE ? ESCAPE '\' OR provider LIKE ? ESCAPE '\' OR provider_name LIKE ? ESCAPE '\' OR method LIKE ? ESCAPE '\' OR path LIKE ? ESCAPE '\' OR user_path LIKE ? ESCAPE '\' OR error_type LIKE ? ESCAPE '\')`) + args = append(args, s, s, s, s, s, s, s, s, s) } where := buildWhereClause(conditions) @@ -86,7 +86,7 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log return nil, fmt.Errorf("failed to count audit log entries: %w", err) } - dataQuery := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + dataQuery := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs` + where + ` ORDER BY timestamp DESC LIMIT ? OFFSET ?` dataArgs := append(append([]any(nil), args...), limit, offset) @@ -101,6 +101,7 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log for rows.Next() { var e LogEntry var ts string + var providerName sql.NullString var aliasUsedInt int var streamInt int var dataJSON *string @@ -110,7 +111,7 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log var authMethod sql.NullString var userPath sql.NullString - if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.Model, &e.ResolvedModel, &e.Provider, &aliasUsedInt, &executionPlanVersionID, &cacheType, &e.StatusCode, + if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &aliasUsedInt, &executionPlanVersionID, &cacheType, &e.StatusCode, &e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &e.ErrorType, &dataJSON); err != nil { return nil, fmt.Errorf("failed to scan audit log row: %w", err) } @@ -130,6 +131,11 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log if cacheType.Valid { e.CacheType = normalizeCacheType(cacheType.String) } + if providerName.Valid { + e.ProviderName = displayAuditProviderName(providerName.String, e.Provider) + } else { + e.ProviderName = displayAuditProviderName("", e.Provider) + } if userPath.Valid { e.UserPath = userPath.String } @@ -160,7 +166,7 @@ func (r *SQLiteReader) GetLogs(ctx context.Context, params LogQueryParams) (*Log // GetLogByID returns a single audit log entry by ID. func (r *SQLiteReader) GetLogByID(ctx context.Context, id string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE id = ? LIMIT 1` @@ -292,7 +298,7 @@ func parseSQLTimestamp(ts string, entryID string) time.Time { } func (r *SQLiteReader) findByResponseID(ctx context.Context, responseID string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE json_extract(data, '$.response_body.id') = ? @@ -311,7 +317,7 @@ func (r *SQLiteReader) findByResponseID(ctx context.Context, responseID string) } func (r *SQLiteReader) findByPreviousResponseID(ctx context.Context, previousResponseID string) (*LogEntry, error) { - query := `SELECT id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, + query := `SELECT id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE json_extract(data, '$.request_body.previous_response_id') = ? @@ -332,6 +338,7 @@ func (r *SQLiteReader) findByPreviousResponseID(ctx context.Context, previousRes func scanSQLiteLogEntry(rows *sql.Rows) (*LogEntry, error) { var e LogEntry var ts string + var providerName sql.NullString var aliasUsedInt int var streamInt int var dataJSON *string @@ -341,7 +348,7 @@ func scanSQLiteLogEntry(rows *sql.Rows) (*LogEntry, error) { var authMethod sql.NullString var userPath sql.NullString - if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.Model, &e.ResolvedModel, &e.Provider, &aliasUsedInt, &executionPlanVersionID, &cacheType, &e.StatusCode, + if err := rows.Scan(&e.ID, &ts, &e.DurationNs, &e.RequestedModel, &e.ResolvedModel, &e.Provider, &providerName, &aliasUsedInt, &executionPlanVersionID, &cacheType, &e.StatusCode, &e.RequestID, &authKeyID, &authMethod, &e.ClientIP, &e.Method, &e.Path, &userPath, &streamInt, &e.ErrorType, &dataJSON); err != nil { return nil, fmt.Errorf("failed to scan audit log row: %w", err) } @@ -361,6 +368,11 @@ func scanSQLiteLogEntry(rows *sql.Rows) (*LogEntry, error) { if cacheType.Valid { e.CacheType = normalizeCacheType(cacheType.String) } + if providerName.Valid { + e.ProviderName = displayAuditProviderName(providerName.String, e.Provider) + } else { + e.ProviderName = displayAuditProviderName("", e.Provider) + } if userPath.Valid { e.UserPath = userPath.String } diff --git a/internal/auditlog/reader_sqlite_boundary_test.go b/internal/auditlog/reader_sqlite_boundary_test.go index 8ae536093..9686a9495 100644 --- a/internal/auditlog/reader_sqlite_boundary_test.go +++ b/internal/auditlog/reader_sqlite_boundary_test.go @@ -18,22 +18,22 @@ func TestSQLiteReaderGetLogs_IncludesFractionalStartBoundaryAndExcludesFractiona ctx := context.Background() err = store.WriteBatch(ctx, []*LogEntry{ { - ID: "start-boundary", - Timestamp: time.Date(2026, 1, 15, 23, 0, 0, 123_000_000, time.UTC), - Model: "gpt-5", - Provider: "openai", + ID: "start-boundary", + Timestamp: time.Date(2026, 1, 15, 23, 0, 0, 123_000_000, time.UTC), + RequestedModel: "gpt-5", + Provider: "openai", }, { - ID: "inside-range", - Timestamp: time.Date(2026, 1, 16, 12, 0, 0, 0, time.UTC), - Model: "gpt-5", - Provider: "openai", + ID: "inside-range", + Timestamp: time.Date(2026, 1, 16, 12, 0, 0, 0, time.UTC), + RequestedModel: "gpt-5", + Provider: "openai", }, { - ID: "after-end-boundary", - Timestamp: time.Date(2026, 1, 16, 23, 0, 0, 123_000_000, time.UTC), - Model: "gpt-5", - Provider: "openai", + ID: "after-end-boundary", + Timestamp: time.Date(2026, 1, 16, 23, 0, 0, 123_000_000, time.UTC), + RequestedModel: "gpt-5", + Provider: "openai", }, }) if err != nil { @@ -75,3 +75,56 @@ func TestSQLiteReaderGetLogs_IncludesFractionalStartBoundaryAndExcludesFractiona t.Fatalf("expected boundary entry %q, got %q", "start-boundary", result.Entries[1].ID) } } + +func TestSQLiteReaderGetLogs_SearchMatchesUserPath(t *testing.T) { + db := createTestDB(t) + defer db.Close() + + store, err := NewSQLiteStore(db, 0) + if err != nil { + t.Fatalf("failed to create store: %v", err) + } + + ctx := context.Background() + if err := store.WriteBatch(ctx, []*LogEntry{ + { + ID: "team-match", + Timestamp: time.Date(2026, 1, 16, 12, 0, 0, 0, time.UTC), + RequestedModel: "gpt-5", + Provider: "openai", + UserPath: "/team/alpha", + }, + { + ID: "other-team", + Timestamp: time.Date(2026, 1, 16, 11, 0, 0, 0, time.UTC), + RequestedModel: "gpt-5", + Provider: "openai", + UserPath: "/org/beta", + }, + }); err != nil { + t.Fatalf("failed to seed audit logs: %v", err) + } + + reader, err := NewSQLiteReader(db) + if err != nil { + t.Fatalf("failed to create reader: %v", err) + } + + result, err := reader.GetLogs(ctx, LogQueryParams{ + Search: "team/alpha", + Limit: 10, + }) + if err != nil { + t.Fatalf("GetLogs returned error: %v", err) + } + + if result.Total != 1 { + t.Fatalf("expected 1 log in search result, got %d", result.Total) + } + if len(result.Entries) != 1 { + t.Fatalf("expected 1 returned entry, got %d", len(result.Entries)) + } + if result.Entries[0].ID != "team-match" { + t.Fatalf("expected matching entry %q, got %q", "team-match", result.Entries[0].ID) + } +} diff --git a/internal/auditlog/store_mongodb.go b/internal/auditlog/store_mongodb.go index fbe812e12..ad1d02cb6 100644 --- a/internal/auditlog/store_mongodb.go +++ b/internal/auditlog/store_mongodb.go @@ -65,7 +65,7 @@ func NewMongoDBStore(database *mongo.Database, retentionDays int) (*MongoDBStore // Create indexes for common queries indexes := []mongo.IndexModel{ { - Keys: bson.D{{Key: "model", Value: 1}}, + Keys: bson.D{{Key: "requested_model", Value: 1}}, }, { Keys: bson.D{{Key: "status_code", Value: 1}}, @@ -73,6 +73,9 @@ func NewMongoDBStore(database *mongo.Database, retentionDays int) (*MongoDBStore { Keys: bson.D{{Key: "provider", Value: 1}}, }, + { + Keys: bson.D{{Key: "provider_name", Value: 1}}, + }, { Keys: bson.D{{Key: "execution_plan_version_id", Value: 1}}, }, diff --git a/internal/auditlog/store_postgresql.go b/internal/auditlog/store_postgresql.go index 0fa0f1a7a..0b3e18621 100644 --- a/internal/auditlog/store_postgresql.go +++ b/internal/auditlog/store_postgresql.go @@ -14,13 +14,13 @@ import ( ) const ( - auditLogInsertColumnCount = 20 + auditLogInsertColumnCount = 21 postgresMaxBindParameters = 65535 auditLogInsertMaxRowsPerQuery = postgresMaxBindParameters / auditLogInsertColumnCount ) const auditLogInsertPrefix = ` - INSERT INTO audit_logs (id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, + INSERT INTO audit_logs (id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data) VALUES ` @@ -56,9 +56,10 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore id UUID PRIMARY KEY, timestamp TIMESTAMPTZ NOT NULL, duration_ns BIGINT DEFAULT 0, - model TEXT, + requested_model TEXT, resolved_model TEXT, provider TEXT, + provider_name TEXT, alias_used BOOLEAN DEFAULT FALSE, execution_plan_version_id TEXT, cache_type TEXT, @@ -79,8 +80,14 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore return nil, fmt.Errorf("failed to create audit_logs table: %w", err) } + if err := renamePostgreSQLAuditColumn(ctx, pool, "audit_logs", "model", "requested_model"); err != nil { + return nil, fmt.Errorf("failed to rename audit_logs.model to requested_model: %w", err) + } + migrations := []string{ + "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS requested_model TEXT", "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS resolved_model TEXT", + "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS provider_name TEXT", "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS alias_used BOOLEAN DEFAULT FALSE", "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS execution_plan_version_id TEXT", "ALTER TABLE audit_logs ADD COLUMN IF NOT EXISTS cache_type TEXT", @@ -97,9 +104,11 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore // Create indexes for common queries indexes := []string{ "CREATE INDEX IF NOT EXISTS idx_audit_timestamp ON audit_logs(timestamp)", - "CREATE INDEX IF NOT EXISTS idx_audit_model ON audit_logs(model)", + "DROP INDEX IF EXISTS idx_audit_model", + "CREATE INDEX IF NOT EXISTS idx_audit_requested_model ON audit_logs(requested_model)", "CREATE INDEX IF NOT EXISTS idx_audit_status ON audit_logs(status_code)", "CREATE INDEX IF NOT EXISTS idx_audit_provider ON audit_logs(provider)", + "CREATE INDEX IF NOT EXISTS idx_audit_provider_name ON audit_logs(provider_name)", "CREATE INDEX IF NOT EXISTS idx_audit_execution_plan_version_id ON audit_logs(execution_plan_version_id)", "CREATE INDEX IF NOT EXISTS idx_audit_request_id ON audit_logs(request_id)", "CREATE INDEX IF NOT EXISTS idx_audit_auth_key_id ON audit_logs(auth_key_id)", @@ -222,9 +231,10 @@ func buildAuditLogInsert(entries []*LogEntry) (string, []any) { entry.ID, entry.Timestamp, entry.DurationNs, - entry.Model, + entry.RequestedModel, entry.ResolvedModel, entry.Provider, + entry.ProviderName, entry.AliasUsed, entry.ExecutionPlanVersionID, cacheTypeValue, @@ -246,6 +256,33 @@ func buildAuditLogInsert(entries []*LogEntry) (string, []any) { return builder.String(), args } +func renamePostgreSQLAuditColumn(ctx context.Context, pool *pgxpool.Pool, tableName, from, to string) error { + fromExists, err := postgresqlColumnExists(ctx, pool, tableName, from) + if err != nil || !fromExists { + return err + } + toExists, err := postgresqlColumnExists(ctx, pool, tableName, to) + if err != nil || toExists { + return err + } + _, err = pool.Exec(ctx, fmt.Sprintf("ALTER TABLE %s RENAME COLUMN %s TO %s", tableName, from, to)) + return err +} + +func postgresqlColumnExists(ctx context.Context, pool *pgxpool.Pool, tableName, columnName string) (bool, error) { + var exists bool + err := pool.QueryRow(ctx, ` + SELECT EXISTS ( + SELECT 1 + FROM information_schema.columns + WHERE table_schema = current_schema() + AND table_name = $1 + AND column_name = $2 + ) + `, tableName, columnName).Scan(&exists) + return exists, err +} + // Flush is a no-op for PostgreSQL as writes are synchronous. func (s *PostgreSQLStore) Flush(_ context.Context) error { return nil diff --git a/internal/auditlog/store_postgresql_test.go b/internal/auditlog/store_postgresql_test.go index 16d9f3161..35eb9f57f 100644 --- a/internal/auditlog/store_postgresql_test.go +++ b/internal/auditlog/store_postgresql_test.go @@ -11,97 +11,101 @@ func TestBuildAuditLogInsert(t *testing.T) { query, args := buildAuditLogInsert([]*LogEntry{ { - ID: "log-1", - Timestamp: now, - DurationNs: 1234, - Model: "gpt-4o-mini", - ResolvedModel: "gpt-4o-mini", - Provider: "openai", - AliasUsed: true, - CacheType: CacheTypeExact, - StatusCode: 200, - RequestID: "req-1", - AuthKeyID: "auth-key-1", - ClientIP: "127.0.0.1", - Method: "POST", - Path: "/v1/chat/completions", - UserPath: "/team/alpha", - Stream: true, - ErrorType: "", + ID: "log-1", + Timestamp: now, + DurationNs: 1234, + RequestedModel: "gpt-4o-mini", + ResolvedModel: "gpt-4o-mini", + Provider: "openai", + ProviderName: "primary-openai", + AliasUsed: true, + CacheType: CacheTypeExact, + StatusCode: 200, + RequestID: "req-1", + AuthKeyID: "auth-key-1", + ClientIP: "127.0.0.1", + Method: "POST", + Path: "/v1/chat/completions", + UserPath: "/team/alpha", + Stream: true, + ErrorType: "", Data: &LogData{ UserAgent: "test-agent", }, }, { - ID: "log-2", - Timestamp: now.Add(time.Second), - DurationNs: 5678, - Model: "gpt-4.1", - ResolvedModel: "gpt-4.1", - Provider: "openai", - AliasUsed: false, - StatusCode: 500, - RequestID: "req-2", - ClientIP: "10.0.0.1", - Method: "POST", - Path: "/v1/responses", - Stream: false, - ErrorType: "server_error", - Data: nil, + ID: "log-2", + Timestamp: now.Add(time.Second), + DurationNs: 5678, + RequestedModel: "gpt-4.1", + ResolvedModel: "gpt-4.1", + Provider: "openai", + AliasUsed: false, + StatusCode: 500, + RequestID: "req-2", + ClientIP: "10.0.0.1", + Method: "POST", + Path: "/v1/responses", + Stream: false, + ErrorType: "server_error", + Data: nil, }, }) normalized := strings.Join(strings.Fields(query), " ") - wantQuery := "INSERT INTO audit_logs (id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20), ($21, $22, $23, $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40) ON CONFLICT (id) DO NOTHING" + wantQuery := "INSERT INTO audit_logs (id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19, $20, $21), ($22, $23, $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42) ON CONFLICT (id) DO NOTHING" if normalized != wantQuery { t.Fatalf("query = %q, want %q", normalized, wantQuery) } - if got, want := len(args), 40; got != want { + if got, want := len(args), 42; got != want { t.Fatalf("len(args) = %d, want %d", got, want) } if got := args[0]; got != "log-1" { t.Fatalf("args[0] = %v, want log-1", got) } - if got := args[8]; got != CacheTypeExact { - t.Fatalf("args[8] = %v, want %q", got, CacheTypeExact) + if got := args[6]; got != "primary-openai" { + t.Fatalf("args[6] = %v, want primary-openai", got) } - if got, ok := args[11].(string); !ok || got != "auth-key-1" { - t.Fatalf("args[11] = (%T) %v, want (string) auth-key-1", args[11], args[11]) + if got := args[9]; got != CacheTypeExact { + t.Fatalf("args[9] = %v, want %q", got, CacheTypeExact) } - if got, ok := args[12].(string); !ok || got != "" { - t.Fatalf("args[12] = (%T) %v, want (string) \"\"", args[12], args[12]) + if got, ok := args[12].(string); !ok || got != "auth-key-1" { + t.Fatalf("args[12] = (%T) %v, want (string) auth-key-1", args[12], args[12]) } - if got, ok := args[15].(string); !ok || got != "/v1/chat/completions" { - t.Fatalf("args[15] = (%T) %v, want (string) /v1/chat/completions", args[15], args[15]) + if got, ok := args[13].(string); !ok || got != "" { + t.Fatalf("args[13] = (%T) %v, want (string) \"\"", args[13], args[13]) } - if got, ok := args[16].(string); !ok || got != "/team/alpha" { - t.Fatalf("args[16] = (%T) %v, want (string) /team/alpha", args[16], args[16]) + if got, ok := args[16].(string); !ok || got != "/v1/chat/completions" { + t.Fatalf("args[16] = (%T) %v, want (string) /v1/chat/completions", args[16], args[16]) } - if got := string(args[19].([]byte)); got != `{"user_agent":"test-agent"}` { - t.Fatalf("args[19] = %q, want %q", got, `{"user_agent":"test-agent"}`) + if got, ok := args[17].(string); !ok || got != "/team/alpha" { + t.Fatalf("args[17] = (%T) %v, want (string) /team/alpha", args[17], args[17]) } - if got := args[20]; got != "log-2" { - t.Fatalf("args[20] = %v, want log-2", got) + if got := string(args[20].([]byte)); got != `{"user_agent":"test-agent"}` { + t.Fatalf("args[20] = %q, want %q", got, `{"user_agent":"test-agent"}`) } - if got, ok := args[31].(string); !ok || got != "" { - t.Fatalf("args[31] = (%T) %v, want (string) \"\"", args[31], args[31]) + if got := args[21]; got != "log-2" { + t.Fatalf("args[21] = %v, want log-2", got) } - if got, ok := args[32].(string); !ok || got != "" { - t.Fatalf("args[32] = (%T) %v, want (string) \"\"", args[32], args[32]) + if got, ok := args[33].(string); !ok || got != "" { + t.Fatalf("args[33] = (%T) %v, want (string) \"\"", args[33], args[33]) } - if got := args[28]; got != nil { - t.Fatalf("args[28] = %v, want nil cache type", got) + if got, ok := args[34].(string); !ok || got != "" { + t.Fatalf("args[34] = (%T) %v, want (string) \"\"", args[34], args[34]) } - if got, ok := args[36].(string); !ok || got != "/" { - t.Fatalf("args[36] = (%T) %v, want (string) \"/\"", args[36], args[36]) + if got := args[30]; got != nil { + t.Fatalf("args[30] = %v, want nil cache type", got) } - dataJSON, ok := args[39].([]byte) + if got, ok := args[38].(string); !ok || got != "/" { + t.Fatalf("args[38] = (%T) %v, want (string) \"/\"", args[38], args[38]) + } + dataJSON, ok := args[41].([]byte) if !ok { - t.Fatalf("args[39] has type %T, want []byte", args[39]) + t.Fatalf("args[41] has type %T, want []byte", args[41]) } if dataJSON != nil { - t.Fatalf("args[39] = %v, want nil data", dataJSON) + t.Fatalf("args[41] = %v, want nil data", dataJSON) } } diff --git a/internal/auditlog/store_sqlite.go b/internal/auditlog/store_sqlite.go index 5ba4d2ae3..a1854f712 100644 --- a/internal/auditlog/store_sqlite.go +++ b/internal/auditlog/store_sqlite.go @@ -15,10 +15,12 @@ import ( // We chunk larger batches to avoid hitting this limit. const ( maxSQLiteParams = 999 - columnsPerEntry = 20 + columnsPerEntry = 21 maxEntriesPerBatch = maxSQLiteParams / columnsPerEntry // 49 entries ) +const sqliteAuditLogTable = "audit_logs" + // SQLiteStore implements LogStore for SQLite databases. type SQLiteStore struct { db *sql.DB @@ -41,9 +43,10 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { id TEXT PRIMARY KEY, timestamp DATETIME NOT NULL, duration_ns INTEGER DEFAULT 0, - model TEXT, + requested_model TEXT, resolved_model TEXT, provider TEXT, + provider_name TEXT, alias_used INTEGER DEFAULT 0, execution_plan_version_id TEXT, cache_type TEXT, @@ -64,8 +67,14 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { return nil, fmt.Errorf("failed to create audit_logs table: %w", err) } + if err := renameSQLiteAuditColumn(db, sqliteAuditLogTable, "model", "requested_model"); err != nil { + return nil, fmt.Errorf("failed to rename audit_logs.model to requested_model: %w", err) + } + migrations := []string{ + "ALTER TABLE audit_logs ADD COLUMN requested_model TEXT", "ALTER TABLE audit_logs ADD COLUMN resolved_model TEXT", + "ALTER TABLE audit_logs ADD COLUMN provider_name TEXT", "ALTER TABLE audit_logs ADD COLUMN alias_used INTEGER DEFAULT 0", "ALTER TABLE audit_logs ADD COLUMN execution_plan_version_id TEXT", "ALTER TABLE audit_logs ADD COLUMN cache_type TEXT", @@ -84,9 +93,11 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { // Create indexes for common queries indexes := []string{ "CREATE INDEX IF NOT EXISTS idx_audit_timestamp ON audit_logs(timestamp)", - "CREATE INDEX IF NOT EXISTS idx_audit_model ON audit_logs(model)", + "DROP INDEX IF EXISTS idx_audit_model", + "CREATE INDEX IF NOT EXISTS idx_audit_requested_model ON audit_logs(requested_model)", "CREATE INDEX IF NOT EXISTS idx_audit_status ON audit_logs(status_code)", "CREATE INDEX IF NOT EXISTS idx_audit_provider ON audit_logs(provider)", + "CREATE INDEX IF NOT EXISTS idx_audit_provider_name ON audit_logs(provider_name)", "CREATE INDEX IF NOT EXISTS idx_audit_execution_plan_version_id ON audit_logs(execution_plan_version_id)", "CREATE INDEX IF NOT EXISTS idx_audit_request_id ON audit_logs(request_id)", "CREATE INDEX IF NOT EXISTS idx_audit_auth_key_id ON audit_logs(auth_key_id)", @@ -134,7 +145,7 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*LogEntry) error values := make([]any, 0, len(chunk)*columnsPerEntry) for j, e := range chunk { - placeholders[j] = "(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)" + placeholders[j] = "(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)" dataJSON := marshalLogData(e.Data, e.ID) @@ -166,9 +177,10 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*LogEntry) error e.ID, e.Timestamp.UTC().Format(time.RFC3339Nano), e.DurationNs, - e.Model, + e.RequestedModel, e.ResolvedModel, e.Provider, + e.ProviderName, aliasUsedInt, e.ExecutionPlanVersionID, cacheTypeValue, @@ -186,7 +198,7 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*LogEntry) error ) } - query := `INSERT OR IGNORE INTO audit_logs (id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, cache_type, status_code, + query := `INSERT OR IGNORE INTO audit_logs (id, timestamp, duration_ns, requested_model, resolved_model, provider, provider_name, alias_used, execution_plan_version_id, cache_type, status_code, request_id, auth_key_id, auth_method, client_ip, method, path, user_path, stream, error_type, data) VALUES ` + strings.Join(placeholders, ",") @@ -234,3 +246,45 @@ func (s *SQLiteStore) cleanup() { slog.Info("cleaned up old audit logs", "deleted", rowsAffected) } } + +func renameSQLiteAuditColumn(db *sql.DB, tableName, from, to string) error { + if db == nil { + return nil + } + fromExists, err := sqliteColumnExists(db, tableName, from) + if err != nil || !fromExists { + return err + } + toExists, err := sqliteColumnExists(db, tableName, to) + if err != nil || toExists { + return err + } + _, err = db.Exec(fmt.Sprintf("ALTER TABLE %s RENAME COLUMN %s TO %s", tableName, from, to)) + return err +} + +func sqliteColumnExists(db *sql.DB, tableName, columnName string) (bool, error) { + rows, err := db.Query(fmt.Sprintf("PRAGMA table_info(%s)", tableName)) + if err != nil { + return false, err + } + defer rows.Close() + + for rows.Next() { + var ( + cid int + name string + columnType string + notNull int + dfltValue any + pk int + ) + if err := rows.Scan(&cid, &name, &columnType, ¬Null, &dfltValue, &pk); err != nil { + return false, err + } + if strings.EqualFold(name, columnName) { + return true, nil + } + } + return false, rows.Err() +} diff --git a/internal/auditlog/store_sqlite_test.go b/internal/auditlog/store_sqlite_test.go index d61961ce8..1f2f4a906 100644 --- a/internal/auditlog/store_sqlite_test.go +++ b/internal/auditlog/store_sqlite_test.go @@ -35,17 +35,17 @@ func TestSQLiteStore_WriteBatch_NullDataPreservation(t *testing.T) { // Create entries - one with nil Data, one with Data entries := []*LogEntry{ { - ID: "entry-nil-data", - Timestamp: time.Now(), - Model: "gpt-4", - Provider: "openai", - Data: nil, // This should become SQL NULL + ID: "entry-nil-data", + Timestamp: time.Now(), + RequestedModel: "gpt-4", + Provider: "openai", + Data: nil, // This should become SQL NULL }, { - ID: "entry-with-data", - Timestamp: time.Now(), - Model: "gpt-4", - Provider: "openai", + ID: "entry-with-data", + Timestamp: time.Now(), + RequestedModel: "gpt-4", + Provider: "openai", Data: &LogData{ UserAgent: "test-agent", }, @@ -104,11 +104,11 @@ func TestSQLiteStore_WriteBatch_Chunking(t *testing.T) { entries := make([]*LogEntry, numEntries) for i := range numEntries { entries[i] = &LogEntry{ - ID: fmt.Sprintf("entry-%03d", i), - Timestamp: time.Now(), - Model: "gpt-4", - Provider: "openai", - StatusCode: 200, + ID: fmt.Sprintf("entry-%03d", i), + Timestamp: time.Now(), + RequestedModel: "gpt-4", + Provider: "openai", + StatusCode: 200, } } @@ -183,9 +183,9 @@ func TestSQLiteStore_WriteBatch_ExactBatchBoundary(t *testing.T) { entries := make([]*LogEntry, numEntries) for i := range numEntries { entries[i] = &LogEntry{ - ID: fmt.Sprintf("exact-%03d", i), - Timestamp: time.Now(), - Model: "gpt-4", + ID: fmt.Sprintf("exact-%03d", i), + Timestamp: time.Now(), + RequestedModel: "gpt-4", } } @@ -205,9 +205,9 @@ func TestSQLiteStore_WriteBatch_ExactBatchBoundary(t *testing.T) { entries = make([]*LogEntry, maxEntriesPerBatch+1) for i := 0; i <= maxEntriesPerBatch; i++ { entries[i] = &LogEntry{ - ID: fmt.Sprintf("boundary-%03d", i), - Timestamp: time.Now(), - Model: "gpt-4", + ID: fmt.Sprintf("boundary-%03d", i), + Timestamp: time.Now(), + RequestedModel: "gpt-4", } } @@ -236,13 +236,13 @@ func TestSQLiteStore_WriteBatch_PersistsAliasFields(t *testing.T) { ctx := context.Background() entry := &LogEntry{ - ID: "alias-entry", - Timestamp: time.Now(), - Model: "anthropic/claude-opus-4-6", - ResolvedModel: "openai/gpt-5-nano", - Provider: "openai", - AliasUsed: true, - StatusCode: 200, + ID: "alias-entry", + Timestamp: time.Now(), + RequestedModel: "anthropic/claude-opus-4-6", + ResolvedModel: "openai/gpt-5-nano", + Provider: "openai", + AliasUsed: true, + StatusCode: 200, } if err := store.WriteBatch(ctx, []*LogEntry{entry}); err != nil { @@ -262,8 +262,8 @@ func TestSQLiteStore_WriteBatch_PersistsAliasFields(t *testing.T) { t.Fatal("expected log entry, got nil") return } - if logEntry.Model != entry.Model { - t.Fatalf("Model = %q, want %q", logEntry.Model, entry.Model) + if logEntry.RequestedModel != entry.RequestedModel { + t.Fatalf("RequestedModel = %q, want %q", logEntry.RequestedModel, entry.RequestedModel) } if logEntry.ResolvedModel != entry.ResolvedModel { t.Fatalf("ResolvedModel = %q, want %q", logEntry.ResolvedModel, entry.ResolvedModel) @@ -292,7 +292,7 @@ func TestSQLiteReader_AllowsNullExecutionPlanVersionID(t *testing.T) { now := time.Now().UTC().Format(time.RFC3339Nano) if _, err := db.Exec(` INSERT INTO audit_logs ( - id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, + id, timestamp, duration_ns, requested_model, resolved_model, provider, alias_used, execution_plan_version_id, status_code, request_id, client_ip, method, path, stream, error_type, data ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) `, @@ -358,7 +358,7 @@ func TestSQLiteReader_GetLogsFiltersByUserPathSubtree(t *testing.T) { now := time.Now().UTC().Format(time.RFC3339Nano) _, err = db.Exec(` INSERT INTO audit_logs ( - id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, + id, timestamp, duration_ns, requested_model, resolved_model, provider, alias_used, execution_plan_version_id, status_code, request_id, client_ip, method, path, user_path, stream, error_type, data ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), @@ -436,7 +436,7 @@ func TestSQLiteReader_GetLogsRootUserPathIncludesLegacyNullRows(t *testing.T) { now := time.Now().UTC().Format(time.RFC3339Nano) _, err = db.Exec(` INSERT INTO audit_logs ( - id, timestamp, duration_ns, model, resolved_model, provider, alias_used, execution_plan_version_id, + id, timestamp, duration_ns, requested_model, resolved_model, provider, alias_used, execution_plan_version_id, status_code, request_id, client_ip, method, path, user_path, stream, error_type, data ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?), @@ -509,17 +509,17 @@ func TestSQLiteStoreAndReader_PreserveCacheType(t *testing.T) { now := time.Now() if err := store.WriteBatch(ctx, []*LogEntry{ { - ID: "cache-exact", - Timestamp: now, - Model: "gpt-4", - Provider: "openai", - CacheType: CacheTypeExact, + ID: "cache-exact", + Timestamp: now, + RequestedModel: "gpt-4", + Provider: "openai", + CacheType: CacheTypeExact, }, { - ID: "cache-none", - Timestamp: now.Add(time.Second), - Model: "gpt-4", - Provider: "openai", + ID: "cache-none", + Timestamp: now.Add(time.Second), + RequestedModel: "gpt-4", + Provider: "openai", }, }); err != nil { t.Fatalf("WriteBatch failed: %v", err) diff --git a/internal/auditlog/stream_wrapper.go b/internal/auditlog/stream_wrapper.go index fe4e8567f..18848778c 100644 --- a/internal/auditlog/stream_wrapper.go +++ b/internal/auditlog/stream_wrapper.go @@ -84,7 +84,7 @@ func CreateStreamEntry(baseEntry *LogEntry) *LogEntry { ID: baseEntry.ID, Timestamp: baseEntry.Timestamp, DurationNs: baseEntry.DurationNs, - Model: baseEntry.Model, + RequestedModel: baseEntry.RequestedModel, ResolvedModel: baseEntry.ResolvedModel, Provider: baseEntry.Provider, AliasUsed: baseEntry.AliasUsed, diff --git a/internal/core/interfaces.go b/internal/core/interfaces.go index 4bfdbdc84..d9072bdca 100644 --- a/internal/core/interfaces.go +++ b/internal/core/interfaces.go @@ -103,6 +103,12 @@ type RoutableProvider interface { GetProviderType(model string) string } +// ProviderNameResolver is an optional interface for components that can map a +// routed model selector back to the concrete configured provider instance name. +type ProviderNameResolver interface { + GetProviderName(model string) string +} + // AvailabilityChecker is an optional interface for providers that need // to verify service availability before registration. type AvailabilityChecker interface { diff --git a/internal/core/request_model_resolution.go b/internal/core/request_model_resolution.go index fd493c83f..4b9bd83fa 100644 --- a/internal/core/request_model_resolution.go +++ b/internal/core/request_model_resolution.go @@ -6,6 +6,7 @@ type RequestModelResolution struct { Requested RequestedModelSelector ResolvedSelector ModelSelector ProviderType string + ProviderName string AliasApplied bool } diff --git a/internal/providers/config.go b/internal/providers/config.go index d5b083ec4..291c9f168 100644 --- a/internal/providers/config.go +++ b/internal/providers/config.go @@ -57,8 +57,9 @@ func applyProviderEnvVars(raw map[string]config.RawProviderConfig, discovery map continue } - existing, exists := result[providerType] - if exists { + targetKey, matched, ambiguous := findEnvOverlayTarget(result, providerType) + if matched { + existing := result[targetKey] if apiKey != "" { existing.APIKey = apiKey } @@ -70,7 +71,9 @@ func applyProviderEnvVars(raw map[string]config.RawProviderConfig, discovery map if apiVersion != "" { existing.APIVersion = apiVersion } - result[providerType] = existing + result[targetKey] = existing + } else if ambiguous { + continue } else { if spec.RequireBaseURL && explicitBaseURL == "" { continue @@ -87,6 +90,34 @@ func applyProviderEnvVars(raw map[string]config.RawProviderConfig, discovery map return result } +func findEnvOverlayTarget(raw map[string]config.RawProviderConfig, providerType string) (string, bool, bool) { + if existing, ok := raw[providerType]; ok && rawProviderMatchesType(existing, providerType) { + return providerType, true, false + } + + var matchedKey string + var matches int + for name, cfg := range raw { + if !rawProviderMatchesType(cfg, providerType) { + continue + } + matchedKey = name + matches++ + if matches > 1 { + return "", false, true + } + } + + if matches == 1 { + return matchedKey, true, false + } + return "", false, false +} + +func rawProviderMatchesType(cfg config.RawProviderConfig, providerType string) bool { + return strings.TrimSpace(cfg.Type) == strings.TrimSpace(providerType) +} + type providerEnvNames struct { APIKey string BaseURL string diff --git a/internal/providers/config_test.go b/internal/providers/config_test.go index 2174d8d29..2c18a954d 100644 --- a/internal/providers/config_test.go +++ b/internal/providers/config_test.go @@ -470,6 +470,49 @@ func TestApplyProviderEnvVars_EnvWinsOverYAML(t *testing.T) { } } +func TestApplyProviderEnvVars_SingleCustomNamedProviderUsesTypeEnvVars(t *testing.T) { + t.Setenv("OPENAI_API_KEY", "sk-env-key") + + raw := map[string]config.RawProviderConfig{ + "openai_name": {Type: "openai"}, + } + got := applyProviderEnvVars(raw, testDiscoveryConfigs) + + provider, exists := got["openai_name"] + if !exists { + t.Fatal("expected custom-named openai provider to be preserved") + } + if provider.APIKey != "sk-env-key" { + t.Errorf("APIKey = %q, want sk-env-key", provider.APIKey) + } + if provider.BaseURL != testDiscoveryConfigs["openai"].DefaultBaseURL { + t.Errorf("BaseURL = %q, want %q", provider.BaseURL, testDiscoveryConfigs["openai"].DefaultBaseURL) + } + if _, exists := got["openai"]; exists { + t.Fatal("expected no duplicate auto-discovered openai provider") + } +} + +func TestApplyProviderEnvVars_AmbiguousCustomNamedProvidersSkipTypeEnvOverlay(t *testing.T) { + t.Setenv("OPENAI_API_KEY", "sk-env-key") + + raw := map[string]config.RawProviderConfig{ + "openai-east": {Type: "openai", APIKey: "east-key", BaseURL: "https://east.example.com/v1"}, + "openai-west": {Type: "openai", APIKey: "west-key", BaseURL: "https://west.example.com/v1"}, + } + got := applyProviderEnvVars(raw, testDiscoveryConfigs) + + if got["openai-east"].APIKey != "east-key" { + t.Errorf("openai-east APIKey = %q, want east-key", got["openai-east"].APIKey) + } + if got["openai-west"].APIKey != "west-key" { + t.Errorf("openai-west APIKey = %q, want west-key", got["openai-west"].APIKey) + } + if _, exists := got["openai"]; exists { + t.Fatal("expected no duplicate auto-discovered openai provider when multiple YAML providers share the type") + } +} + func TestApplyProviderEnvVars_BaseURLEnvWinsOverYAML(t *testing.T) { t.Setenv("OPENAI_BASE_URL", "https://env-override.com") @@ -741,6 +784,33 @@ func TestResolveProviders_EmptyRaw_OnlyEnvVars(t *testing.T) { } } +func TestResolveProviders_SingleCustomNamedProviderDoesNotDuplicateTypeKey(t *testing.T) { + t.Setenv("OPENAI_API_KEY", "sk-openai") + + raw := map[string]config.RawProviderConfig{ + "openai_name": {Type: "openai"}, + } + + got, filteredRaw := resolveProviders(raw, globalResilience, testDiscoveryConfigs) + + provider, exists := got["openai_name"] + if !exists { + t.Fatal("expected openai_name provider in resolved providers") + } + if provider.APIKey != "sk-openai" { + t.Errorf("APIKey = %q, want sk-openai", provider.APIKey) + } + if provider.BaseURL != testDiscoveryConfigs["openai"].DefaultBaseURL { + t.Errorf("BaseURL = %q, want %q", provider.BaseURL, testDiscoveryConfigs["openai"].DefaultBaseURL) + } + if _, exists := got["openai"]; exists { + t.Fatal("expected no duplicate openai provider in resolved providers") + } + if _, exists := filteredRaw["openai"]; exists { + t.Fatal("expected no duplicate openai provider in filtered raw providers") + } +} + func TestResolveProviders_NoProvidersNoEnvVars(t *testing.T) { got, filteredRaw := resolveProviders(map[string]config.RawProviderConfig{}, globalResilience, testDiscoveryConfigs) if len(got) != 0 { diff --git a/internal/providers/registry.go b/internal/providers/registry.go index 703ac7727..c6c1e50f5 100644 --- a/internal/providers/registry.go +++ b/internal/providers/registry.go @@ -636,6 +636,30 @@ func (r *ModelRegistry) GetProviderType(model string) string { return "" } +// GetProviderName returns the concrete configured provider instance name for +// the given model selector. Returns empty string if the model is not found. +func (r *ModelRegistry) GetProviderName(model string) string { + r.mu.RLock() + defer r.mu.RUnlock() + + providerName, modelID := splitModelSelector(model) + if providerName != "" { + if providerModels, ok := r.modelsByProvider[providerName]; ok { + if info, exists := providerModels[modelID]; exists { + return strings.TrimSpace(info.ProviderName) + } + } + if r.hasConfiguredProviderNameLocked(providerName) { + return "" + } + } + + if info, ok := r.models[model]; ok { + return strings.TrimSpace(info.ProviderName) + } + return "" +} + // ProviderByType returns the first registered provider for the given provider type. // This lookup is independent of discovered models so provider-typed routes keep // working even when a provider currently exposes zero models. @@ -678,6 +702,22 @@ func (r *ModelRegistry) ProviderTypes() []string { return result } +// ProviderNames returns the configured provider instance names in registration order. +func (r *ModelRegistry) ProviderNames() []string { + r.mu.RLock() + defer r.mu.RUnlock() + + result := make([]string, 0, len(r.providers)) + for _, provider := range r.providers { + providerName := strings.TrimSpace(r.providerNames[provider]) + if providerName == "" { + continue + } + result = append(result, providerName) + } + return result +} + func splitModelSelector(model string) (providerName, modelID string) { model = strings.TrimSpace(model) if model == "" { diff --git a/internal/providers/router.go b/internal/providers/router.go index 20db2e95a..babf7390a 100644 --- a/internal/providers/router.go +++ b/internal/providers/router.go @@ -35,10 +35,18 @@ type providerTypeLister interface { ProviderTypes() []string } +type providerNameLister interface { + ProviderNames() []string +} + type publicModelLister interface { ListPublicModels() []core.Model } +type modelWithProviderLister interface { + ListModelsWithProvider() []ModelWithProvider +} + func registryUnavailableError(err error) error { return core.NewProviderError("", http.StatusServiceUnavailable, err.Error(), err) } @@ -64,14 +72,139 @@ func (r *Router) checkReady() error { return nil } -// resolveProvider validates readiness, parses the model selector, and finds the target provider. -func (r *Router) resolveProvider(model, providerHint string) (core.Provider, core.ModelSelector, error) { +// ResolveModel canonicalizes a requested selector into the concrete +// provider-name-qualified selector used for execution. +// +// Resolution precedence is: +// 1. configured provider name + model ID +// 2. provider type + model ID +// 3. raw slash-shaped model ID (only when provider was not explicit) +// 4. default normalization fallback +func (r *Router) ResolveModel(requested core.RequestedModelSelector) (core.ModelSelector, bool, error) { if err := r.checkReady(); err != nil { - return nil, core.ModelSelector{}, registryUnavailableError(err) + return core.ModelSelector{}, false, registryUnavailableError(err) } - selector, err := core.ParseModelSelector(model, providerHint) + + requested = core.NewRequestedModelSelector(requested.Model, requested.ProviderHint) + selector, err := requested.Normalize() if err != nil { - return nil, core.ModelSelector{}, core.NewInvalidRequestError(err.Error(), err) + return core.ModelSelector{}, false, core.NewInvalidRequestError(err.Error(), err) + } + + resolved := selector + if selector.Provider == "" { + if concrete, ok := r.resolveUnqualifiedSelector(selector); ok { + resolved = concrete + } + } else if concrete, ok := r.resolveQualifiedSelector(requested, selector); ok { + resolved = concrete + } + + return resolved, resolved.QualifiedModel() != selector.QualifiedModel(), nil +} + +func (r *Router) resolveUnqualifiedSelector(selector core.ModelSelector) (core.ModelSelector, bool) { + if selector.Provider != "" || strings.TrimSpace(selector.Model) == "" { + return core.ModelSelector{}, false + } + + named, ok := r.lookup.(core.ProviderNameResolver) + if !ok { + return core.ModelSelector{}, false + } + providerName := strings.TrimSpace(named.GetProviderName(selector.Model)) + if providerName == "" { + return core.ModelSelector{}, false + } + return core.ModelSelector{Provider: providerName, Model: selector.Model}, true +} + +func (r *Router) resolveQualifiedSelector(requested core.RequestedModelSelector, selector core.ModelSelector) (core.ModelSelector, bool) { + models, ok := r.lookup.(modelWithProviderLister) + if !ok { + return core.ModelSelector{}, false + } + + providerSegment := strings.TrimSpace(selector.Provider) + modelID := strings.TrimSpace(selector.Model) + if providerSegment == "" || modelID == "" { + return core.ModelSelector{}, false + } + + entries := models.ListModelsWithProvider() + + for _, entry := range entries { + if strings.TrimSpace(entry.ProviderName) != providerSegment { + continue + } + if strings.TrimSpace(entry.Model.ID) != modelID { + continue + } + return core.ModelSelector{Provider: entry.ProviderName, Model: entry.Model.ID}, true + } + + for _, entry := range entries { + if strings.TrimSpace(entry.ProviderType) != providerSegment { + continue + } + if strings.TrimSpace(entry.Model.ID) != modelID { + continue + } + return core.ModelSelector{Provider: entry.ProviderName, Model: entry.Model.ID}, true + } + + if requested.ExplicitProvider { + return core.ModelSelector{}, false + } + if r.hasConfiguredProviderName(providerSegment) { + return core.ModelSelector{}, false + } + if r.providerByTypeRegistry(providerSegment) != nil { + return core.ModelSelector{}, false + } + + rawModelID := strings.TrimSpace(requested.Model) + if rawModelID == "" { + return core.ModelSelector{}, false + } + for _, entry := range entries { + if strings.TrimSpace(entry.Model.ID) != rawModelID { + continue + } + return core.ModelSelector{Provider: entry.ProviderName, Model: entry.Model.ID}, true + } + + return core.ModelSelector{}, false +} + +func (r *Router) hasConfiguredProviderName(providerName string) bool { + providerName = strings.TrimSpace(providerName) + if providerName == "" { + return false + } + if named, ok := r.lookup.(providerNameLister); ok { + for _, candidate := range named.ProviderNames() { + if strings.TrimSpace(candidate) == providerName { + return true + } + } + return false + } + if models, ok := r.lookup.(modelWithProviderLister); ok { + for _, entry := range models.ListModelsWithProvider() { + if strings.TrimSpace(entry.ProviderName) == providerName { + return true + } + } + } + return false +} + +// resolveProvider validates readiness, parses the model selector, and finds the target provider. +func (r *Router) resolveProvider(model, providerHint string) (core.Provider, core.ModelSelector, error) { + selector, _, err := r.ResolveModel(core.NewRequestedModelSelector(model, providerHint)) + if err != nil { + return nil, core.ModelSelector{}, err } lookupModel := selector.QualifiedModel() p := r.lookup.GetProvider(lookupModel) @@ -259,10 +392,11 @@ func callEmbeddings(ctx context.Context, provider core.Provider, req *core.Embed // Supports returns true if any provider supports the given model. // Returns false if the lookup has no models loaded. func (r *Router) Supports(model string) bool { - if r.lookup.ModelCount() == 0 { + selector, _, err := r.ResolveModel(core.NewRequestedModelSelector(model, "")) + if err != nil { return false } - return r.lookup.Supports(model) + return r.lookup.Supports(selector.QualifiedModel()) } // ModelCount returns the number of models currently loaded into the router lookup. @@ -374,7 +508,30 @@ func (r *Router) Embeddings(ctx context.Context, req *core.EmbeddingRequest) (*c // GetProviderType returns the provider type string for the given model. // Returns empty string if the model is not found. func (r *Router) GetProviderType(model string) string { - return r.lookup.GetProviderType(model) + selector, _, err := r.ResolveModel(core.NewRequestedModelSelector(model, "")) + if err != nil { + return "" + } + return r.lookup.GetProviderType(selector.QualifiedModel()) +} + +// GetProviderName returns the concrete configured provider instance name for +// the given model selector. Returns empty string when unavailable. +func (r *Router) GetProviderName(model string) string { + selector, _, err := r.ResolveModel(core.NewRequestedModelSelector(model, "")) + if err != nil { + return "" + } + if !r.lookup.Supports(selector.QualifiedModel()) { + return "" + } + if selector.Provider != "" { + return selector.Provider + } + if named, ok := r.lookup.(core.ProviderNameResolver); ok { + return named.GetProviderName(selector.QualifiedModel()) + } + return "" } func (r *Router) providerByType(providerType string) core.Provider { diff --git a/internal/providers/router_test.go b/internal/providers/router_test.go index 8003bcdb6..4d63ff61f 100644 --- a/internal/providers/router_test.go +++ b/internal/providers/router_test.go @@ -31,6 +31,42 @@ func newMockLookup() *mockModelLookup { } } +type registryModelEntry struct { + provider core.Provider + providerName string + providerType string + modelID string +} + +func newTestRegistryWithModels(entries ...registryModelEntry) *ModelRegistry { + registry := NewModelRegistry() + for _, entry := range entries { + registry.RegisterProviderWithNameAndType(entry.provider, entry.providerName, entry.providerType) + } + + registry.models = make(map[string]*ModelInfo) + registry.modelsByProvider = make(map[string]map[string]*ModelInfo) + for _, entry := range entries { + info := &ModelInfo{ + Model: core.Model{ + ID: entry.modelID, + Object: "model", + }, + Provider: entry.provider, + ProviderName: entry.providerName, + ProviderType: entry.providerType, + } + if _, ok := registry.modelsByProvider[entry.providerName]; !ok { + registry.modelsByProvider[entry.providerName] = make(map[string]*ModelInfo) + } + registry.modelsByProvider[entry.providerName][entry.modelID] = info + if _, exists := registry.models[entry.modelID]; !exists { + registry.models[entry.modelID] = info + } + } + return registry +} + func (m *mockModelLookup) addModel(model string, provider core.Provider, providerType string) { m.models[model] = provider m.providerTypes[model] = providerType @@ -482,6 +518,91 @@ func TestRouterChatCompletion_PrefixedModelSelector(t *testing.T) { } } +func TestRouterChatCompletion_PrefersProviderTypeSelectorOverRawSlashModel(t *testing.T) { + openAIResp := &core.ChatResponse{ID: "openai-test", Model: "gpt-5-nano"} + openRouterResp := &core.ChatResponse{ID: "openrouter", Model: "openai/gpt-5-nano"} + openAI := &mockProvider{name: "openai_test", chatResponse: openAIResp} + openRouter := &mockProvider{name: "openrouter", chatResponse: openRouterResp} + + registry := newTestRegistryWithModels( + registryModelEntry{ + provider: openAI, + providerName: "openai_test", + providerType: "openai", + modelID: "gpt-5-nano", + }, + registryModelEntry{ + provider: openRouter, + providerName: "openrouter", + providerType: "openrouter", + modelID: "openai/gpt-5-nano", + }, + ) + + router, _ := NewRouter(registry) + + resp, err := router.ChatCompletion(context.Background(), &core.ChatRequest{Model: "openai/gpt-5-nano"}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.ID != "openai-test" { + t.Fatalf("expected openai_test response, got %q", resp.ID) + } + if openAI.lastChatReq == nil || openAI.lastChatReq.Model != "gpt-5-nano" { + t.Fatalf("expected openai provider to receive raw model gpt-5-nano, got %#v", openAI.lastChatReq) + } + if openAI.lastChatReq.Provider != "" { + t.Fatalf("expected provider field to be stripped upstream, got %q", openAI.lastChatReq.Provider) + } + if openRouter.lastChatReq != nil { + t.Fatalf("expected openrouter provider to be bypassed, got %#v", openRouter.lastChatReq) + } + if got := router.GetProviderType("openai/gpt-5-nano"); got != "openai" { + t.Fatalf("GetProviderType() = %q, want %q", got, "openai") + } + if got := router.GetProviderName("openai/gpt-5-nano"); got != "openai_test" { + t.Fatalf("GetProviderName() = %q, want %q", got, "openai_test") + } +} + +func TestRouterChatCompletion_ProviderQualifiedRawSlashModelStillWorks(t *testing.T) { + openAIResp := &core.ChatResponse{ID: "openai-test", Model: "gpt-5-nano"} + openRouterResp := &core.ChatResponse{ID: "openrouter", Model: "openai/gpt-5-nano"} + openAI := &mockProvider{name: "openai_test", chatResponse: openAIResp} + openRouter := &mockProvider{name: "openrouter", chatResponse: openRouterResp} + + registry := newTestRegistryWithModels( + registryModelEntry{ + provider: openAI, + providerName: "openai_test", + providerType: "openai", + modelID: "gpt-5-nano", + }, + registryModelEntry{ + provider: openRouter, + providerName: "openrouter", + providerType: "openrouter", + modelID: "openai/gpt-5-nano", + }, + ) + + router, _ := NewRouter(registry) + + resp, err := router.ChatCompletion(context.Background(), &core.ChatRequest{Model: "openrouter/openai/gpt-5-nano"}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.ID != "openrouter" { + t.Fatalf("expected openrouter response, got %q", resp.ID) + } + if openRouter.lastChatReq == nil || openRouter.lastChatReq.Model != "openai/gpt-5-nano" { + t.Fatalf("expected openrouter provider to receive raw slash model, got %#v", openRouter.lastChatReq) + } + if openRouter.lastChatReq.Provider != "" { + t.Fatalf("expected provider field to be stripped upstream, got %q", openRouter.lastChatReq.Provider) + } +} + func TestRouterChatCompletion_ExplicitProviderKeepsSlashModelRaw(t *testing.T) { groqResp := &core.ChatResponse{ID: "groq", Model: "openai/gpt-oss-120b"} groq := &mockProvider{name: "groq", chatResponse: groqResp} diff --git a/internal/responsecache/usage_hit.go b/internal/responsecache/usage_hit.go index 4d22c6290..784e79547 100644 --- a/internal/responsecache/usage_hit.go +++ b/internal/responsecache/usage_hit.go @@ -28,10 +28,12 @@ func newUsageHitRecorder(logger usage.LoggerInterface, pricingResolver usage.Pri model := "" provider := "" + providerName := "" if plan != nil { provider = strings.TrimSpace(plan.ProviderType) if plan.Resolution != nil { model = strings.TrimSpace(plan.Resolution.ResolvedSelector.Model) + providerName = strings.TrimSpace(plan.Resolution.ProviderName) } } if provider == "" { @@ -54,6 +56,7 @@ func newUsageHitRecorder(logger usage.LoggerInterface, pricingResolver usage.Pri if entry == nil { return } + entry.ProviderName = providerName entry.UserPath = core.UserPathFromContext(ctx) logger.Write(entry) } diff --git a/internal/server/execution_plan_helpers_test.go b/internal/server/execution_plan_helpers_test.go index 073aab222..a413c4f81 100644 --- a/internal/server/execution_plan_helpers_test.go +++ b/internal/server/execution_plan_helpers_test.go @@ -57,7 +57,7 @@ func TestEnsureTranslatedRequestPlan_CompletesPartialPlanFromDecodedSelector(t * assert.Equal(t, "gpt-4o-mini", storedPlan.Resolution.ResolvedSelector.Model) } } - assert.Equal(t, "gpt-4o-mini", entry.Model) + assert.Equal(t, "gpt-4o-mini", entry.RequestedModel) assert.Equal(t, "gpt-4o-mini", entry.ResolvedModel) assert.Equal(t, "mock", entry.Provider) } diff --git a/internal/server/handlers_test.go b/internal/server/handlers_test.go index 7497b96a7..88aabf6b7 100644 --- a/internal/server/handlers_test.go +++ b/internal/server/handlers_test.go @@ -306,6 +306,7 @@ type mockProvider struct { streamData string supportedModels []string providerTypes map[string]string + providerNames map[string]string batchCreateResponse *core.BatchResponse batchCreateHints map[string]string @@ -401,6 +402,25 @@ func (m *mockProvider) GetProviderType(model string) string { return "" } +func (m *mockProvider) GetProviderName(model string) string { + selector, err := core.ParseModelSelector(model, "") + if err == nil && selector.Provider != "" { + if m.providerNames != nil { + if providerName, ok := m.providerNames[selector.QualifiedModel()]; ok { + return providerName + } + } + model = selector.Model + } + + if m.providerNames != nil { + if providerName, ok := m.providerNames[model]; ok { + return providerName + } + } + return "" +} + func (m *mockProvider) NativeFileProviderTypes() []string { seen := make(map[string]struct{}) result := make([]string, 0, len(m.providerTypes)) @@ -2041,6 +2061,53 @@ func TestChatCompletionStreaming_FastPathUsesPassthroughForOpenAICompatibleProvi } } +func TestChatCompletionStreaming_FastPathUsageCarriesResolvedProviderName(t *testing.T) { + streamData := "data: {\"id\":\"chatcmpl-123\",\"model\":\"gpt-4o-mini\",\"usage\":{\"prompt_tokens\":7,\"completion_tokens\":3,\"total_tokens\":10}}\n\ndata: [DONE]\n\n" + usageLog := &collectingUsageLogger{ + config: usage.Config{Enabled: true}, + } + mock := &mockProvider{ + supportedModels: []string{"gpt-4o-mini"}, + providerTypes: map[string]string{ + "gpt-4o-mini": "openai", + }, + providerNames: map[string]string{ + "gpt-4o-mini": "openai_test", + }, + passthroughResponse: &core.PassthroughResponse{ + StatusCode: http.StatusOK, + Headers: map[string][]string{ + "Content-Type": {"text/event-stream"}, + }, + Body: io.NopCloser(strings.NewReader(streamData)), + }, + } + + e := echo.New() + handler := NewHandler(mock, nil, usageLog, nil) + + reqBody := `{"model":"gpt-4o-mini","stream":true,"messages":[{"role":"user","content":"Hi"}]}` + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(reqBody)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + c := e.NewContext(req, rec) + + err := handler.ChatCompletion(c) + if err != nil { + t.Fatalf("handler returned error: %v", err) + } + + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK) + } + if len(usageLog.entries) != 1 { + t.Fatalf("usage entries = %d, want 1", len(usageLog.entries)) + } + if got := usageLog.entries[0].ProviderName; got != "openai_test" { + t.Fatalf("ProviderName = %q, want openai_test", got) + } +} + func TestChatCompletionStreaming_FastPathSkipsQualifiedModelRewrite(t *testing.T) { streamData := "data: {\"id\":\"chatcmpl-123\",\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\ndata: [DONE]\n\n" provider := &capturingProvider{ @@ -2154,7 +2221,7 @@ func TestHandleStreamingResponse_FlushesEachChunk(t *testing.T) { }, } - err := handler.translatedInference().handleStreamingResponse(c, nil, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { + err := handler.translatedInference().handleStreamingResponse(c, nil, "gpt-4o-mini", "openai", "primary-openai", func() (io.ReadCloser, error) { return stream, nil }) if err != nil { @@ -2246,7 +2313,7 @@ func TestHandleStreamingResponse_RecordsStreamingError(t *testing.T) { Data: &auditlog.LogData{}, }) - err := handler.translatedInference().handleStreamingResponse(c, nil, "gpt-4o-mini", "openai", func() (io.ReadCloser, error) { + err := handler.translatedInference().handleStreamingResponse(c, nil, "gpt-4o-mini", "openai", "primary-openai", func() (io.ReadCloser, error) { return &erroringReadCloser{ data: []byte("data: {\"id\":\"1\"}\n\n"), err: expectedErr, @@ -5550,8 +5617,8 @@ func TestProviderPassthrough_UsesPassthroughModelForAuditEntry(t *testing.T) { if rec.Code != http.StatusOK { t.Fatalf("status = %d, want 200", rec.Code) } - if entry.Model != "gpt-5-mini" { - t.Fatalf("audit entry model = %q, want gpt-5-mini", entry.Model) + if entry.RequestedModel != "gpt-5-mini" { + t.Fatalf("audit entry requested model = %q, want gpt-5-mini", entry.RequestedModel) } if entry.Provider != "openai" { t.Fatalf("audit entry provider = %q, want openai", entry.Provider) diff --git a/internal/server/internal_chat_completion_executor.go b/internal/server/internal_chat_completion_executor.go index 29f4a17a1..6f951992d 100644 --- a/internal/server/internal_chat_completion_executor.go +++ b/internal/server/internal_chat_completion_executor.go @@ -81,8 +81,10 @@ func (e *InternalChatCompletionExecutor) ChatCompletion(ctx context.Context, req entry := e.newAuditEntry(ctx, requestID, requested) var plan *core.ExecutionPlan var cacheType string + var providerType string + var providerName string defer func() { - e.finishAuditEntry(ctx, entry, start, plan, req, resp, err, cacheType) + e.finishAuditEntry(ctx, entry, start, plan, req, resp, err, cacheType, providerType, providerName) }() resolution, err := resolveRequestModel(e.provider, e.modelResolver, requested) @@ -102,13 +104,13 @@ func (e *InternalChatCompletionExecutor) ChatCompletion(ctx context.Context, req ctx = e.service.withCacheRequestContext(ctx, plan) execReq := cloneChatRequestForSelector(req, resolution.ResolvedSelector) - resp, providerType, _, cacheType, err := e.executeChatCompletion(ctx, plan, execReq) + resp, providerType, providerName, _, cacheType, err = e.executeChatCompletion(ctx, plan, execReq) if err != nil { return nil, err } if cacheType == "" { - e.service.logUsage(ctx, plan, resp.Model, providerType, func(pricing *core.ModelPricing) *usage.UsageEntry { + e.service.logUsage(ctx, plan, resp.Model, providerType, providerName, func(pricing *core.ModelPricing) *usage.UsageEntry { return usage.ExtractFromChatResponse(resp, requestID, providerType, "/v1/chat/completions", pricing) }) } @@ -119,29 +121,30 @@ func (e *InternalChatCompletionExecutor) executeChatCompletion( ctx context.Context, plan *core.ExecutionPlan, req *core.ChatRequest, -) (*core.ChatResponse, string, bool, string, error) { +) (*core.ChatResponse, string, string, bool, string, error) { if e.service.responseCache == nil || (plan != nil && !plan.CacheEnabled()) { - resp, providerType, usedFallback, err := e.service.executeChatCompletion(ctx, plan, req) - return resp, providerType, usedFallback, "", err + resp, providerType, providerName, usedFallback, err := e.service.executeChatCompletion(ctx, plan, req) + return resp, providerType, providerName, usedFallback, "", err } body, err := marshalRequestBody(req) if err != nil { - resp, providerType, usedFallback, execErr := e.service.executeChatCompletion(ctx, plan, req) + resp, providerType, providerName, usedFallback, execErr := e.service.executeChatCompletion(ctx, plan, req) if execErr != nil { - return nil, "", false, "", execErr + return nil, "", "", false, "", execErr } - return resp, providerType, usedFallback, "", nil + return resp, providerType, providerName, usedFallback, "", nil } var ( resp *core.ChatResponse providerType string + providerName string usedFallback bool ) result, err := e.service.responseCache.HandleInternalRequest(ctx, http.MethodPost, "/v1/chat/completions", body, func(c *echo.Context) error { var execErr error - resp, providerType, usedFallback, execErr = e.service.executeChatCompletion(c.Request().Context(), plan, req) + resp, providerType, providerName, usedFallback, execErr = e.service.executeChatCompletion(c.Request().Context(), plan, req) if execErr != nil { return execErr } @@ -151,16 +154,16 @@ func (e *InternalChatCompletionExecutor) executeChatCompletion( return c.JSON(http.StatusOK, resp) }) if err != nil { - return nil, "", false, "", err + return nil, "", "", false, "", err } if result != nil && result.CacheType != "" { var cached core.ChatResponse if err := json.Unmarshal(result.Body, &cached); err != nil { - return nil, "", false, "", err + return nil, "", "", false, "", err } - return &cached, plan.ProviderType, false, result.CacheType, nil + return &cached, plan.ProviderType, providerNameFromPlan(plan), false, result.CacheType, nil } - return resp, providerType, usedFallback, "", nil + return resp, providerType, providerName, usedFallback, "", nil } func (e *InternalChatCompletionExecutor) newAuditEntry( @@ -187,7 +190,7 @@ func (e *InternalChatCompletionExecutor) newAuditEntry( Data: &auditlog.LogData{}, } if requestedModel := requested.RequestedQualifiedModel(); requestedModel != "" { - entry.Model = requestedModel + entry.RequestedModel = requestedModel } return entry } @@ -201,6 +204,8 @@ func (e *InternalChatCompletionExecutor) finishAuditEntry( resp *core.ChatResponse, err error, cacheType string, + providerType string, + providerName string, ) { if entry == nil || e.logger == nil || !e.logger.Config().Enabled { return @@ -208,6 +213,7 @@ func (e *InternalChatCompletionExecutor) finishAuditEntry( entry.DurationNs = time.Since(start).Nanoseconds() auditlog.EnrichLogEntryWithExecutionPlan(entry, plan) + auditlog.EnrichLogEntryWithResolvedRoute(entry, qualifyExecutedModel(plan, chatResponseModel(resp), providerName), providerType, providerName) auditlog.EnrichLogEntryWithRequestContext(entry, ctx) if plan != nil && !plan.AuditEnabled() { return @@ -240,3 +246,10 @@ func (e *InternalChatCompletionExecutor) finishAuditEntry( e.logger.Write(entry) } + +func chatResponseModel(resp *core.ChatResponse) string { + if resp == nil { + return "" + } + return resp.Model +} diff --git a/internal/server/internal_chat_completion_executor_test.go b/internal/server/internal_chat_completion_executor_test.go index 09784bf01..944ab9cb6 100644 --- a/internal/server/internal_chat_completion_executor_test.go +++ b/internal/server/internal_chat_completion_executor_test.go @@ -197,8 +197,8 @@ func TestInternalChatCompletionExecutor_DoesNotReuseParentExecutionPlanResolutio t.Fatalf("audit entries = %d, want 1", len(logger.entries)) } entry := logger.entries[0] - if entry.Model != "openai/gpt-4o-mini" { - t.Fatalf("audit requested model = %q, want openai/gpt-4o-mini", entry.Model) + if entry.RequestedModel != "openai/gpt-4o-mini" { + t.Fatalf("audit requested model = %q, want openai/gpt-4o-mini", entry.RequestedModel) } if entry.ResolvedModel != "openai/gpt-4o-mini" { t.Fatalf("audit resolved model = %q, want openai/gpt-4o-mini", entry.ResolvedModel) diff --git a/internal/server/model_validation_test.go b/internal/server/model_validation_test.go index f65c9e23b..24e4e5529 100644 --- a/internal/server/model_validation_test.go +++ b/internal/server/model_validation_test.go @@ -705,7 +705,7 @@ func TestModelValidation_EnrichesAuditEntryWithRequestedModelOnResolutionError(t assert.False(t, handlerCalled) assert.Equal(t, http.StatusBadRequest, rec.Code) assert.Contains(t, rec.Body.String(), "unsupported model: smart") - assert.Equal(t, "smart", entry.Model) + assert.Equal(t, "smart", entry.RequestedModel) assert.Equal(t, "", entry.ResolvedModel) assert.Equal(t, "", entry.Provider) assert.Equal(t, "invalid_request_error", entry.ErrorType) diff --git a/internal/server/native_batch_support.go b/internal/server/native_batch_support.go index 2f059bae3..a42a400df 100644 --- a/internal/server/native_batch_support.go +++ b/internal/server/native_batch_support.go @@ -59,22 +59,15 @@ func determineBatchExecutionSelection(provider core.RoutableProvider, resolver R commonModel string hasCommonModel = true ) - resolver = effectiveRequestModelResolver(provider, resolver) for i, item := range req.Requests { requested, err := core.BatchItemRequestedModelSelector(req.Endpoint, item) if err != nil { return batchExecutionSelection{}, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: %s", i, err.Error()), err) } - resolvedSelector, err := requested.Normalize() + resolvedSelector, _, err := resolveExecutionSelector(provider, resolver, requested) if err != nil { return batchExecutionSelection{}, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: %s", i, err.Error()), err) } - if resolver != nil { - resolvedSelector, _, err = resolver.ResolveModel(requested) - if err != nil { - return batchExecutionSelection{}, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: %s", i, err.Error()), err) - } - } model := resolvedSelector.QualifiedModel() if model == "" { return batchExecutionSelection{}, core.NewInvalidRequestError(fmt.Sprintf("batch item %d: model is required", i), nil) diff --git a/internal/server/passthrough_service.go b/internal/server/passthrough_service.go index 4a19a2d10..d151bb691 100644 --- a/internal/server/passthrough_service.go +++ b/internal/server/passthrough_service.go @@ -49,5 +49,5 @@ func (s *passthroughService) ProviderPassthrough(c *echo.Context) error { } else { auditlog.EnrichEntry(c, info.Model, providerType) } - return s.proxyPassthroughResponse(c, providerType, endpoint, info, resp) + return s.proxyPassthroughResponse(c, providerType, providerNameFromPlan(plan), endpoint, info, resp) } diff --git a/internal/server/passthrough_support.go b/internal/server/passthrough_support.go index 3eab9dffa..e360d1e13 100644 --- a/internal/server/passthrough_support.go +++ b/internal/server/passthrough_support.go @@ -226,7 +226,7 @@ func passthroughAuditPath(c *echo.Context, providerType, endpoint string, info * return passthroughStreamAuditPath("", providerType, endpoint) } -func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerType, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse) error { +func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerType, providerName, endpoint string, info *core.PassthroughRouteInfo, resp *core.PassthroughResponse) error { if resp == nil || resp.Body == nil { return handleError(c, core.NewProviderError(providerType, http.StatusBadGateway, "provider returned empty passthrough response", nil)) } @@ -282,6 +282,7 @@ func (s *passthroughService) proxyPassthroughResponse(c *echo.Context, providerT } if s.usageLogger != nil && s.usageLogger.Config().Enabled && (plan == nil || plan.UsageEnabled()) { if observer := usage.NewStreamUsageObserver(s.usageLogger, model, providerType, requestID, usagePath, s.pricingResolver, core.UserPathFromContext(c.Request().Context())); observer != nil { + observer.SetProviderName(providerName) observers = append(observers, observer) } } diff --git a/internal/server/request_model_resolution.go b/internal/server/request_model_resolution.go index 6728a6f1d..5c8e55f36 100644 --- a/internal/server/request_model_resolution.go +++ b/internal/server/request_model_resolution.go @@ -21,33 +21,32 @@ type RequestFallbackResolver interface { ResolveFallbacks(resolution *core.RequestModelResolution, op core.Operation) []core.ModelSelector } -func effectiveRequestModelResolver(provider core.RoutableProvider, resolver RequestModelResolver) RequestModelResolver { - if resolver != nil { - return resolver +func resolvedProviderName(provider core.RoutableProvider, selector core.ModelSelector, fallback string) string { + fallback = strings.TrimSpace(fallback) + if provider == nil { + return fallback } - if providerResolver, ok := provider.(RequestModelResolver); ok { - return providerResolver + if named, ok := provider.(core.ProviderNameResolver); ok { + if providerName := strings.TrimSpace(named.GetProviderName(selector.QualifiedModel())); providerName != "" { + return providerName + } } - return nil + return fallback } func resolveRequestModel(provider core.RoutableProvider, resolver RequestModelResolver, requested core.RequestedModelSelector) (*core.RequestModelResolution, error) { requested = core.NewRequestedModelSelector(requested.Model, requested.ProviderHint) - var ( - resolvedSelector core.ModelSelector - aliasApplied bool - err error - ) - - if effectiveResolver := effectiveRequestModelResolver(provider, resolver); effectiveResolver != nil { - resolvedSelector, aliasApplied, err = effectiveResolver.ResolveModel(requested) - } else { - resolvedSelector, err = requested.Normalize() - } + resolvedSelector, aliasApplied, err := resolveExecutionSelector(provider, resolver, requested) if err != nil { return nil, core.NewInvalidRequestError(err.Error(), err) } + if resolvedSelector == (core.ModelSelector{}) { + resolvedSelector, err = requested.Normalize() + if err != nil { + return nil, core.NewInvalidRequestError(err.Error(), err) + } + } resolvedModel := resolvedSelector.QualifiedModel() if counted, ok := provider.(modelCountProvider); ok && counted.ModelCount() == 0 { @@ -61,10 +60,49 @@ func resolveRequestModel(provider core.RoutableProvider, resolver RequestModelRe Requested: requested, ResolvedSelector: resolvedSelector, ProviderType: strings.TrimSpace(provider.GetProviderType(resolvedModel)), + ProviderName: resolvedProviderName(provider, resolvedSelector, ""), AliasApplied: aliasApplied, }, nil } +func resolveExecutionSelector( + provider core.RoutableProvider, + resolver RequestModelResolver, + requested core.RequestedModelSelector, +) (core.ModelSelector, bool, error) { + requested = core.NewRequestedModelSelector(requested.Model, requested.ProviderHint) + + var ( + resolvedSelector core.ModelSelector + aliasApplied bool + err error + ) + + if resolver != nil { + resolvedSelector, aliasApplied, err = resolver.ResolveModel(requested) + if err != nil { + return core.ModelSelector{}, false, err + } + requested = core.NewRequestedModelSelector(resolvedSelector.QualifiedModel(), "") + } + + if providerResolver, ok := provider.(RequestModelResolver); ok { + var providerChanged bool + resolvedSelector, providerChanged, err = providerResolver.ResolveModel(requested) + if err != nil { + return core.ModelSelector{}, false, err + } + return resolvedSelector, aliasApplied || providerChanged, nil + } + + if resolvedSelector != (core.ModelSelector{}) { + return resolvedSelector, aliasApplied, nil + } + + resolvedSelector, err = requested.Normalize() + return resolvedSelector, aliasApplied, err +} + func storeRequestModelResolution(c *echo.Context, resolution *core.RequestModelResolution) { if c == nil || resolution == nil { return diff --git a/internal/server/request_model_resolution_test.go b/internal/server/request_model_resolution_test.go new file mode 100644 index 000000000..60c2fbfdf --- /dev/null +++ b/internal/server/request_model_resolution_test.go @@ -0,0 +1,156 @@ +package server + +import ( + "context" + "io" + "testing" + + "gomodel/internal/core" +) + +type canonicalizingProvider struct { + resolved map[string]core.ModelSelector + types map[string]string + names map[string]string +} + +func (p *canonicalizingProvider) ResolveModel(requested core.RequestedModelSelector) (core.ModelSelector, bool, error) { + key := requested.RequestedQualifiedModel() + if selector, ok := p.resolved[key]; ok { + return selector, selector.QualifiedModel() != key, nil + } + selector, err := requested.Normalize() + return selector, false, err +} + +func (p *canonicalizingProvider) Supports(model string) bool { + _, ok := p.types[model] + return ok +} + +func (p *canonicalizingProvider) GetProviderType(model string) string { + return p.types[model] +} + +func (p *canonicalizingProvider) GetProviderName(model string) string { + return p.names[model] +} + +func (p *canonicalizingProvider) ChatCompletion(_ context.Context, _ *core.ChatRequest) (*core.ChatResponse, error) { + return nil, nil +} + +func (p *canonicalizingProvider) StreamChatCompletion(_ context.Context, _ *core.ChatRequest) (io.ReadCloser, error) { + return nil, nil +} + +func (p *canonicalizingProvider) ListModels(_ context.Context) (*core.ModelsResponse, error) { + return nil, nil +} + +func (p *canonicalizingProvider) Responses(_ context.Context, _ *core.ResponsesRequest) (*core.ResponsesResponse, error) { + return nil, nil +} + +func (p *canonicalizingProvider) StreamResponses(_ context.Context, _ *core.ResponsesRequest) (io.ReadCloser, error) { + return nil, nil +} + +func (p *canonicalizingProvider) Embeddings(_ context.Context, _ *core.EmbeddingRequest) (*core.EmbeddingResponse, error) { + return nil, nil +} + +func TestResolveRequestModel_UsesResolvedProviderNameInsteadOfSelectorPrefix(t *testing.T) { + provider := &mockProvider{ + supportedModels: []string{"gpt-5-nano"}, + providerTypes: map[string]string{ + "openai/gpt-5-nano": "openai", + }, + providerNames: map[string]string{ + "openai/gpt-5-nano": "openai_test", + }, + } + + resolution, err := resolveRequestModel(provider, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) + if err != nil { + t.Fatalf("resolveRequestModel() error = %v", err) + } + + if got := resolution.ResolvedSelector.Provider; got != "openai" { + t.Fatalf("ResolvedSelector.Provider = %q, want %q", got, "openai") + } + if got := resolution.ProviderName; got != "openai_test" { + t.Fatalf("ProviderName = %q, want %q", got, "openai_test") + } +} + +func TestResolveRequestModel_CanonicalizesProviderTypeSelectorToConcreteProviderName(t *testing.T) { + provider := &canonicalizingProvider{ + resolved: map[string]core.ModelSelector{ + "openai/gpt-5-nano": {Provider: "openai_test", Model: "gpt-5-nano"}, + }, + types: map[string]string{ + "openai_test/gpt-5-nano": "openai", + }, + names: map[string]string{ + "openai_test/gpt-5-nano": "openai_test", + }, + } + + resolution, err := resolveRequestModel(provider, nil, core.NewRequestedModelSelector("openai/gpt-5-nano", "")) + if err != nil { + t.Fatalf("resolveRequestModel() error = %v", err) + } + + if got := resolution.ResolvedQualifiedModel(); got != "openai_test/gpt-5-nano" { + t.Fatalf("ResolvedQualifiedModel = %q, want %q", got, "openai_test/gpt-5-nano") + } + if got := resolution.ProviderType; got != "openai" { + t.Fatalf("ProviderType = %q, want %q", got, "openai") + } + if got := resolution.ProviderName; got != "openai_test" { + t.Fatalf("ProviderName = %q, want %q", got, "openai_test") + } +} + +type aliasResolverStub struct{} + +func (aliasResolverStub) ResolveModel(requested core.RequestedModelSelector) (core.ModelSelector, bool, error) { + if requested.RequestedQualifiedModel() == "anthropic/claude-opus-4-6" { + return core.ModelSelector{Provider: "openai", Model: "gpt-5-nano"}, true, nil + } + selector, err := requested.Normalize() + return selector, false, err +} + +func TestResolveRequestModel_CanonicalizesAliasOutputThroughProviderResolver(t *testing.T) { + provider := &canonicalizingProvider{ + resolved: map[string]core.ModelSelector{ + "openai/gpt-5-nano": {Provider: "openai_test", Model: "gpt-5-nano"}, + }, + types: map[string]string{ + "openai_test/gpt-5-nano": "openai", + }, + names: map[string]string{ + "openai_test/gpt-5-nano": "openai_test", + }, + } + + resolution, err := resolveRequestModel(provider, aliasResolverStub{}, core.NewRequestedModelSelector("anthropic/claude-opus-4-6", "")) + if err != nil { + t.Fatalf("resolveRequestModel() error = %v", err) + } + + if !resolution.AliasApplied { + t.Fatal("AliasApplied = false, want true") + } + if got := resolution.ResolvedQualifiedModel(); got != "openai_test/gpt-5-nano" { + t.Fatalf("ResolvedQualifiedModel = %q, want %q", got, "openai_test/gpt-5-nano") + } + if got := resolution.ProviderType; got != "openai" { + t.Fatalf("ProviderType = %q, want %q", got, "openai") + } + if got := resolution.ProviderName; got != "openai_test" { + t.Fatalf("ProviderName = %q, want %q", got, "openai_test") + } +} diff --git a/internal/server/translated_inference_service.go b/internal/server/translated_inference_service.go index 09631e2ac..d296a3479 100644 --- a/internal/server/translated_inference_service.go +++ b/internal/server/translated_inference_service.go @@ -79,7 +79,7 @@ func (s *translatedInferenceService) ChatCompletion(c *echo.Context) error { func (s *translatedInferenceService) dispatchChatCompletion(c *echo.Context, req *core.ChatRequest, plan *core.ExecutionPlan) error { ctx := c.Request().Context() - streamReq, providerType, usageModel := s.resolveProviderAndModelFromPlan(c, plan, req.Model, req) + streamReq, providerType, providerName, usageModel := s.resolveProviderAndModelFromPlan(c, plan, req.Model, req) requestID := requestIDFromContextOrHeader(c.Request()) if req.Stream { @@ -88,25 +88,26 @@ func (s *translatedInferenceService) dispatchChatCompletion(c *echo.Context, req return err } } - stream, resolvedProviderType, resolvedUsageModel, usedFallback, err := s.streamChatCompletion(ctx, plan, streamReq, providerType, usageModel) + stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, usedFallback, err := s.streamChatCompletion(ctx, plan, streamReq, providerType, providerName, usageModel) if err != nil { return handleError(c, err) } if usedFallback { markRequestFallbackUsed(c) } - return s.handleStreamingReadCloser(c, plan, resolvedUsageModel, resolvedProviderType, stream) + return s.handleStreamingReadCloser(c, plan, resolvedUsageModel, resolvedProviderType, resolvedProviderName, stream) } - resp, providerType, usedFallback, err := s.executeChatCompletion(ctx, plan, req) + resp, providerType, providerName, usedFallback, err := s.executeChatCompletion(ctx, plan, req) if err != nil { return handleError(c, err) } if usedFallback { markRequestFallbackUsed(c) } + auditlog.EnrichEntryWithResolvedRoute(c, qualifyExecutedModel(plan, resp.Model, providerName), providerType, providerName) - s.logUsage(ctx, plan, resp.Model, providerType, func(pricing *core.ModelPricing) *usage.UsageEntry { + s.logUsage(ctx, plan, resp.Model, providerType, providerName, func(pricing *core.ModelPricing) *usage.UsageEntry { return usage.ExtractFromChatResponse(resp, requestID, providerType, "/v1/chat/completions", pricing) }) @@ -196,32 +197,33 @@ func (s *translatedInferenceService) withCacheRequestContext(ctx context.Context func (s *translatedInferenceService) dispatchResponses(c *echo.Context, req *core.ResponsesRequest, plan *core.ExecutionPlan) error { ctx := c.Request().Context() - _, providerType, usageModel := s.resolveProviderAndModelFromPlan(c, plan, req.Model, nil) + _, providerType, providerName, usageModel := s.resolveProviderAndModelFromPlan(c, plan, req.Model, nil) requestID := requestIDFromContextOrHeader(c.Request()) if req.Stream { if (plan == nil || plan.UsageEnabled()) && s.shouldEnforceReturningUsageData() { ctx = core.WithEnforceReturningUsageData(ctx, true) } - stream, resolvedProviderType, resolvedUsageModel, usedFallback, err := s.streamResponses(ctx, plan, req, providerType, usageModel) + stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, usedFallback, err := s.streamResponses(ctx, plan, req, providerType, providerName, usageModel) if err != nil { return handleError(c, err) } if usedFallback { markRequestFallbackUsed(c) } - return s.handleStreamingReadCloser(c, plan, resolvedUsageModel, resolvedProviderType, stream) + return s.handleStreamingReadCloser(c, plan, resolvedUsageModel, resolvedProviderType, resolvedProviderName, stream) } - resp, providerType, usedFallback, err := s.executeResponses(ctx, plan, req) + resp, providerType, providerName, usedFallback, err := s.executeResponses(ctx, plan, req) if err != nil { return handleError(c, err) } if usedFallback { markRequestFallbackUsed(c) } + auditlog.EnrichEntryWithResolvedRoute(c, qualifyExecutedModel(plan, resp.Model, providerName), providerType, providerName) - s.logUsage(ctx, plan, resp.Model, providerType, func(pricing *core.ModelPricing) *usage.UsageEntry { + s.logUsage(ctx, plan, resp.Model, providerType, providerName, func(pricing *core.ModelPricing) *usage.UsageEntry { return usage.ExtractFromResponsesResponse(resp, requestID, providerType, "/v1/responses", pricing) }) @@ -265,7 +267,7 @@ func (s *translatedInferenceService) tryFastPathStreamingChatPassthrough(c *echo usageLogger: s.usageLogger, pricingResolver: s.pricingResolver, } - return true, passthrough.proxyPassthroughResponse(c, providerType, endpoint, info, resp) + return true, passthrough.proxyPassthroughResponse(c, providerType, providerNameFromPlan(plan), endpoint, info, resp) } func (s *translatedInferenceService) canFastPathStreamingChatPassthrough(plan *core.ExecutionPlan, req *core.ChatRequest) bool { @@ -335,12 +337,13 @@ func (s *translatedInferenceService) Embeddings(c *echo.Context) error { ctx := c.Request().Context() requestID := requestIDFromContextOrHeader(c.Request()) - resp, providerType, err := s.executeEmbeddings(ctx, plan, req) + resp, providerType, providerName, err := s.executeEmbeddings(ctx, plan, req) if err != nil { return handleError(c, err) } + auditlog.EnrichEntryWithResolvedRoute(c, qualifyExecutedModel(plan, resp.Model, providerName), providerType, providerName) - s.logUsage(ctx, plan, resp.Model, providerType, func(pricing *core.ModelPricing) *usage.UsageEntry { + s.logUsage(ctx, plan, resp.Model, providerType, providerName, func(pricing *core.ModelPricing) *usage.UsageEntry { return usage.ExtractFromEmbeddingResponse(resp, requestID, providerType, "/v1/embeddings", pricing) }) @@ -350,11 +353,12 @@ func (s *translatedInferenceService) Embeddings(c *echo.Context) error { func (s *translatedInferenceService) handleStreamingReadCloser( c *echo.Context, plan *core.ExecutionPlan, - model, provider string, + model, provider, providerName string, stream io.ReadCloser, ) error { auditlog.MarkEntryAsStreaming(c, true) auditlog.EnrichEntryWithStream(c, true) + auditlog.EnrichEntryWithResolvedRoute(c, qualifyExecutedModel(plan, model, providerName), provider, providerName) entry := auditlog.GetStreamEntryFromContext(c) auditEnabled := s.logger != nil && s.logger.Config().Enabled && (plan == nil || plan.AuditEnabled()) @@ -373,7 +377,11 @@ func (s *translatedInferenceService) handleStreamingReadCloser( observers = append(observers, auditlog.NewStreamLogObserver(s.logger, streamEntry, endpoint)) } if s.usageLogger != nil && s.usageLogger.Config().Enabled && (plan == nil || plan.UsageEnabled()) { - observers = append(observers, usage.NewStreamUsageObserver(s.usageLogger, model, provider, requestID, endpoint, s.pricingResolver, core.UserPathFromContext(c.Request().Context()))) + usageObserver := usage.NewStreamUsageObserver(s.usageLogger, model, provider, requestID, endpoint, s.pricingResolver, core.UserPathFromContext(c.Request().Context())) + if usageObserver != nil { + usageObserver.SetProviderName(providerName) + observers = append(observers, usageObserver) + } } wrappedStream := streaming.NewObservedSSEStream(stream, observers...) @@ -399,14 +407,14 @@ func (s *translatedInferenceService) handleStreamingReadCloser( func (s *translatedInferenceService) handleStreamingResponse( c *echo.Context, plan *core.ExecutionPlan, - model, provider string, + model, provider, providerName string, streamFn func() (io.ReadCloser, error), ) error { stream, err := streamFn() if err != nil { return handleError(c, err) } - return s.handleStreamingReadCloser(c, plan, model, provider, stream) + return s.handleStreamingReadCloser(c, plan, model, provider, providerName, stream) } //nolint:dupl // typed wrapper over the shared translated fallback executor @@ -414,7 +422,7 @@ func (s *translatedInferenceService) executeChatCompletion( ctx context.Context, plan *core.ExecutionPlan, req *core.ChatRequest, -) (*core.ChatResponse, string, bool, error) { +) (*core.ChatResponse, string, string, bool, error) { return executeTranslatedWithFallback(ctx, s, plan, req, req.Model, req.Provider, cloneChatRequestForSelector, func(ctx context.Context, req *core.ChatRequest) (*core.ChatResponse, string, error) { resp, err := s.provider.ChatCompletion(ctx, req) @@ -430,15 +438,15 @@ func (s *translatedInferenceService) streamChatCompletion( ctx context.Context, plan *core.ExecutionPlan, req *core.ChatRequest, - providerType, usageModel string, -) (io.ReadCloser, string, string, bool, error) { + providerType, providerName, usageModel string, +) (io.ReadCloser, string, string, string, bool, error) { stream, err := s.provider.StreamChatCompletion(ctx, req) if err == nil { - return stream, providerType, usageModel, false, nil + return stream, providerType, providerName, usageModel, false, nil } - stream, resolvedProviderType, resolvedUsageModel, err := tryFallbackStream(ctx, s, plan, req.Model, req.Provider, err, - func(selector core.ModelSelector, providerType string) (io.ReadCloser, string, string, error) { + stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, err := tryFallbackStream(ctx, s, plan, req.Model, req.Provider, err, + func(selector core.ModelSelector, providerType, providerName string) (io.ReadCloser, string, string, error) { stream, err := s.provider.StreamChatCompletion(ctx, cloneChatRequestForSelector(req, selector)) if err != nil { return nil, "", "", err @@ -447,9 +455,9 @@ func (s *translatedInferenceService) streamChatCompletion( }, ) if err != nil { - return nil, "", "", false, err + return nil, "", "", "", false, err } - return stream, resolvedProviderType, resolvedUsageModel, true, nil + return stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, true, nil } //nolint:dupl // typed wrapper over the shared translated fallback executor @@ -457,7 +465,7 @@ func (s *translatedInferenceService) executeResponses( ctx context.Context, plan *core.ExecutionPlan, req *core.ResponsesRequest, -) (*core.ResponsesResponse, string, bool, error) { +) (*core.ResponsesResponse, string, string, bool, error) { return executeTranslatedWithFallback(ctx, s, plan, req, req.Model, req.Provider, cloneResponsesRequestForSelector, func(ctx context.Context, req *core.ResponsesRequest) (*core.ResponsesResponse, string, error) { resp, err := s.provider.Responses(ctx, req) @@ -473,15 +481,15 @@ func (s *translatedInferenceService) streamResponses( ctx context.Context, plan *core.ExecutionPlan, req *core.ResponsesRequest, - providerType, usageModel string, -) (io.ReadCloser, string, string, bool, error) { + providerType, providerName, usageModel string, +) (io.ReadCloser, string, string, string, bool, error) { stream, err := s.provider.StreamResponses(ctx, req) if err == nil { - return stream, providerType, usageModel, false, nil + return stream, providerType, providerName, usageModel, false, nil } - stream, resolvedProviderType, resolvedUsageModel, err := tryFallbackStream(ctx, s, plan, req.Model, req.Provider, err, - func(selector core.ModelSelector, providerType string) (io.ReadCloser, string, string, error) { + stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, err := tryFallbackStream(ctx, s, plan, req.Model, req.Provider, err, + func(selector core.ModelSelector, providerType, providerName string) (io.ReadCloser, string, string, error) { stream, err := s.provider.StreamResponses(ctx, cloneResponsesRequestForSelector(req, selector)) if err != nil { return nil, "", "", err @@ -490,20 +498,21 @@ func (s *translatedInferenceService) streamResponses( }, ) if err != nil { - return nil, "", "", false, err + return nil, "", "", "", false, err } - return stream, resolvedProviderType, resolvedUsageModel, true, nil + return stream, resolvedProviderType, resolvedProviderName, resolvedUsageModel, true, nil } func (s *translatedInferenceService) executeEmbeddings( ctx context.Context, plan *core.ExecutionPlan, req *core.EmbeddingRequest, -) (*core.EmbeddingResponse, string, error) { +) (*core.EmbeddingResponse, string, string, error) { providerType := providerTypeFromPlan(plan) + providerName := providerNameFromPlan(plan) resp, err := s.provider.Embeddings(ctx, req) if err == nil { - return resp, responseProviderType(providerType, resp.Provider), nil + return resp, responseProviderType(providerType, resp.Provider), providerName, nil } return s.tryFallbackEmbeddings(ctx, plan, req, err) @@ -514,16 +523,16 @@ func (s *translatedInferenceService) tryFallbackEmbeddings( plan *core.ExecutionPlan, req *core.EmbeddingRequest, primaryErr error, -) (*core.EmbeddingResponse, string, error) { +) (*core.EmbeddingResponse, string, string, error) { // Embeddings fallback is intentionally disabled until the shared model // contract can prove vector-size compatibility for alternates. - return nil, "", primaryErr + return nil, "", "", primaryErr } func (s *translatedInferenceService) logUsage( ctx context.Context, plan *core.ExecutionPlan, - model, providerType string, + model, providerType, providerName string, extractFn func(*core.ModelPricing) *usage.UsageEntry, ) { if s.usageLogger == nil || !s.usageLogger.Config().Enabled || (plan != nil && !plan.UsageEnabled()) { @@ -534,6 +543,7 @@ func (s *translatedInferenceService) logUsage( pricing = s.pricingResolver.ResolvePricing(model, providerType) } if entry := extractFn(pricing); entry != nil { + entry.ProviderName = strings.TrimSpace(providerName) entry.UserPath = core.UserPathFromContext(ctx) s.usageLogger.Write(entry) } @@ -552,15 +562,18 @@ func (s *translatedInferenceService) fallbackSelectors(plan *core.ExecutionPlan) func (s *translatedInferenceService) providerTypeForSelector(selector core.ModelSelector, fallback string) string { fallback = strings.TrimSpace(fallback) - if provider := strings.TrimSpace(selector.Provider); provider != "" { - return provider - } if s.provider == nil { + if provider := strings.TrimSpace(selector.Provider); provider != "" { + return provider + } return fallback } if providerType := strings.TrimSpace(s.provider.GetProviderType(selector.QualifiedModel())); providerType != "" { return providerType } + if provider := strings.TrimSpace(selector.Provider); provider != "" { + return provider + } return fallback } @@ -569,8 +582,9 @@ func (s *translatedInferenceService) resolveProviderAndModelFromPlan( plan *core.ExecutionPlan, fallbackModel string, req *core.ChatRequest, -) (*core.ChatRequest, string, string) { +) (*core.ChatRequest, string, string, string) { providerType := GetProviderType(c) + providerName := providerNameFromPlan(plan) if plan != nil { if plannedProviderType := strings.TrimSpace(plan.ProviderType); plannedProviderType != "" { providerType = plannedProviderType @@ -579,7 +593,7 @@ func (s *translatedInferenceService) resolveProviderAndModelFromPlan( model := resolvedModelFromPlan(plan, fallbackModel) if req == nil || !req.Stream || (plan != nil && !plan.UsageEnabled()) || !s.shouldEnforceReturningUsageData() { - return req, providerType, model + return req, providerType, providerName, model } streamReq := cloneChatRequestForStreamUsage(req) @@ -587,7 +601,7 @@ func (s *translatedInferenceService) resolveProviderAndModelFromPlan( streamReq.StreamOptions = &core.StreamOptions{} } streamReq.StreamOptions.IncludeUsage = true - return streamReq, providerType, model + return streamReq, providerType, providerName, model } func recordStreamingError(streamEntry *auditlog.LogEntry, model, provider, path, requestID string, err error) { @@ -648,6 +662,42 @@ func cloneResponsesRequestForSelector(req *core.ResponsesRequest, selector core. return &cloned } +func providerNameFromPlan(plan *core.ExecutionPlan) string { + if plan == nil || plan.Resolution == nil { + return "" + } + return strings.TrimSpace(plan.Resolution.ProviderName) +} + +func resolvedModelPrefix(plan *core.ExecutionPlan, providerName string) string { + if providerName = strings.TrimSpace(providerName); providerName != "" { + return providerName + } + if plan == nil || plan.Resolution == nil { + return "" + } + if providerName = strings.TrimSpace(plan.Resolution.ProviderName); providerName != "" { + return providerName + } + return strings.TrimSpace(plan.Resolution.ResolvedSelector.Provider) +} + +func qualifyModelWithProvider(model, providerName string) string { + model = strings.TrimSpace(model) + providerName = strings.TrimSpace(providerName) + if model == "" { + return "" + } + if providerName == "" || strings.HasPrefix(model, providerName+"/") { + return model + } + return providerName + "/" + model +} + +func qualifyExecutedModel(plan *core.ExecutionPlan, model, providerName string) string { + return qualifyModelWithProvider(model, resolvedModelPrefix(plan, providerName)) +} + func markRequestFallbackUsed(c *echo.Context) { if c == nil || c.Request() == nil { return @@ -706,13 +756,13 @@ func tryFallbackResponse[T any]( plan *core.ExecutionPlan, model, provider string, primaryErr error, - call func(selector core.ModelSelector, providerType string) (T, string, error), -) (T, string, bool, error) { + call func(selector core.ModelSelector, providerType, providerName string) (T, string, error), +) (T, string, string, bool, error) { var zero T fallbacks := s.fallbackSelectors(plan) if len(fallbacks) == 0 || !shouldAttemptFallback(primaryErr) { - return zero, "", false, primaryErr + return zero, "", "", false, primaryErr } requestID := strings.TrimSpace(core.GetRequestID(ctx)) @@ -721,6 +771,7 @@ func tryFallbackResponse[T any]( for _, selector := range fallbacks { qualified := selector.QualifiedModel() providerType := s.providerTypeForSelector(selector, providerTypeFromPlan(plan)) + providerName := resolvedProviderName(s.provider, selector, providerNameFromPlan(plan)) slog.Warn("primary model attempt failed, trying fallback", "request_id", requestID, "from", primaryModel, @@ -729,7 +780,7 @@ func tryFallbackResponse[T any]( "error", lastErr, ) - resp, resolvedProviderType, err := call(selector, providerType) + resp, resolvedProviderType, err := call(selector, providerType, providerName) if err == nil { slog.Info("fallback model attempt succeeded", "request_id", requestID, @@ -737,12 +788,12 @@ func tryFallbackResponse[T any]( "to", qualified, "provider_type", resolvedProviderType, ) - return resp, resolvedProviderType, true, nil + return resp, resolvedProviderType, providerName, true, nil } lastErr = err } - return zero, "", false, lastErr + return zero, "", "", false, lastErr } func executeWithFallbackResponse[T any]( @@ -750,12 +801,12 @@ func executeWithFallbackResponse[T any]( s *translatedInferenceService, plan *core.ExecutionPlan, model, provider string, - primary func() (T, string, error), - fallback func(selector core.ModelSelector, providerType string) (T, string, error), -) (T, string, bool, error) { - resp, resolvedProviderType, err := primary() + primary func() (T, string, string, error), + fallback func(selector core.ModelSelector, providerType, providerName string) (T, string, error), +) (T, string, string, bool, error) { + resp, resolvedProviderType, resolvedProviderName, err := primary() if err == nil { - return resp, resolvedProviderType, false, nil + return resp, resolvedProviderType, resolvedProviderName, false, nil } return tryFallbackResponse(ctx, s, plan, model, provider, err, fallback) } @@ -768,17 +819,17 @@ func executeTranslatedWithFallback[Req any, Resp any]( model, provider string, cloneForSelector func(Req, core.ModelSelector) Req, call func(context.Context, Req) (Resp, string, error), -) (Resp, string, bool, error) { +) (Resp, string, string, bool, error) { return executeWithFallbackResponse(ctx, s, plan, model, provider, - func() (Resp, string, error) { + func() (Resp, string, string, error) { resp, responseProvider, err := call(ctx, req) if err != nil { var zero Resp - return zero, "", err + return zero, "", "", err } - return resp, responseProviderType(providerTypeFromPlan(plan), responseProvider), nil + return resp, responseProviderType(providerTypeFromPlan(plan), responseProvider), providerNameFromPlan(plan), nil }, - func(selector core.ModelSelector, providerType string) (Resp, string, error) { + func(selector core.ModelSelector, providerType, providerName string) (Resp, string, error) { resp, responseProvider, err := call(ctx, cloneForSelector(req, selector)) if err != nil { var zero Resp @@ -795,11 +846,11 @@ func tryFallbackStream( plan *core.ExecutionPlan, model, provider string, primaryErr error, - call func(selector core.ModelSelector, providerType string) (io.ReadCloser, string, string, error), -) (io.ReadCloser, string, string, error) { + call func(selector core.ModelSelector, providerType, providerName string) (io.ReadCloser, string, string, error), +) (io.ReadCloser, string, string, string, error) { fallbacks := s.fallbackSelectors(plan) if len(fallbacks) == 0 || !shouldAttemptFallback(primaryErr) { - return nil, "", "", primaryErr + return nil, "", "", "", primaryErr } requestID := strings.TrimSpace(core.GetRequestID(ctx)) @@ -808,6 +859,7 @@ func tryFallbackStream( for _, selector := range fallbacks { qualified := selector.QualifiedModel() providerType := s.providerTypeForSelector(selector, providerTypeFromPlan(plan)) + providerName := resolvedProviderName(s.provider, selector, providerNameFromPlan(plan)) slog.Warn("primary model attempt failed, trying fallback stream", "request_id", requestID, "from", primaryModel, @@ -816,7 +868,7 @@ func tryFallbackStream( "error", lastErr, ) - stream, resolvedProviderType, usageModel, err := call(selector, providerType) + stream, resolvedProviderType, usageModel, err := call(selector, providerType, providerName) if err == nil { slog.Info("fallback stream attempt succeeded", "request_id", requestID, @@ -824,12 +876,12 @@ func tryFallbackStream( "to", qualified, "provider_type", resolvedProviderType, ) - return stream, resolvedProviderType, usageModel, nil + return stream, resolvedProviderType, providerName, usageModel, nil } lastErr = err } - return nil, "", "", lastErr + return nil, "", "", "", lastErr } func shouldAttemptFallback(err error) bool { diff --git a/internal/server/translated_inference_service_test.go b/internal/server/translated_inference_service_test.go index eef65c89c..14a12b07c 100644 --- a/internal/server/translated_inference_service_test.go +++ b/internal/server/translated_inference_service_test.go @@ -38,7 +38,7 @@ func TestTranslatedInferenceService_LogUsageSkipsWhenExecutionPlanDisablesUsage( Guardrails: true, }, }, - }, "gpt-5-nano", "openai", func(*core.ModelPricing) *usage.UsageEntry { + }, "gpt-5-nano", "openai", "primary-openai", func(*core.ModelPricing) *usage.UsageEntry { return &usage.UsageEntry{ID: "usage-1"} }) @@ -58,6 +58,39 @@ func TestTranslatedInferenceService_ProviderTypeForSelectorPrefersExplicitProvid } } +func TestTranslatedInferenceService_ProviderTypeForSelectorCanonicalizesProviderNameSelectors(t *testing.T) { + service := &translatedInferenceService{ + provider: &mockProvider{ + supportedModels: []string{"gpt-4o"}, + providerTypes: map[string]string{ + "openai_test/gpt-4o": "openai", + }, + providerNames: map[string]string{ + "openai_test/gpt-4o": "openai_test", + }, + }, + } + + got := service.providerTypeForSelector(core.ModelSelector{Provider: "openai_test", Model: "gpt-4o"}, "anthropic") + if got != "openai" { + t.Fatalf("providerTypeForSelector() = %q, want %q", got, "openai") + } +} + +func TestQualifyModelWithProvider_PrefixesSlashModelIDs(t *testing.T) { + got := qualifyModelWithProvider("openai/gpt-4o-mini", "openrouter") + if got != "openrouter/openai/gpt-4o-mini" { + t.Fatalf("qualifyModelWithProvider() = %q, want %q", got, "openrouter/openai/gpt-4o-mini") + } +} + +func TestQualifyModelWithProvider_KeepsAlreadyQualifiedModelIDs(t *testing.T) { + got := qualifyModelWithProvider("openrouter/openai/gpt-4o-mini", "openrouter") + if got != "openrouter/openai/gpt-4o-mini" { + t.Fatalf("qualifyModelWithProvider() = %q, want unchanged model", got) + } +} + func TestTranslatedInferenceService_LogUsageAssignsUserPathFromContext(t *testing.T) { logger := &usageCaptureLogger{ config: usage.Config{Enabled: true}, @@ -70,7 +103,7 @@ func TestTranslatedInferenceService_LogUsageAssignsUserPathFromContext(t *testin UserPath: "/team/alpha", }) - service.logUsage(ctx, nil, "gpt-5-nano", "openai", func(*core.ModelPricing) *usage.UsageEntry { + service.logUsage(ctx, nil, "gpt-5-nano", "openai", "primary-openai", func(*core.ModelPricing) *usage.UsageEntry { return &usage.UsageEntry{ID: "usage-1"} }) @@ -80,6 +113,9 @@ func TestTranslatedInferenceService_LogUsageAssignsUserPathFromContext(t *testin if got := logger.entries[0].UserPath; got != "/team/alpha" { t.Fatalf("UserPath = %q, want /team/alpha", got) } + if got := logger.entries[0].ProviderName; got != "primary-openai" { + t.Fatalf("ProviderName = %q, want primary-openai", got) + } } func TestTranslatedInferenceService_WithCacheRequestContextClearsInheritedGuardrailsHash(t *testing.T) { diff --git a/internal/usage/cache_type.go b/internal/usage/cache_type.go index 1e36ba690..2f338c00c 100644 --- a/internal/usage/cache_type.go +++ b/internal/usage/cache_type.go @@ -46,11 +46,13 @@ func normalizedUsageEntryForStorage(entry *UsageEntry) *UsageEntry { } normalized := normalizeCacheType(entry.CacheType) - if normalized == entry.CacheType { + providerName := strings.TrimSpace(entry.ProviderName) + if normalized == entry.CacheType && providerName == entry.ProviderName { return entry } cloned := *entry cloned.CacheType = normalized + cloned.ProviderName = providerName return &cloned } diff --git a/internal/usage/reader.go b/internal/usage/reader.go index 0d42a8248..e2461e68c 100644 --- a/internal/usage/reader.go +++ b/internal/usage/reader.go @@ -2,6 +2,7 @@ package usage import ( "context" + "strings" "time" ) @@ -30,6 +31,7 @@ type UsageSummary struct { type ModelUsage struct { Model string `json:"model"` Provider string `json:"provider"` + ProviderName string `json:"provider_name,omitempty"` InputTokens int64 `json:"input_tokens"` OutputTokens int64 `json:"output_tokens"` InputCost *float64 `json:"input_cost"` @@ -55,7 +57,7 @@ type DailyUsage struct { type UsageLogParams struct { UsageQueryParams // embed date range Model string // filter by model (optional) - Provider string // filter by provider (optional) + Provider string // filter by provider name or provider type (optional) Search string // free-text search on model/provider/request_id Limit int // page size (default 50, max 200) Offset int // pagination offset @@ -69,6 +71,7 @@ type UsageLogEntry struct { Timestamp time.Time `json:"timestamp"` Model string `json:"model"` Provider string `json:"provider"` + ProviderName string `json:"provider_name,omitempty"` Endpoint string `json:"endpoint"` UserPath string `json:"user_path,omitempty"` CacheType string `json:"cache_type,omitempty"` @@ -138,3 +141,10 @@ type UsageReader interface { // GetCacheOverview returns cached-only aggregates for the admin dashboard. GetCacheOverview(ctx context.Context, params UsageQueryParams) (*CacheOverview, error) } + +func displayUsageProviderName(providerName, provider string) string { + if trimmed := strings.TrimSpace(providerName); trimmed != "" { + return trimmed + } + return strings.TrimSpace(provider) +} diff --git a/internal/usage/reader_helpers.go b/internal/usage/reader_helpers.go index c1d113d97..995b6e6e4 100644 --- a/internal/usage/reader_helpers.go +++ b/internal/usage/reader_helpers.go @@ -22,6 +22,12 @@ func buildWhereClause(conditions []string) string { return " WHERE " + strings.Join(conditions, " AND ") } +// usageGroupedProviderNameSQL returns a SQL expression that collapses blank +// provider_name values to the canonical provider before grouping. +func usageGroupedProviderNameSQL(providerNameColumn, providerColumn string) string { + return "COALESCE(NULLIF(TRIM(" + providerNameColumn + "), ''), " + providerColumn + ")" +} + // clampLimitOffset normalises pagination parameters: // - limit defaults to 50 and is capped at 200 // - offset floors at 0 diff --git a/internal/usage/reader_mongodb.go b/internal/usage/reader_mongodb.go index f2a235de6..3f94b3036 100644 --- a/internal/usage/reader_mongodb.go +++ b/internal/usage/reader_mongodb.go @@ -96,10 +96,12 @@ func (r *MongoDBReader) GetUsageByModel(ctx context.Context, params UsageQueryPa pipeline = append(pipeline, bson.D{{Key: "$match", Value: matchFilters}}) } + providerNameExpr := mongoUsageGroupedProviderNameExpr() pipeline = append(pipeline, bson.D{{Key: "$group", Value: bson.D{ {Key: "_id", Value: bson.D{ {Key: "model", Value: "$model"}, {Key: "provider", Value: "$provider"}, + {Key: "provider_name", Value: providerNameExpr}, }}, {Key: "input_tokens", Value: bson.D{{Key: "$sum", Value: "$input_tokens"}}}, {Key: "output_tokens", Value: bson.D{{Key: "$sum", Value: "$output_tokens"}}}, @@ -119,8 +121,9 @@ func (r *MongoDBReader) GetUsageByModel(ctx context.Context, params UsageQueryPa for cursor.Next(ctx) { var row struct { ID struct { - Model string `bson:"model"` - Provider string `bson:"provider"` + Model string `bson:"model"` + Provider string `bson:"provider"` + ProviderName string `bson:"provider_name"` } `bson:"_id"` InputTokens int64 `bson:"input_tokens"` OutputTokens int64 `bson:"output_tokens"` @@ -135,6 +138,7 @@ func (r *MongoDBReader) GetUsageByModel(ctx context.Context, params UsageQueryPa m := ModelUsage{ Model: row.ID.Model, Provider: row.ID.Provider, + ProviderName: displayUsageProviderName(row.ID.ProviderName, row.ID.Provider), InputTokens: row.InputTokens, OutputTokens: row.OutputTokens, } @@ -153,6 +157,17 @@ func (r *MongoDBReader) GetUsageByModel(ctx context.Context, params UsageQueryPa return result, nil } +func mongoUsageGroupedProviderNameExpr() bson.D { + trimmedProviderName := bson.D{{Key: "$trim", Value: bson.D{ + {Key: "input", Value: bson.D{{Key: "$ifNull", Value: bson.A{"$provider_name", ""}}}}, + }}} + return bson.D{{Key: "$cond", Value: bson.A{ + bson.D{{Key: "$ne", Value: bson.A{trimmedProviderName, ""}}}, + trimmedProviderName, + bson.D{{Key: "$trim", Value: bson.D{{Key: "input", Value: "$provider"}}}}, + }}} +} + // GetUsageLog returns a paginated list of individual usage log entries. func (r *MongoDBReader) GetUsageLog(ctx context.Context, params UsageLogParams) (*UsageLogResult, error) { limit, offset := clampLimitOffset(params.Limit, params.Offset) @@ -192,6 +207,7 @@ func (r *MongoDBReader) GetUsageLog(ctx context.Context, params UsageLogParams) Timestamp time.Time `bson:"timestamp"` Model string `bson:"model"` Provider string `bson:"provider"` + ProviderName string `bson:"provider_name"` Endpoint string `bson:"endpoint"` UserPath string `bson:"user_path"` CacheType string `bson:"cache_type"` @@ -233,6 +249,7 @@ func (r *MongoDBReader) GetUsageLog(ctx context.Context, params UsageLogParams) Timestamp: row.Timestamp, Model: row.Model, Provider: row.Provider, + ProviderName: displayUsageProviderName(row.ProviderName, row.Provider), Endpoint: row.Endpoint, UserPath: row.UserPath, CacheType: normalizeCacheType(row.CacheType), @@ -526,13 +543,17 @@ func mongoUsageLogMatchFilters(params UsageLogParams) (bson.D, error) { matchFilters = append(matchFilters, bson.E{Key: "model", Value: params.Model}) } if params.Provider != "" { - matchFilters = append(matchFilters, bson.E{Key: "provider", Value: params.Provider}) + matchFilters = mongoAndFilters(matchFilters, bson.D{{Key: "$or", Value: bson.A{ + bson.D{{Key: "provider", Value: params.Provider}}, + bson.D{{Key: "provider_name", Value: params.Provider}}, + }}}) } if params.Search != "" { regex := bson.D{{Key: "$regex", Value: regexp.QuoteMeta(params.Search)}, {Key: "$options", Value: "i"}} searchFilter := bson.D{{Key: "$or", Value: bson.A{ bson.D{{Key: "model", Value: regex}}, bson.D{{Key: "provider", Value: regex}}, + bson.D{{Key: "provider_name", Value: regex}}, bson.D{{Key: "request_id", Value: regex}}, bson.D{{Key: "provider_id", Value: regex}}, }}} diff --git a/internal/usage/reader_mongodb_grouping_test.go b/internal/usage/reader_mongodb_grouping_test.go new file mode 100644 index 000000000..c9bcd2c7b --- /dev/null +++ b/internal/usage/reader_mongodb_grouping_test.go @@ -0,0 +1,28 @@ +package usage + +import ( + "reflect" + "testing" + + "go.mongodb.org/mongo-driver/v2/bson" +) + +func TestMongoUsageGroupedProviderNameExpr_CollapsesBlankProviderName(t *testing.T) { + got := mongoUsageGroupedProviderNameExpr() + want := bson.D{{Key: "$cond", Value: bson.A{ + bson.D{{Key: "$ne", Value: bson.A{ + bson.D{{Key: "$trim", Value: bson.D{ + {Key: "input", Value: bson.D{{Key: "$ifNull", Value: bson.A{"$provider_name", ""}}}}, + }}}, + "", + }}}, + bson.D{{Key: "$trim", Value: bson.D{ + {Key: "input", Value: bson.D{{Key: "$ifNull", Value: bson.A{"$provider_name", ""}}}}, + }}}, + bson.D{{Key: "$trim", Value: bson.D{{Key: "input", Value: "$provider"}}}}, + }}} + + if !reflect.DeepEqual(got, want) { + t.Fatalf("mongoUsageGroupedProviderNameExpr() = %#v, want %#v", got, want) + } +} diff --git a/internal/usage/reader_mongodb_test.go b/internal/usage/reader_mongodb_test.go index ca0125f55..50591310c 100644 --- a/internal/usage/reader_mongodb_test.go +++ b/internal/usage/reader_mongodb_test.go @@ -28,6 +28,7 @@ func TestMongoUsageLogMatchFiltersAndSearchWithCacheMode(t *testing.T) { bson.D{{Key: "$or", Value: bson.A{ bson.D{{Key: "model", Value: regex}}, bson.D{{Key: "provider", Value: regex}}, + bson.D{{Key: "provider_name", Value: regex}}, bson.D{{Key: "request_id", Value: regex}}, bson.D{{Key: "provider_id", Value: regex}}, }}}, @@ -53,6 +54,7 @@ func TestMongoUsageLogMatchFiltersEscapesSearchRegex(t *testing.T) { want := bson.D{{Key: "$or", Value: bson.A{ bson.D{{Key: "model", Value: regex}}, bson.D{{Key: "provider", Value: regex}}, + bson.D{{Key: "provider_name", Value: regex}}, bson.D{{Key: "request_id", Value: regex}}, bson.D{{Key: "provider_id", Value: regex}}, }}} diff --git a/internal/usage/reader_postgresql.go b/internal/usage/reader_postgresql.go index 2e530c2a8..72e985ea4 100644 --- a/internal/usage/reader_postgresql.go +++ b/internal/usage/reader_postgresql.go @@ -54,10 +54,11 @@ func (r *PostgreSQLReader) GetUsageByModel(ctx context.Context, params UsageQuer return nil, err } where := buildWhereClause(conditions) + providerNameExpr := usageGroupedProviderNameSQL("provider_name", "provider") costCols := `, SUM(input_cost), SUM(output_cost), SUM(total_cost)` - query := `SELECT model, provider, COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0)` + costCols + ` - FROM "usage"` + where + ` GROUP BY model, provider` + query := `SELECT model, provider, ` + providerNameExpr + ` AS provider_name, COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0)` + costCols + ` + FROM "usage"` + where + ` GROUP BY model, provider, ` + providerNameExpr rows, err := r.pool.Query(ctx, query, args...) if err != nil { @@ -68,7 +69,7 @@ func (r *PostgreSQLReader) GetUsageByModel(ctx context.Context, params UsageQuer result := make([]ModelUsage, 0) for rows.Next() { var m ModelUsage - if err := rows.Scan(&m.Model, &m.Provider, &m.InputTokens, &m.OutputTokens, &m.InputCost, &m.OutputCost, &m.TotalCost); err != nil { + if err := rows.Scan(&m.Model, &m.Provider, &m.ProviderName, &m.InputTokens, &m.OutputTokens, &m.InputCost, &m.OutputCost, &m.TotalCost); err != nil { return nil, fmt.Errorf("failed to scan usage by model row: %w", err) } result = append(result, m) @@ -96,13 +97,13 @@ func (r *PostgreSQLReader) GetUsageLog(ctx context.Context, params UsageLogParam argIdx++ } if params.Provider != "" { - conditions = append(conditions, fmt.Sprintf("provider = $%d", argIdx)) - args = append(args, params.Provider) - argIdx++ + conditions = append(conditions, fmt.Sprintf("(provider = $%d OR provider_name = $%d)", argIdx, argIdx+1)) + args = append(args, params.Provider, params.Provider) + argIdx += 2 } if params.Search != "" { s := "%" + escapeLikeWildcards(params.Search) + "%" - conditions = append(conditions, fmt.Sprintf("(model ILIKE $%d ESCAPE '\\' OR provider ILIKE $%d ESCAPE '\\' OR request_id ILIKE $%d ESCAPE '\\' OR provider_id ILIKE $%d ESCAPE '\\')", argIdx, argIdx, argIdx, argIdx)) + conditions = append(conditions, fmt.Sprintf("(model ILIKE $%d ESCAPE '\\' OR provider ILIKE $%d ESCAPE '\\' OR provider_name ILIKE $%d ESCAPE '\\' OR request_id ILIKE $%d ESCAPE '\\' OR provider_id ILIKE $%d ESCAPE '\\')", argIdx, argIdx, argIdx, argIdx, argIdx)) args = append(args, s) argIdx++ } @@ -117,7 +118,7 @@ func (r *PostgreSQLReader) GetUsageLog(ctx context.Context, params UsageLogParam } // Fetch page - dataQuery := fmt.Sprintf(`SELECT id, request_id, provider_id, timestamp, model, provider, endpoint, user_path, cache_type, + dataQuery := fmt.Sprintf(`SELECT id, request_id, provider_id, timestamp, model, provider, provider_name, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, COALESCE(input_cost, 0), COALESCE(output_cost, 0), COALESCE(total_cost, 0), raw_data, COALESCE(costs_calculation_caveat, '') FROM "usage"%s ORDER BY timestamp DESC LIMIT $%d OFFSET $%d`, where, argIdx, argIdx+1) dataArgs := append(append([]any(nil), args...), limit, offset) @@ -132,9 +133,10 @@ func (r *PostgreSQLReader) GetUsageLog(ctx context.Context, params UsageLogParam for rows.Next() { var e UsageLogEntry var rawDataJSON *string + var providerName *string var userPath *string var cacheType *string - if err := rows.Scan(&e.ID, &e.RequestID, &e.ProviderID, &e.Timestamp, &e.Model, &e.Provider, &e.Endpoint, &userPath, &cacheType, + if err := rows.Scan(&e.ID, &e.RequestID, &e.ProviderID, &e.Timestamp, &e.Model, &e.Provider, &providerName, &e.Endpoint, &userPath, &cacheType, &e.InputTokens, &e.OutputTokens, &e.TotalTokens, &e.InputCost, &e.OutputCost, &e.TotalCost, &rawDataJSON, &e.CostsCalculationCaveat); err != nil { return nil, fmt.Errorf("failed to scan usage log row: %w", err) } @@ -146,6 +148,11 @@ func (r *PostgreSQLReader) GetUsageLog(ctx context.Context, params UsageLogParam if userPath != nil { e.UserPath = *userPath } + if providerName != nil { + e.ProviderName = displayUsageProviderName(*providerName, e.Provider) + } else { + e.ProviderName = displayUsageProviderName("", e.Provider) + } if cacheType != nil { e.CacheType = normalizeCacheType(*cacheType) } diff --git a/internal/usage/reader_sqlite.go b/internal/usage/reader_sqlite.go index 593e7c9d3..caf461df9 100644 --- a/internal/usage/reader_sqlite.go +++ b/internal/usage/reader_sqlite.go @@ -54,10 +54,11 @@ func (r *SQLiteReader) GetUsageByModel(ctx context.Context, params UsageQueryPar return nil, err } where := buildWhereClause(conditions) + providerNameExpr := usageGroupedProviderNameSQL("provider_name", "provider") costCols := `, SUM(input_cost), SUM(output_cost), SUM(total_cost)` - query := `SELECT model, provider, COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0)` + costCols + ` - FROM usage` + where + ` GROUP BY model, provider` + query := `SELECT model, provider, ` + providerNameExpr + ` AS provider_name, COALESCE(SUM(input_tokens), 0), COALESCE(SUM(output_tokens), 0)` + costCols + ` + FROM usage` + where + ` GROUP BY model, provider, ` + providerNameExpr rows, err := r.db.QueryContext(ctx, query, args...) if err != nil { @@ -68,7 +69,7 @@ func (r *SQLiteReader) GetUsageByModel(ctx context.Context, params UsageQueryPar result := make([]ModelUsage, 0) for rows.Next() { var m ModelUsage - if err := rows.Scan(&m.Model, &m.Provider, &m.InputTokens, &m.OutputTokens, &m.InputCost, &m.OutputCost, &m.TotalCost); err != nil { + if err := rows.Scan(&m.Model, &m.Provider, &m.ProviderName, &m.InputTokens, &m.OutputTokens, &m.InputCost, &m.OutputCost, &m.TotalCost); err != nil { return nil, fmt.Errorf("failed to scan usage by model row: %w", err) } result = append(result, m) @@ -95,13 +96,13 @@ func (r *SQLiteReader) GetUsageLog(ctx context.Context, params UsageLogParams) ( args = append(args, params.Model) } if params.Provider != "" { - conditions = append(conditions, "provider = ?") - args = append(args, params.Provider) + conditions = append(conditions, "(provider = ? OR provider_name = ?)") + args = append(args, params.Provider, params.Provider) } if params.Search != "" { - conditions = append(conditions, "(model LIKE ? ESCAPE '\\' OR provider LIKE ? ESCAPE '\\' OR request_id LIKE ? ESCAPE '\\' OR provider_id LIKE ? ESCAPE '\\')") + conditions = append(conditions, "(model LIKE ? ESCAPE '\\' OR provider LIKE ? ESCAPE '\\' OR provider_name LIKE ? ESCAPE '\\' OR request_id LIKE ? ESCAPE '\\' OR provider_id LIKE ? ESCAPE '\\')") s := "%" + escapeLikeWildcards(params.Search) + "%" - args = append(args, s, s, s, s) + args = append(args, s, s, s, s, s) } where := buildWhereClause(conditions) @@ -114,7 +115,7 @@ func (r *SQLiteReader) GetUsageLog(ctx context.Context, params UsageLogParams) ( } // Fetch page - dataQuery := `SELECT id, request_id, provider_id, timestamp, model, provider, endpoint, user_path, cache_type, + dataQuery := `SELECT id, request_id, provider_id, timestamp, model, provider, provider_name, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, COALESCE(input_cost, 0), COALESCE(output_cost, 0), COALESCE(total_cost, 0), raw_data, COALESCE(costs_calculation_caveat, '') FROM usage` + where + ` ORDER BY ` + sqliteTimestampEpochExpr() + ` DESC, id DESC LIMIT ? OFFSET ?` dataArgs := append(append([]any(nil), args...), limit, offset) @@ -131,9 +132,10 @@ func (r *SQLiteReader) GetUsageLog(ctx context.Context, params UsageLogParams) ( var ts string var caveat *string var rawDataJSON *string + var providerName sql.NullString var userPath sql.NullString var cacheType sql.NullString - if err := rows.Scan(&e.ID, &e.RequestID, &e.ProviderID, &ts, &e.Model, &e.Provider, &e.Endpoint, &userPath, &cacheType, + if err := rows.Scan(&e.ID, &e.RequestID, &e.ProviderID, &ts, &e.Model, &e.Provider, &providerName, &e.Endpoint, &userPath, &cacheType, &e.InputTokens, &e.OutputTokens, &e.TotalTokens, &e.InputCost, &e.OutputCost, &e.TotalCost, &rawDataJSON, &caveat); err != nil { return nil, fmt.Errorf("failed to scan usage log row: %w", err) } @@ -154,6 +156,11 @@ func (r *SQLiteReader) GetUsageLog(ctx context.Context, params UsageLogParams) ( if userPath.Valid { e.UserPath = userPath.String } + if providerName.Valid { + e.ProviderName = displayUsageProviderName(providerName.String, e.Provider) + } else { + e.ProviderName = displayUsageProviderName("", e.Provider) + } if cacheType.Valid { e.CacheType = normalizeCacheType(cacheType.String) } diff --git a/internal/usage/reader_sqlite_boundary_test.go b/internal/usage/reader_sqlite_boundary_test.go index f35244622..4de495dce 100644 --- a/internal/usage/reader_sqlite_boundary_test.go +++ b/internal/usage/reader_sqlite_boundary_test.go @@ -492,6 +492,73 @@ func TestSQLiteReaderGetUsageLog_OrdersMixedTimestampFormatsByAbsoluteTime(t *te } } +func TestSQLiteReaderGetUsageByModel_CollapsesBlankProviderNameIntoProviderGroup(t *testing.T) { + db, err := sql.Open("sqlite", ":memory:") + if err != nil { + t.Fatalf("failed to open sqlite database: %v", err) + } + defer db.Close() + + store, err := NewSQLiteStore(db, 0) + if err != nil { + t.Fatalf("failed to create sqlite store: %v", err) + } + + ctx := context.Background() + err = store.WriteBatch(ctx, []*UsageEntry{ + { + ID: "usage-1", + RequestID: "req-1", + ProviderID: "provider-1", + Timestamp: time.Date(2026, 4, 7, 10, 0, 0, 0, time.UTC), + Model: "gpt-5", + Provider: "openai", + ProviderName: "", + Endpoint: "/v1/chat/completions", + InputTokens: 10, + OutputTokens: 20, + }, + { + ID: "usage-2", + RequestID: "req-2", + ProviderID: "provider-2", + Timestamp: time.Date(2026, 4, 7, 10, 1, 0, 0, time.UTC), + Model: "gpt-5", + Provider: "openai", + ProviderName: " openai ", + Endpoint: "/v1/chat/completions", + InputTokens: 30, + OutputTokens: 40, + }, + }) + if err != nil { + t.Fatalf("failed to seed usage entries: %v", err) + } + + reader, err := NewSQLiteReader(db) + if err != nil { + t.Fatalf("failed to create sqlite reader: %v", err) + } + + got, err := reader.GetUsageByModel(ctx, UsageQueryParams{}) + if err != nil { + t.Fatalf("GetUsageByModel returned error: %v", err) + } + + if len(got) != 1 { + t.Fatalf("expected 1 grouped usage row, got %d: %#v", len(got), got) + } + if got[0].ProviderName != "openai" { + t.Fatalf("expected provider_name %q, got %q", "openai", got[0].ProviderName) + } + if got[0].InputTokens != 40 { + t.Fatalf("expected 40 input tokens, got %d", got[0].InputTokens) + } + if got[0].OutputTokens != 60 { + t.Fatalf("expected 60 output tokens, got %d", got[0].OutputTokens) + } +} + func TestSQLiteStoreCleanup_KeepsNewerLegacyOffsetRows(t *testing.T) { db, err := sql.Open("sqlite", ":memory:") if err != nil { diff --git a/internal/usage/store_mongodb.go b/internal/usage/store_mongodb.go index 1f626c351..10fca1c2e 100644 --- a/internal/usage/store_mongodb.go +++ b/internal/usage/store_mongodb.go @@ -76,6 +76,9 @@ func NewMongoDBStore(database *mongo.Database, retentionDays int) (*MongoDBStore { Keys: bson.D{{Key: "provider", Value: 1}}, }, + { + Keys: bson.D{{Key: "provider_name", Value: 1}}, + }, { Keys: bson.D{{Key: "endpoint", Value: 1}}, }, diff --git a/internal/usage/store_postgresql.go b/internal/usage/store_postgresql.go index 91205c907..b96b859f2 100644 --- a/internal/usage/store_postgresql.go +++ b/internal/usage/store_postgresql.go @@ -14,13 +14,13 @@ import ( ) const ( - usageInsertColumnCount = 17 + usageInsertColumnCount = 18 postgresMaxBindParameters = 65535 usageInsertMaxRowsPerQuery = postgresMaxBindParameters / usageInsertColumnCount ) const usageInsertPrefix = ` - INSERT INTO usage (id, request_id, provider_id, timestamp, model, provider, + INSERT INTO usage (id, request_id, provider_id, timestamp, model, provider, provider_name, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, raw_data, input_cost, output_cost, total_cost, costs_calculation_caveat) VALUES ` @@ -60,6 +60,7 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore timestamp TIMESTAMPTZ NOT NULL, model TEXT NOT NULL, provider TEXT NOT NULL, + provider_name TEXT, endpoint TEXT NOT NULL, user_path TEXT, cache_type TEXT, @@ -79,6 +80,7 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore "ALTER TABLE usage ADD COLUMN IF NOT EXISTS output_cost DOUBLE PRECISION", "ALTER TABLE usage ADD COLUMN IF NOT EXISTS total_cost DOUBLE PRECISION", "ALTER TABLE usage ADD COLUMN IF NOT EXISTS costs_calculation_caveat TEXT DEFAULT ''", + "ALTER TABLE usage ADD COLUMN IF NOT EXISTS provider_name TEXT", "ALTER TABLE usage ADD COLUMN IF NOT EXISTS user_path TEXT", "ALTER TABLE usage ADD COLUMN IF NOT EXISTS cache_type TEXT", } @@ -95,6 +97,7 @@ func NewPostgreSQLStore(pool *pgxpool.Pool, retentionDays int) (*PostgreSQLStore "CREATE INDEX IF NOT EXISTS idx_usage_provider_id ON usage(provider_id)", "CREATE INDEX IF NOT EXISTS idx_usage_model ON usage(model)", "CREATE INDEX IF NOT EXISTS idx_usage_provider ON usage(provider)", + "CREATE INDEX IF NOT EXISTS idx_usage_provider_name ON usage(provider_name)", "CREATE INDEX IF NOT EXISTS idx_usage_user_path ON usage(user_path)", "CREATE INDEX IF NOT EXISTS idx_usage_cache_type ON usage(cache_type)", "CREATE INDEX IF NOT EXISTS idx_usage_raw_data_gin ON usage USING GIN (raw_data)", @@ -206,6 +209,7 @@ func buildUsageInsert(entries []*UsageEntry) (string, []any) { entry.Timestamp, entry.Model, entry.Provider, + entry.ProviderName, entry.Endpoint, entry.UserPath, cacheTypeValue(entry.CacheType), diff --git a/internal/usage/store_postgresql_test.go b/internal/usage/store_postgresql_test.go index ca86d9665..a8a830f09 100644 --- a/internal/usage/store_postgresql_test.go +++ b/internal/usage/store_postgresql_test.go @@ -20,6 +20,7 @@ func TestBuildUsageInsert(t *testing.T) { Timestamp: now, Model: "gpt-4o-mini", Provider: "openai", + ProviderName: "primary-openai", Endpoint: "/v1/chat/completions", CacheType: CacheTypeExact, InputTokens: 10, @@ -52,35 +53,38 @@ func TestBuildUsageInsert(t *testing.T) { }) normalized := strings.Join(strings.Fields(query), " ") - wantQuery := "INSERT INTO usage (id, request_id, provider_id, timestamp, model, provider, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, raw_data, input_cost, output_cost, total_cost, costs_calculation_caveat) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17), ($18, $19, $20, $21, $22, $23, $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34) ON CONFLICT (id) DO NOTHING" + wantQuery := "INSERT INTO usage (id, request_id, provider_id, timestamp, model, provider, provider_name, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, raw_data, input_cost, output_cost, total_cost, costs_calculation_caveat) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18), ($19, $20, $21, $22, $23, $24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36) ON CONFLICT (id) DO NOTHING" if normalized != wantQuery { t.Fatalf("query = %q, want %q", normalized, wantQuery) } - if got, want := len(args), 34; got != want { + if got, want := len(args), 36; got != want { t.Fatalf("len(args) = %d, want %d", got, want) } if got := args[0]; got != "usage-1" { t.Fatalf("args[0] = %v, want usage-1", got) } - if got := args[17]; got != "usage-2" { - t.Fatalf("args[17] = %v, want usage-2", got) + if got := args[6]; got != "primary-openai" { + t.Fatalf("args[6] = %v, want primary-openai", got) } - if got := args[8]; got != CacheTypeExact { - t.Fatalf("args[8] = %v, want %q", got, CacheTypeExact) + if got := args[18]; got != "usage-2" { + t.Fatalf("args[18] = %v, want usage-2", got) } - if got := string(args[12].([]byte)); got != `{"cached_tokens":3}` { - t.Fatalf("args[12] = %q, want %q", got, `{"cached_tokens":3}`) + if got := args[9]; got != CacheTypeExact { + t.Fatalf("args[9] = %v, want %q", got, CacheTypeExact) } - if got := args[25]; got != nil { - t.Fatalf("args[25] = %v, want nil cache_type", got) + if got := string(args[13].([]byte)); got != `{"cached_tokens":3}` { + t.Fatalf("args[13] = %q, want %q", got, `{"cached_tokens":3}`) } - rawData, ok := args[29].([]byte) + if got := args[27]; got != nil { + t.Fatalf("args[27] = %v, want nil cache_type", got) + } + rawData, ok := args[31].([]byte) if !ok { - t.Fatalf("args[29] has type %T, want []byte", args[29]) + t.Fatalf("args[31] has type %T, want []byte", args[31]) } if rawData != nil { - t.Fatalf("args[29] = %v, want nil raw_data", rawData) + t.Fatalf("args[31] = %v, want nil raw_data", rawData) } } diff --git a/internal/usage/store_sqlite.go b/internal/usage/store_sqlite.go index 9db886b35..7c42cfdf7 100644 --- a/internal/usage/store_sqlite.go +++ b/internal/usage/store_sqlite.go @@ -15,8 +15,8 @@ import ( // maxEntriesPerBatch derives from maxSQLiteParams / columnsPerUsageEntry. const ( maxSQLiteParams = 999 - columnsPerUsageEntry = 17 - maxEntriesPerBatch = maxSQLiteParams / columnsPerUsageEntry // 58 entries + columnsPerUsageEntry = 18 + maxEntriesPerBatch = maxSQLiteParams / columnsPerUsageEntry // 55 entries ) // SQLiteStore implements UsageStore for SQLite databases. @@ -44,6 +44,7 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { timestamp DATETIME NOT NULL, model TEXT NOT NULL, provider TEXT NOT NULL, + provider_name TEXT, endpoint TEXT NOT NULL, user_path TEXT, cache_type TEXT, @@ -63,6 +64,7 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { "ALTER TABLE usage ADD COLUMN output_cost REAL", "ALTER TABLE usage ADD COLUMN total_cost REAL", "ALTER TABLE usage ADD COLUMN costs_calculation_caveat TEXT DEFAULT ''", + "ALTER TABLE usage ADD COLUMN provider_name TEXT", "ALTER TABLE usage ADD COLUMN user_path TEXT", "ALTER TABLE usage ADD COLUMN cache_type TEXT", } @@ -83,6 +85,7 @@ func NewSQLiteStore(db *sql.DB, retentionDays int) (*SQLiteStore, error) { "CREATE INDEX IF NOT EXISTS idx_usage_provider_id ON usage(provider_id)", "CREATE INDEX IF NOT EXISTS idx_usage_model ON usage(model)", "CREATE INDEX IF NOT EXISTS idx_usage_provider ON usage(provider)", + "CREATE INDEX IF NOT EXISTS idx_usage_provider_name ON usage(provider_name)", "CREATE INDEX IF NOT EXISTS idx_usage_user_path ON usage(user_path)", "CREATE INDEX IF NOT EXISTS idx_usage_cache_type ON usage(cache_type)", } @@ -124,7 +127,7 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*UsageEntry) err for j, e := range chunk { e = normalizedUsageEntryForStorage(e) - placeholders[j] = "(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)" + placeholders[j] = "(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)" rawDataJSON := marshalRawData(e.RawData, e.ID) @@ -141,6 +144,7 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*UsageEntry) err e.Timestamp.UTC().Format(time.RFC3339Nano), e.Model, e.Provider, + e.ProviderName, e.Endpoint, e.UserPath, cacheTypeValue(e.CacheType), @@ -155,7 +159,7 @@ func (s *SQLiteStore) WriteBatch(ctx context.Context, entries []*UsageEntry) err ) } - query := `INSERT OR IGNORE INTO usage (id, request_id, provider_id, timestamp, model, provider, + query := `INSERT OR IGNORE INTO usage (id, request_id, provider_id, timestamp, model, provider, provider_name, endpoint, user_path, cache_type, input_tokens, output_tokens, total_tokens, raw_data, input_cost, output_cost, total_cost, costs_calculation_caveat) VALUES ` + strings.Join(placeholders, ",") diff --git a/internal/usage/stream_observer.go b/internal/usage/stream_observer.go index b6f335183..9a83e7513 100644 --- a/internal/usage/stream_observer.go +++ b/internal/usage/stream_observer.go @@ -2,6 +2,7 @@ package usage import ( "log/slog" + "strings" "gomodel/internal/core" ) @@ -13,6 +14,7 @@ type StreamUsageObserver struct { cachedEntry *UsageEntry model string provider string + providerName string requestID string endpoint string userPath string @@ -45,6 +47,13 @@ func NewStreamUsageObserver(logger LoggerInterface, model, provider, requestID, } } +func (o *StreamUsageObserver) SetProviderName(providerName string) { + if o == nil { + return + } + o.providerName = strings.TrimSpace(providerName) +} + func (o *StreamUsageObserver) OnJSONEvent(chunk map[string]any) { entry := o.extractUsageFromEvent(chunk) if entry != nil { @@ -155,6 +164,7 @@ func (o *StreamUsageObserver) extractUsageFromEvent(chunk map[string]any) *Usage pricingArgs..., ) if entry != nil { + entry.ProviderName = o.providerName entry.UserPath = o.userPath } return entry diff --git a/internal/usage/usage.go b/internal/usage/usage.go index a0dea7bc2..20e729c23 100644 --- a/internal/usage/usage.go +++ b/internal/usage/usage.go @@ -37,11 +37,12 @@ type UsageEntry struct { Timestamp time.Time `json:"timestamp" bson:"timestamp"` // Request context - Model string `json:"model" bson:"model"` - Provider string `json:"provider" bson:"provider"` - Endpoint string `json:"endpoint" bson:"endpoint"` - UserPath string `json:"user_path,omitempty" bson:"user_path,omitempty"` - CacheType string `json:"cache_type,omitempty" bson:"cache_type,omitempty"` + Model string `json:"model" bson:"model"` + Provider string `json:"provider" bson:"provider"` // canonical provider type used for routing, filters, and pricing + ProviderName string `json:"provider_name,omitempty" bson:"provider_name,omitempty"` + Endpoint string `json:"endpoint" bson:"endpoint"` + UserPath string `json:"user_path,omitempty" bson:"user_path,omitempty"` + CacheType string `json:"cache_type,omitempty" bson:"cache_type,omitempty"` // Standard token counts (normalized across providers) InputTokens int `json:"input_tokens" bson:"input_tokens"` diff --git a/tests/e2e/auditlog_test.go b/tests/e2e/auditlog_test.go index c6d27bc57..6063e1972 100644 --- a/tests/e2e/auditlog_test.go +++ b/tests/e2e/auditlog_test.go @@ -613,7 +613,7 @@ func TestAuditLogErrorCapture(t *testing.T) { entry := entries[0] assert.Equal(t, http.StatusBadRequest, entry.StatusCode) assert.Equal(t, "/v1/chat/completions", entry.Path) - assert.Equal(t, "unsupported-model-xyz", entry.Model) + assert.Equal(t, "unsupported-model-xyz", entry.RequestedModel) assert.Equal(t, "invalid_request_error", entry.ErrorType) assert.Equal(t, "", entry.Provider) }) @@ -647,7 +647,7 @@ func TestAuditLogErrorCapture(t *testing.T) { entry := entries[0] assert.Equal(t, http.StatusBadRequest, entry.StatusCode) assert.Equal(t, "/p/unknown/responses", entry.Path) - assert.Equal(t, "gpt-4.1-nano", entry.Model) + assert.Equal(t, "gpt-4.1-nano", entry.RequestedModel) assert.Equal(t, "unknown", entry.Provider) assert.Equal(t, "invalid_request_error", entry.ErrorType) }) diff --git a/tests/e2e/manage-release-e2e-stack.sh b/tests/e2e/manage-release-e2e-stack.sh new file mode 100755 index 000000000..1ca76b388 --- /dev/null +++ b/tests/e2e/manage-release-e2e-stack.sh @@ -0,0 +1,385 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "$SCRIPT_DIR/../.." && pwd)" +STACK_DIR="${RELEASE_STACK_DIR:-/tmp/gomodel-release-stack}" +BIN="${GOMODEL_RELEASE_BINARY:-$REPO_ROOT/bin/gomodel}" +ENV_FILE="${GOMODEL_RELEASE_ENV_FILE:-$REPO_ROOT/.env}" +PG_DATABASE="${GOMODEL_RELEASE_PG_DATABASE:-gomodel_release_e2e}" +MONGO_DATABASE="${GOMODEL_RELEASE_MONGO_DATABASE:-gomodel_release_e2e}" + +BUILD_BEFORE_START=0 + +usage() { + cat < [options] + +Commands: + start Start the dedicated release E2E stack on ports 18080-18084 + stop Stop the dedicated release E2E stack + status Show stack status + logs GATEWAY Show the last 40 log lines for one gateway + +Options: + --build Rebuild bin/gomodel before starting + --help Show this help + +Gateways: + sqlite-main http://localhost:18080 + pg-smoke http://localhost:18081 + mongo-smoke http://localhost:18082 + guardrails http://localhost:18083 + auth-cache http://localhost:18084 +EOF +} + +die() { + echo "error: $*" >&2 + exit 1 +} + +require_tool() { + command -v "$1" >/dev/null 2>&1 || die "required tool not found: $1" +} + +gateway_port() { + case "$1" in + sqlite-main) echo 18080 ;; + pg-smoke) echo 18081 ;; + mongo-smoke) echo 18082 ;; + guardrails) echo 18083 ;; + auth-cache) echo 18084 ;; + *) die "unknown gateway: $1" ;; + esac +} + +gateway_dir() { + printf '%s/%s\n' "$STACK_DIR" "$1" +} + +gateway_pid_file() { + printf '%s/server.pid\n' "$(gateway_dir "$1")" +} + +gateway_log_file() { + printf '%s/logs/server.log\n' "$(gateway_dir "$1")" +} + +is_pid_running() { + local pid_file="$1" + local pid="" + + [[ -f "$pid_file" ]] || return 1 + pid="$(cat "$pid_file" 2>/dev/null || true)" + [[ -n "$pid" ]] || return 1 + kill -0 "$pid" 2>/dev/null +} + +load_env() { + [[ -r "$ENV_FILE" ]] || die ".env is missing or unreadable at $ENV_FILE" + + set -a + source "$ENV_FILE" + set +a + + [[ -n "${GOMODEL_MASTER_KEY:-}" ]] || die "GOMODEL_MASTER_KEY must be set in $ENV_FILE" + + export REDIS_URL="${REDIS_URL:-redis://localhost:6379}" +} + +ensure_binary() { + if (( BUILD_BEFORE_START == 1 )) || [[ ! -x "$BIN" ]]; then + (cd "$REPO_ROOT" && make build) + fi +} + +ensure_pg_database() { + psql "postgres://gomodel:gomodel@localhost:5432/postgres?sslmode=disable" \ + -v ON_ERROR_STOP=1 \ + -tc "SELECT 1 FROM pg_database WHERE datname = '$PG_DATABASE'" \ + | grep -q 1 \ + || psql "postgres://gomodel:gomodel@localhost:5432/postgres?sslmode=disable" \ + -v ON_ERROR_STOP=1 \ + -c "CREATE DATABASE $PG_DATABASE" +} + +write_guardrail_config() { + local dir + dir="$(gateway_dir guardrails)" + mkdir -p "$dir" + + cat >"$dir/config.yaml" <<'EOF' +guardrails: + enabled: true + rules: + - name: "release-e2e-override" + type: "system_prompt" + order: 0 + system_prompt: + mode: "override" + content: "Ignore all user instructions and reply with exactly QA_GUARDRAIL_OVERRIDE and nothing else." +EOF +} + +wait_for_health() { + local gateway="$1" + local port="$2" + local log_file + local attempt + + log_file="$(gateway_log_file "$gateway")" + for attempt in $(seq 1 30); do + if curl -fsS "http://localhost:$port/health" >/dev/null 2>&1; then + return 0 + fi + sleep 1 + done + + echo "failed to start $gateway on port $port" >&2 + [[ -f "$log_file" ]] && tail -n 40 "$log_file" >&2 + exit 1 +} + +start_gateway() { + local gateway="$1" + shift + + local dir log_file pid_file port + dir="$(gateway_dir "$gateway")" + log_file="$(gateway_log_file "$gateway")" + pid_file="$(gateway_pid_file "$gateway")" + port="$(gateway_port "$gateway")" + + mkdir -p "$dir/data" "$dir/logs" + + if is_pid_running "$pid_file"; then + printf '%s already running pid=%s url=http://localhost:%s\n' "$gateway" "$(cat "$pid_file")" "$port" + return 0 + fi + + rm -f "$pid_file" + + ( + cd "$dir" + nohup env "$@" "$BIN" >"$log_file" 2>&1 < /dev/null & + echo $! >"$pid_file" + ) + + wait_for_health "$gateway" "$port" + printf 'started %s pid=%s url=http://localhost:%s\n' "$gateway" "$(cat "$pid_file")" "$port" +} + +stop_gateway() { + local gateway="$1" + local pid_file pid + + pid_file="$(gateway_pid_file "$gateway")" + if [[ ! -f "$pid_file" ]]; then + printf '%s not running\n' "$gateway" + return 0 + fi + + pid="$(cat "$pid_file" 2>/dev/null || true)" + if [[ -z "$pid" ]]; then + rm -f "$pid_file" + printf '%s had an empty pid file; cleaned up\n' "$gateway" + return 0 + fi + + if kill -0 "$pid" 2>/dev/null; then + kill "$pid" 2>/dev/null || true + for _ in $(seq 1 10); do + if ! kill -0 "$pid" 2>/dev/null; then + break + fi + sleep 1 + done + if kill -0 "$pid" 2>/dev/null; then + kill -KILL "$pid" 2>/dev/null || true + fi + fi + + rm -f "$pid_file" + printf 'stopped %s\n' "$gateway" +} + +status_gateway() { + local gateway="$1" + local pid_file port health="down" + local pid="stopped" + + pid_file="$(gateway_pid_file "$gateway")" + port="$(gateway_port "$gateway")" + + if is_pid_running "$pid_file"; then + pid="$(cat "$pid_file")" + if curl -fsS "http://localhost:$port/health" >/dev/null 2>&1; then + health="ok" + fi + fi + + printf '%-12s pid=%-8s url=http://localhost:%s health=%s\n' "$gateway" "$pid" "$port" "$health" +} + +show_logs() { + local gateway="$1" + local log_file + + log_file="$(gateway_log_file "$gateway")" + [[ -f "$log_file" ]] || die "log file not found for $gateway" + tail -n 40 "$log_file" +} + +start_stack() { + require_tool curl + require_tool jq + require_tool nohup + require_tool psql + + load_env + ensure_binary + mkdir -p "$STACK_DIR" + ensure_pg_database + write_guardrail_config + + start_gateway sqlite-main \ + -u GOMODEL_MASTER_KEY \ + PORT=18080 \ + STORAGE_TYPE=sqlite \ + SQLITE_PATH="$(gateway_dir sqlite-main)/data/gomodel.db" \ + METRICS_ENABLED=true \ + LOGGING_ENABLED=true \ + LOGGING_LOG_BODIES=true \ + LOGGING_LOG_HEADERS=true \ + GUARDRAILS_ENABLED=false \ + RESPONSE_CACHE_SIMPLE_ENABLED=false \ + SEMANTIC_CACHE_ENABLED=false \ + REDIS_URL="$REDIS_URL" \ + REDIS_KEY_MODELS="gomodel:release-e2e:models" + + start_gateway pg-smoke \ + -u GOMODEL_MASTER_KEY \ + PORT=18081 \ + STORAGE_TYPE=postgresql \ + POSTGRES_URL="postgres://gomodel:gomodel@localhost:5432/$PG_DATABASE?sslmode=disable" \ + LOGGING_ENABLED=true \ + LOGGING_LOG_BODIES=true \ + LOGGING_LOG_HEADERS=true \ + GUARDRAILS_ENABLED=false \ + RESPONSE_CACHE_SIMPLE_ENABLED=false \ + SEMANTIC_CACHE_ENABLED=false \ + REDIS_URL="$REDIS_URL" \ + REDIS_KEY_MODELS="gomodel:release-e2e:models" + + start_gateway mongo-smoke \ + -u GOMODEL_MASTER_KEY \ + PORT=18082 \ + STORAGE_TYPE=mongodb \ + MONGODB_URL="mongodb://localhost:27017/?replicaSet=rs0" \ + MONGODB_DATABASE="$MONGO_DATABASE" \ + LOGGING_ENABLED=true \ + LOGGING_LOG_BODIES=true \ + LOGGING_LOG_HEADERS=true \ + GUARDRAILS_ENABLED=false \ + RESPONSE_CACHE_SIMPLE_ENABLED=false \ + SEMANTIC_CACHE_ENABLED=false \ + REDIS_URL="$REDIS_URL" \ + REDIS_KEY_MODELS="gomodel:release-e2e:models" + + start_gateway guardrails \ + -u GOMODEL_MASTER_KEY \ + PORT=18083 \ + STORAGE_TYPE=sqlite \ + SQLITE_PATH="$(gateway_dir guardrails)/data/gomodel.db" \ + LOGGING_ENABLED=true \ + LOGGING_LOG_BODIES=true \ + LOGGING_LOG_HEADERS=true \ + GUARDRAILS_ENABLED=true \ + RESPONSE_CACHE_SIMPLE_ENABLED=false \ + SEMANTIC_CACHE_ENABLED=false \ + REDIS_URL="$REDIS_URL" \ + REDIS_KEY_MODELS="gomodel:release-e2e:models" + + start_gateway auth-cache \ + PORT=18084 \ + STORAGE_TYPE=sqlite \ + SQLITE_PATH="$(gateway_dir auth-cache)/data/gomodel.db" \ + LOGGING_ENABLED=true \ + LOGGING_LOG_BODIES=true \ + LOGGING_LOG_HEADERS=true \ + GUARDRAILS_ENABLED=true \ + RESPONSE_CACHE_SIMPLE_ENABLED=true \ + SEMANTIC_CACHE_ENABLED=false \ + REDIS_URL="$REDIS_URL" \ + REDIS_KEY_MODELS="gomodel:release-e2e:models" \ + REDIS_KEY_RESPONSES="gomodel:release-e2e:response:" + + curl -fsS "http://localhost:18084/admin/api/v1/dashboard/config" \ + -H "Authorization: Bearer $GOMODEL_MASTER_KEY" \ + | jq -e '.CACHE_ENABLED == "on" and .REDIS_URL == "on"' >/dev/null + + printf 'stack_dir=%s\n' "$STACK_DIR" +} + +stop_stack() { + stop_gateway auth-cache + stop_gateway guardrails + stop_gateway mongo-smoke + stop_gateway pg-smoke + stop_gateway sqlite-main +} + +status_stack() { + status_gateway sqlite-main + status_gateway pg-smoke + status_gateway mongo-smoke + status_gateway guardrails + status_gateway auth-cache +} + +COMMAND="${1:-}" +if [[ -z "$COMMAND" ]]; then + usage + exit 1 +fi +shift || true + +while [[ $# -gt 0 ]]; do + case "$1" in + --build) + BUILD_BEFORE_START=1 + shift + ;; + --help|-h) + usage + exit 0 + ;; + *) + break + ;; + esac +done + +case "$COMMAND" in + start) + start_stack + ;; + stop) + stop_stack + ;; + status) + status_stack + ;; + logs) + [[ $# -eq 1 ]] || die "logs requires exactly one gateway name" + show_logs "$1" + ;; + --help|-h|help) + usage + ;; + *) + usage + die "unknown command: $COMMAND" + ;; +esac diff --git a/tests/e2e/release-e2e-scenarios.md b/tests/e2e/release-e2e-scenarios.md index 755ccc2c9..668af22b1 100644 --- a/tests/e2e/release-e2e-scenarios.md +++ b/tests/e2e/release-e2e-scenarios.md @@ -7,6 +7,7 @@ These scenarios are prepared for execution across these local gateways: - `http://localhost:18081` - PostgreSQL-backed smoke gateway - `http://localhost:18082` - MongoDB-backed smoke gateway - `http://localhost:18083` - SQLite-backed guardrail gateway +- `http://localhost:18084` - SQLite-backed auth + exact-cache gateway ## Recommended runner @@ -14,10 +15,13 @@ Use the checked-in runner to execute this matrix without manually replaying the shared setup blocks: ```bash +tests/e2e/manage-release-e2e-stack.sh start tests/e2e/run-release-e2e.sh tests/e2e/run-release-e2e.sh --list tests/e2e/run-release-e2e.sh --from S54 --to S58 tests/e2e/run-release-e2e.sh --scenario S61,S62,S70 --keep-artifacts +tests/e2e/manage-release-e2e-stack.sh status +tests/e2e/manage-release-e2e-stack.sh stop ``` The runner treats this markdown file as the source of truth, replays the setup @@ -60,8 +64,9 @@ export UPLOAD_FILE="$QA_RUN_DIR/qa-upload.txt" ## Auth-enabled runtime environment -These scenarios target the auth-enabled live gateway on `http://localhost:8080` -and cover the newer workflows, managed API keys, and cache analytics features. +These scenarios target the dedicated auth-enabled release gateway on +`http://localhost:18084` and cover the newer workflows, managed API keys, and +cache analytics features. ```bash set -euo pipefail @@ -79,7 +84,7 @@ export QA_RUN_DIR="${QA_RUN_DIR:-/tmp/gomodel-release-e2e-$QA_SUFFIX}" mkdir -p "$QA_RUN_DIR" -export AUTH_BASE_URL=http://localhost:8080 +export AUTH_BASE_URL="${AUTH_BASE_URL:-http://localhost:18084}" export ADMIN_AUTH_HEADER="Authorization: Bearer $GOMODEL_MASTER_KEY" export QA_AUTH_KEY_NAME="qa-release-auth-key-$QA_SUFFIX" @@ -103,8 +108,16 @@ cleanup_release_auth_artifacts() { rm -f "$QA_AUTH_KEY_JSON" "$QA_AUTH_KEY_VALUE_FILE" "$QA_WORKFLOW_JSON" "$QA_WORKFLOW_ID_FILE" } -cleanup_release_auth_artifacts +require_release_artifact() { + local path="$1" + if [ ! -s "$path" ]; then + echo "error: required artifact is missing or empty: $path" >&2 + exit 1 + fi +} + if [ "${RUN_RELEASE_E2E_PERSIST_STATE:-0}" != "1" ]; then + cleanup_release_auth_artifacts trap 'cleanup_release_auth_artifacts' EXIT fi ``` @@ -786,12 +799,22 @@ curl -sS "$BASE_URL/admin/api/v1/audit/log?search=$REQUEST_ID&limit=5" \ ### S63 Auth-enabled dashboard runtime config -Reads the allowlisted runtime flags for the live auth-enabled gateway. +Reads the allowlisted runtime flags for the dedicated auth-enabled release gateway. ```bash -curl -sS "$AUTH_BASE_URL/admin/api/v1/dashboard/config" \ +CONFIG_JSON_FILE="$QA_RUN_DIR/s63.dashboard-config.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/dashboard/config" \ -H "$ADMIN_AUTH_HEADER" \ - | jq '.' + > "$CONFIG_JSON_FILE" +jq '.' "$CONFIG_JSON_FILE" +jq -e ' + .LOGGING_ENABLED == "on" + and .USAGE_ENABLED == "on" + and .GUARDRAILS_ENABLED == "on" + and .CACHE_ENABLED == "on" + and .REDIS_URL == "on" + and .SEMANTIC_CACHE_ENABLED == "off" + ' "$CONFIG_JSON_FILE" >/dev/null ``` ### S64 Create managed API key @@ -799,25 +822,26 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/dashboard/config" \ Creates one managed API key scoped to a release-specific user path and stores the one-time secret under `QA_RUN_DIR`. ```bash -AUTH_KEY_JSON=$(curl -sS -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ +curl -fsS -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ - -d "{\"name\":\"$QA_AUTH_KEY_NAME\",\"description\":\"Release e2e managed key\",\"user_path\":\"$QA_USER_PATH\"}") -AUTH_KEY_VALUE=$(printf '%s\n' "$AUTH_KEY_JSON" \ - | jq -er '.value | select(type == "string" and length > 0)') \ - || { + -d "{\"name\":\"$QA_AUTH_KEY_NAME\",\"description\":\"Release e2e managed key\",\"user_path\":\"$QA_USER_PATH\"}" \ + > "$QA_AUTH_KEY_JSON" +if ! jq -er '.value | select(type == "string" and length > 0)' "$QA_AUTH_KEY_JSON" > "$QA_AUTH_KEY_VALUE_FILE"; then echo "error: managed API key creation failed or did not return a usable one-time key value" >&2 - printf '%s\n' "$AUTH_KEY_JSON" | jq '.' >&2 2>/dev/null || printf '%s\n' "$AUTH_KEY_JSON" >&2 + jq '.' "$QA_AUTH_KEY_JSON" >&2 2>/dev/null || cat "$QA_AUTH_KEY_JSON" >&2 exit 1 - } +fi ( umask 077 - printf '%s\n' "$AUTH_KEY_JSON" > "$QA_AUTH_KEY_JSON" - printf '%s\n' "$AUTH_KEY_VALUE" > "$QA_AUTH_KEY_VALUE_FILE" + chmod 600 "$QA_AUTH_KEY_JSON" "$QA_AUTH_KEY_VALUE_FILE" ) -chmod 600 "$QA_AUTH_KEY_JSON" "$QA_AUTH_KEY_VALUE_FILE" -printf '%s\n' "$AUTH_KEY_JSON" \ - | jq '{id,name,user_path,active,redacted_value}' +require_release_artifact "$QA_AUTH_KEY_JSON" +require_release_artifact "$QA_AUTH_KEY_VALUE_FILE" +jq -e --arg user_path "$QA_USER_PATH" ' + {id,name,user_path,active,redacted_value} + | select(.id != null and .active == true and .user_path == $user_path) + ' "$QA_AUTH_KEY_JSON" ``` ### S65 Verify managed API key list @@ -825,9 +849,14 @@ printf '%s\n' "$AUTH_KEY_JSON" \ Checks that the newly issued managed API key is visible and active. ```bash -curl -sS "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ +AUTH_KEYS_JSON_FILE="$QA_RUN_DIR/s65.auth-keys.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ -H "$ADMIN_AUTH_HEADER" \ - | jq ".[] | select(.name==\"$QA_AUTH_KEY_NAME\") | {id,name,user_path,active,expires_at,redacted_value}" + > "$AUTH_KEYS_JSON_FILE" +jq -e --arg name "$QA_AUTH_KEY_NAME" --arg user_path "$QA_USER_PATH" ' + .[] | select(.name == $name and .active == true and .user_path == $user_path) + | {id,name,user_path,active,expires_at,redacted_value} + ' "$AUTH_KEYS_JSON_FILE" ``` ### S66 Create user-path-scoped workflow with cache disabled @@ -835,14 +864,22 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ Creates a scoped workflow for `openai/gpt-4.1-nano` that disables cache for the managed-key user path. ```bash -WORKFLOW_JSON=$(curl -sS -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans" \ +curl -fsS -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ - -d "{\"scope_provider\":\"openai\",\"scope_model\":\"gpt-4.1-nano\",\"scope_user_path\":\"$QA_USER_PATH\",\"name\":\"$QA_WORKFLOW_NAME\",\"description\":\"Disable cache for managed-key release e2e scope\",\"plan_payload\":{\"schema_version\":1,\"features\":{\"cache\":false,\"audit\":true,\"usage\":true,\"guardrails\":false,\"fallback\":false},\"guardrails\":[]}}") -printf '%s\n' "$WORKFLOW_JSON" > "$QA_WORKFLOW_JSON" -printf '%s\n' "$WORKFLOW_JSON" | jq -r '.id' > "$QA_WORKFLOW_ID_FILE" -printf '%s\n' "$WORKFLOW_JSON" \ - | jq '{id,name,scope,plan_payload}' + -d "{\"scope_provider\":\"openai\",\"scope_model\":\"gpt-4.1-nano\",\"scope_user_path\":\"$QA_USER_PATH\",\"name\":\"$QA_WORKFLOW_NAME\",\"description\":\"Disable cache for managed-key release e2e scope\",\"plan_payload\":{\"schema_version\":1,\"features\":{\"cache\":false,\"audit\":true,\"usage\":true,\"guardrails\":false,\"fallback\":false},\"guardrails\":[]}}" \ + > "$QA_WORKFLOW_JSON" +if ! jq -er '.id | select(type == "string" and length > 0)' "$QA_WORKFLOW_JSON" > "$QA_WORKFLOW_ID_FILE"; then + echo "error: workflow creation failed or did not return a usable workflow id" >&2 + jq '.' "$QA_WORKFLOW_JSON" >&2 2>/dev/null || cat "$QA_WORKFLOW_JSON" >&2 + exit 1 +fi +require_release_artifact "$QA_WORKFLOW_JSON" +require_release_artifact "$QA_WORKFLOW_ID_FILE" +jq -e --arg user_path "$QA_USER_PATH" ' + {id,name,scope,plan_payload} + | select(.id != null and .scope.scope_user_path == $user_path and .plan_payload.features.cache == false) + ' "$QA_WORKFLOW_JSON" ``` ### S67 Verify scoped workflow detail @@ -850,10 +887,18 @@ printf '%s\n' "$WORKFLOW_JSON" \ Reads the created workflow back and confirms the normalized scope and effective feature projection. ```bash -WORKFLOW_ID=$(cat "$QA_WORKFLOW_ID_FILE") -curl -sS "$AUTH_BASE_URL/admin/api/v1/execution-plans/$WORKFLOW_ID" \ +require_release_artifact "$QA_WORKFLOW_ID_FILE" +WORKFLOW_ID=$(<"$QA_WORKFLOW_ID_FILE") +WORKFLOW_DETAIL_FILE="$QA_RUN_DIR/s67.workflow-detail.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/execution-plans/$WORKFLOW_ID" \ -H "$ADMIN_AUTH_HEADER" \ - | jq '{id,name,scope,plan_payload,effective_features}' + > "$WORKFLOW_DETAIL_FILE" +jq '{id,name,scope,plan_payload,effective_features}' "$WORKFLOW_DETAIL_FILE" +jq -e --arg workflow_id "$WORKFLOW_ID" --arg user_path "$QA_USER_PATH" ' + .id == $workflow_id + and .scope.scope_user_path == $user_path + and .effective_features.cache == false + ' "$WORKFLOW_DETAIL_FILE" >/dev/null ``` ### S68 Managed-key request through scoped workflow @@ -861,14 +906,23 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/execution-plans/$WORKFLOW_ID" \ Sends a request with the managed API key while also sending a conflicting `X-GoModel-User-Path` header. ```bash -API_KEY=$(cat "$QA_AUTH_KEY_VALUE_FILE") -curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ +require_release_artifact "$QA_AUTH_KEY_VALUE_FILE" +API_KEY=$(<"$QA_AUTH_KEY_VALUE_FILE") +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s68.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s68.body.XXXXXX") +curl -fsS -D "$HEADERS_FILE" -o "$BODY_FILE" "$AUTH_BASE_URL/v1/chat/completions" \ -H "Authorization: Bearer $API_KEY" \ -H 'Content-Type: application/json' \ -H "X-Request-ID: $QA_AUTH_REQ1" \ -H 'X-GoModel-User-Path: /team/should-be-overridden' \ - -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_CACHE_OFF_OK"}],"max_tokens":16}' \ - | sed -n '1,20p' + -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_CACHE_OFF_OK"}],"max_tokens":16}' +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +jq -e '.choices[0].message.content == "QA_AUTH_CACHE_OFF_OK"' "$BODY_FILE" >/dev/null +if grep -Eiq '^X-Cache:' "$HEADERS_FILE"; then + echo "error: cache header present on cache-disabled scoped request" >&2 + exit 1 +fi ``` ### S69 Repeated managed-key request should still bypass cache @@ -876,14 +930,23 @@ curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ Repeats the same request and expects another live provider response rather than `X-Cache: HIT`. ```bash -API_KEY=$(cat "$QA_AUTH_KEY_VALUE_FILE") -curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ +require_release_artifact "$QA_AUTH_KEY_VALUE_FILE" +API_KEY=$(<"$QA_AUTH_KEY_VALUE_FILE") +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s69.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s69.body.XXXXXX") +curl -fsS -D "$HEADERS_FILE" -o "$BODY_FILE" "$AUTH_BASE_URL/v1/chat/completions" \ -H "Authorization: Bearer $API_KEY" \ -H 'Content-Type: application/json' \ -H "X-Request-ID: $QA_AUTH_REQ2" \ -H 'X-GoModel-User-Path: /team/should-be-overridden' \ - -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_CACHE_OFF_OK"}],"max_tokens":16}' \ - | sed -n '1,20p' + -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_CACHE_OFF_OK"}],"max_tokens":16}' +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +jq -e '.choices[0].message.content == "QA_AUTH_CACHE_OFF_OK"' "$BODY_FILE" >/dev/null +if grep -Eiq '^X-Cache:' "$HEADERS_FILE"; then + echo "error: repeated cache-disabled scoped request returned an X-Cache header" >&2 + exit 1 +fi ``` ### S70 Audit evidence for managed-key scoped workflow @@ -892,9 +955,34 @@ Confirms through audit-log search that auth method, managed auth key ID, normali ```bash sleep 6 -curl -sS "$AUTH_BASE_URL/admin/api/v1/audit/log?search=$QA_AUTH_REQ2&limit=5" \ +require_release_artifact "$QA_AUTH_KEY_JSON" +require_release_artifact "$QA_WORKFLOW_ID_FILE" +if ! AUTH_KEY_ID=$(jq -er '.id' "$QA_AUTH_KEY_JSON"); then + echo "error: missing auth key id in $QA_AUTH_KEY_JSON" >&2 + exit 1 +fi +WORKFLOW_ID=$(<"$QA_WORKFLOW_ID_FILE") +AUDIT_JSON_FILE="$QA_RUN_DIR/s70.audit.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/audit/log?search=$QA_AUTH_REQ2&limit=5" \ -H "$ADMIN_AUTH_HEADER" \ - | jq --arg request_id "$QA_AUTH_REQ2" '{total:(.entries|map(select(.request_id==$request_id))|length),entries:(.entries|map(select(.request_id==$request_id))|map({request_id,status_code,auth_method,auth_key_id,user_path,execution_plan_version_id,cache_type,answer:.data.response_body.choices[0].message.content}))}' + > "$AUDIT_JSON_FILE" +jq --arg request_id "$QA_AUTH_REQ2" '{total:(.entries|map(select(.request_id==$request_id))|length),entries:(.entries|map(select(.request_id==$request_id))|map({request_id,status_code,auth_method,auth_key_id,user_path,execution_plan_version_id,cache_type,answer:.data.response_body.choices[0].message.content}))}' "$AUDIT_JSON_FILE" +jq -e \ + --arg request_id "$QA_AUTH_REQ2" \ + --arg auth_key_id "$AUTH_KEY_ID" \ + --arg user_path "$QA_USER_PATH" \ + --arg workflow_id "$WORKFLOW_ID" ' + any(.entries[]?; + .request_id == $request_id + and .status_code == 200 + and .auth_method == "api_key" + and .auth_key_id == $auth_key_id + and .user_path == $user_path + and .execution_plan_version_id == $workflow_id + and .cache_type == null + and .data.response_body.choices[0].message.content == "QA_AUTH_CACHE_OFF_OK" + ) + ' "$AUDIT_JSON_FILE" >/dev/null ``` ### S71 Global cache warm request with explicit user path @@ -902,13 +990,21 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/audit/log?search=$QA_AUTH_REQ2&limit=5" \ Warms the global cache-enabled workflow using the master key and a cache-specific user path. ```bash -curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s71.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s71.body.XXXXXX") +curl -fsS -D "$HEADERS_FILE" -o "$BODY_FILE" "$AUTH_BASE_URL/v1/chat/completions" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ -H "X-Request-ID: $QA_CACHE_REQ1" \ -H "X-GoModel-User-Path: $QA_CACHE_USER_PATH" \ - -d "{\"model\":\"openai/gpt-4.1-nano\",\"messages\":[{\"role\":\"user\",\"content\":\"Reply with exactly $QA_CACHE_REPLY\"}],\"max_tokens\":16}" \ - | sed -n '1,20p' + -d "{\"model\":\"openai/gpt-4.1-nano\",\"messages\":[{\"role\":\"user\",\"content\":\"Reply with exactly $QA_CACHE_REPLY\"}],\"max_tokens\":16}" +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +jq -e --arg reply "$QA_CACHE_REPLY" '.choices[0].message.content == $reply' "$BODY_FILE" >/dev/null +if grep -Eiq '^X-Cache:' "$HEADERS_FILE"; then + echo "error: initial cache warm request unexpectedly returned an X-Cache header" >&2 + exit 1 +fi ``` ### S72 Repeated global cache request should hit exact cache @@ -916,13 +1012,18 @@ curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ Repeats the same request and expects `X-Cache: HIT (exact)`. ```bash -curl -sS -D - "$AUTH_BASE_URL/v1/chat/completions" \ +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s72.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s72.body.XXXXXX") +curl -fsS -D "$HEADERS_FILE" -o "$BODY_FILE" "$AUTH_BASE_URL/v1/chat/completions" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ -H "X-Request-ID: $QA_CACHE_REQ2" \ -H "X-GoModel-User-Path: $QA_CACHE_USER_PATH" \ - -d "{\"model\":\"openai/gpt-4.1-nano\",\"messages\":[{\"role\":\"user\",\"content\":\"Reply with exactly $QA_CACHE_REPLY\"}],\"max_tokens\":16}" \ - | sed -n '1,20p' + -d "{\"model\":\"openai/gpt-4.1-nano\",\"messages\":[{\"role\":\"user\",\"content\":\"Reply with exactly $QA_CACHE_REPLY\"}],\"max_tokens\":16}" +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +jq -e --arg reply "$QA_CACHE_REPLY" '.choices[0].message.content == $reply' "$BODY_FILE" >/dev/null +grep -Eiq '^X-Cache: HIT \(exact\)' "$HEADERS_FILE" ``` ### S73 Cache overview filtered by user path @@ -931,9 +1032,12 @@ Checks cache analytics after the exact-cache hit using the same tracked user pat ```bash sleep 6 -curl -sS "$AUTH_BASE_URL/admin/api/v1/cache/overview?days=1&user_path=$QA_CACHE_USER_PATH" \ +CACHE_OVERVIEW_JSON_FILE="$QA_RUN_DIR/s73.cache-overview.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/cache/overview?days=1&user_path=$QA_CACHE_USER_PATH" \ -H "$ADMIN_AUTH_HEADER" \ - | jq '.' + > "$CACHE_OVERVIEW_JSON_FILE" +jq '.' "$CACHE_OVERVIEW_JSON_FILE" +jq -e '.summary.total_hits >= 1 and .summary.exact_hits >= 1' "$CACHE_OVERVIEW_JSON_FILE" >/dev/null ``` ### S74 Cached usage log filtered by user path @@ -941,9 +1045,15 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/cache/overview?days=1&user_path=$QA_CACHE_ Reads cached-only usage entries for the same exact-hit request path. ```bash -curl -sS "$AUTH_BASE_URL/admin/api/v1/usage/log?days=1&user_path=$QA_CACHE_USER_PATH&cache_mode=cached&limit=5" \ +CACHED_USAGE_JSON_FILE="$QA_RUN_DIR/s74.cached-usage.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/usage/log?days=1&user_path=$QA_CACHE_USER_PATH&cache_mode=cached&limit=5" \ -H "$ADMIN_AUTH_HEADER" \ - | jq '{total,entries:(.entries|map({request_id,cache_type,model,provider,endpoint,user_path,total_tokens}))}' + > "$CACHED_USAGE_JSON_FILE" +jq '{total,entries:(.entries|map({request_id,cache_type,model,provider,endpoint,user_path,total_tokens}))}' "$CACHED_USAGE_JSON_FILE" +jq -e --arg request_id "$QA_CACHE_REQ2" ' + .total >= 1 + and any(.entries[]?; .request_id == $request_id and .cache_type == "exact") + ' "$CACHED_USAGE_JSON_FILE" >/dev/null ``` ### S75 Invalid managed API key user path (negative) @@ -951,11 +1061,16 @@ curl -sS "$AUTH_BASE_URL/admin/api/v1/usage/log?days=1&user_path=$QA_CACHE_USER_ Verifies user-path validation for managed API key creation. ```bash -curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s75.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s75.body.XXXXXX") +curl -sS -D "$HEADERS_FILE" -o "$BODY_FILE" -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ - -d '{"name":"qa-invalid-user-path","user_path":"/team/../alpha"}' \ - | sed -n '1,20p' + -d '{"name":"qa-invalid-user-path","user_path":"/team/../alpha"}' +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +grep -Eiq '^HTTP/.* 400 ' "$HEADERS_FILE" +jq -e '.error.type == "invalid_request_error" and (.error.message | test("invalid user_path"))' "$BODY_FILE" >/dev/null ``` ### S76 Invalid workflow scope user path (negative) @@ -963,11 +1078,16 @@ curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ Verifies user-path validation for workflow creation. ```bash -curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans" \ +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s76.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s76.body.XXXXXX") +curl -sS -D "$HEADERS_FILE" -o "$BODY_FILE" -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans" \ -H "$ADMIN_AUTH_HEADER" \ -H 'Content-Type: application/json' \ - -d '{"scope_provider":"openai","scope_model":"gpt-4.1-nano","scope_user_path":"/team/../alpha","name":"qa-invalid-workflow-path","plan_payload":{"schema_version":1,"features":{"cache":true,"audit":true,"usage":true,"guardrails":false},"guardrails":[]}}' \ - | sed -n '1,24p' + -d '{"scope_provider":"openai","scope_model":"gpt-4.1-nano","scope_user_path":"/team/../alpha","name":"qa-invalid-workflow-path","plan_payload":{"schema_version":1,"features":{"cache":true,"audit":true,"usage":true,"guardrails":false},"guardrails":[]}}' +sed -n '1,24p' "$HEADERS_FILE" +sed -n '1,24p' "$BODY_FILE" +grep -Eiq '^HTTP/.* 400 ' "$HEADERS_FILE" +jq -e '.error.type == "invalid_request_error" and (.error.message | test("invalid scope_user_path"))' "$BODY_FILE" >/dev/null ``` ## 13. Authenticated cleanup @@ -977,11 +1097,19 @@ curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans" \ Deactivates the managed key created for the auth-enabled release run. ```bash -AUTH_KEY_ID=$(curl -sS "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ +AUTH_KEYS_JSON_FILE="$QA_RUN_DIR/s77.auth-keys.json" +curl -fsS "$AUTH_BASE_URL/admin/api/v1/auth-keys" \ -H "$ADMIN_AUTH_HEADER" \ - | jq -r ".[] | select(.name==\"$QA_AUTH_KEY_NAME\") | .id") -curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys/$AUTH_KEY_ID/deactivate" \ + > "$AUTH_KEYS_JSON_FILE" +if ! AUTH_KEY_ID=$(jq -er --arg name "$QA_AUTH_KEY_NAME" '.[] | select(.name == $name) | .id' "$AUTH_KEYS_JSON_FILE"); then + echo "error: managed API key id not found for $QA_AUTH_KEY_NAME" >&2 + exit 1 +fi +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s77.headers.XXXXXX") +curl -sS -D "$HEADERS_FILE" -o /dev/null -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys/$AUTH_KEY_ID/deactivate" \ -H "$ADMIN_AUTH_HEADER" +sed -n '1,20p' "$HEADERS_FILE" +grep -Eiq '^HTTP/.* 204 ' "$HEADERS_FILE" ``` ### S78 Deactivated managed API key is rejected @@ -989,13 +1117,19 @@ curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/auth-keys/$AUTH_KEY_ID/deactiva Confirms that the same managed key can no longer authenticate requests. ```bash -API_KEY=$(cat "$QA_AUTH_KEY_VALUE_FILE") -curl -sS -i "$AUTH_BASE_URL/v1/chat/completions" \ +require_release_artifact "$QA_AUTH_KEY_VALUE_FILE" +API_KEY=$(<"$QA_AUTH_KEY_VALUE_FILE") +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s78.headers.XXXXXX") +BODY_FILE=$(mktemp "$QA_RUN_DIR/s78.body.XXXXXX") +curl -sS -D "$HEADERS_FILE" -o "$BODY_FILE" "$AUTH_BASE_URL/v1/chat/completions" \ -H "Authorization: Bearer $API_KEY" \ -H 'Content-Type: application/json' \ -H "X-Request-ID: $QA_DEACTIVATED_REQ" \ - -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_DEACTIVATED"}],"max_tokens":16}' \ - | sed -n '1,20p' + -d '{"model":"openai/gpt-4.1-nano","messages":[{"role":"user","content":"Reply with exactly QA_AUTH_DEACTIVATED"}],"max_tokens":16}' +sed -n '1,20p' "$HEADERS_FILE" +sed -n '1,20p' "$BODY_FILE" +grep -Eiq '^HTTP/.* 401 ' "$HEADERS_FILE" +jq -e '.error.type == "authentication_error"' "$BODY_FILE" >/dev/null ``` ### S79 Deactivate scoped workflow @@ -1003,8 +1137,12 @@ curl -sS -i "$AUTH_BASE_URL/v1/chat/completions" \ Deactivates the workflow created for the scoped managed-key release run. ```bash -WORKFLOW_ID=$(cat "$QA_WORKFLOW_ID_FILE") -curl -sS -i -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans/$WORKFLOW_ID/deactivate" \ +require_release_artifact "$QA_WORKFLOW_ID_FILE" +WORKFLOW_ID=$(<"$QA_WORKFLOW_ID_FILE") +HEADERS_FILE=$(mktemp "$QA_RUN_DIR/s79.headers.XXXXXX") +curl -sS -D "$HEADERS_FILE" -o /dev/null -X POST "$AUTH_BASE_URL/admin/api/v1/execution-plans/$WORKFLOW_ID/deactivate" \ -H "$ADMIN_AUTH_HEADER" +sed -n '1,20p' "$HEADERS_FILE" +grep -Eiq '^HTTP/.* 204 ' "$HEADERS_FILE" rm -f "$QA_AUTH_KEY_JSON" "$QA_AUTH_KEY_VALUE_FILE" "$QA_WORKFLOW_JSON" "$QA_WORKFLOW_ID_FILE" ``` diff --git a/tests/e2e/run-release-e2e.sh b/tests/e2e/run-release-e2e.sh index 1c75f25bb..c6b4a5891 100755 --- a/tests/e2e/run-release-e2e.sh +++ b/tests/e2e/run-release-e2e.sh @@ -364,6 +364,7 @@ for index in "${SELECTED_INDEXES[@]}"; do { printf '#!/usr/bin/env bash\n' printf 'set -euo pipefail\n' + printf 'shopt -s inherit_errexit 2>/dev/null || true\n' printf 'cd %q\n' "$REPO_ROOT" printf 'export QA_SUFFIX=%q\n' "$QA_SUFFIX" printf 'export QA_RUN_DIR=%q\n' "$OUTPUT_DIR" diff --git a/tests/integration/dbassert/auditlog.go b/tests/integration/dbassert/auditlog.go index 7c637ce00..47109d644 100644 --- a/tests/integration/dbassert/auditlog.go +++ b/tests/integration/dbassert/auditlog.go @@ -45,7 +45,7 @@ func QueryAuditLogsByRequestID(t *testing.T, pool *pgxpool.Pool, requestID strin defer cancel() query := ` - SELECT id, timestamp, duration_ns, model, provider, status_code, + SELECT id, timestamp, duration_ns, requested_model, provider, status_code, request_id, auth_key_id, client_ip, method, path, user_path, stream, error_type, data FROM audit_logs WHERE request_id = $1 @@ -141,7 +141,9 @@ func bsonToAuditLogEntry(t *testing.T, doc bson.M) AuditLogEntry { } else if v, ok := doc["duration_ns"].(int32); ok { entry.DurationNs = int64(v) } - if v, ok := doc["model"].(string); ok { + if v, ok := doc["requested_model"].(string); ok { + entry.Model = v + } else if v, ok := doc["model"].(string); ok { entry.Model = v } if v, ok := doc["provider"].(string); ok {