diff --git a/containers/api-proxy/guards/ai-credits-guard.js b/containers/api-proxy/guards/ai-credits-guard.js index d2aa7871e..8025a69b2 100644 --- a/containers/api-proxy/guards/ai-credits-guard.js +++ b/containers/api-proxy/guards/ai-credits-guard.js @@ -27,6 +27,45 @@ const BUILTIN_FALLBACK_PRICING = Object.freeze({ output: 15.00, }); +// Static ceiling for recognized dynamic selectors whose concrete runtime model +// is unknown at accounting time. These rates bound every priced Copilot model +// in the curated catalog; retain the ceiling if the catalog is unavailable. +const DYNAMIC_SELECTOR_PRICING_CEILING = Object.freeze({ + input: 10.00, + cachedInput: 1.00, + cacheWrite: 12.50, + output: 50.00, +}); + +function buildDynamicSelectorFallbackPricing() { + const pricing = { ...DYNAMIC_SELECTOR_PRICING_CEILING }; + for (const catalogPricing of Object.values(pricingByModel)) { + for (const field of Object.keys(pricing)) { + if (typeof catalogPricing[field] === 'number') { + pricing[field] = Math.max(pricing[field], catalogPricing[field]); + } + } + } + return Object.freeze(pricing); +} + +const DYNAMIC_SELECTOR_FALLBACK_PRICING = buildDynamicSelectorFallbackPricing(); + +// Kept separate from the generic unknown-model fallback because a dynamic +// selector is an explicit provider-supported request, not an unknown model. +const DYNAMIC_SELECTOR_FALLBACK_PRICING_SOURCE = 'dynamic_selector_fallback'; + +function getDynamicSelectorDescriptor(model, provider = undefined) { + if (typeof model !== 'string') return null; + if (provider !== PROVIDER_COPILOT) return null; + if (model.toLowerCase() !== 'auto') return null; + return { name: 'copilot:auto' }; +} + +function isRecognizedDynamicSelector(model, provider = undefined) { + return !!getDynamicSelectorDescriptor(model, provider); +} + function roundCredits(value) { return Math.round((value + Number.EPSILON) * 1_000_000) / 1_000_000; } @@ -102,7 +141,7 @@ function resolveModelPricing(model, state = aiCreditsState, provider = undefined .every(field => Object.hasOwn(runtime.pricing, field))) { return runtime; } - const fallback = resolveLowerPriorityPricing(model, state, options); + const fallback = resolveLowerPriorityPricing(model, state, { ...options, provider }); if (!runtime) return fallback; const mergedPricing = {}; for (const field of ['input', 'cachedInput', 'cacheWrite', 'output']) { @@ -147,6 +186,18 @@ function resolveLowerPriorityPricing(model, state, options = {}) { return { pricing: catalogModel.pricing, source: 'models.dev', tier: 'default' }; } + const dynamicSelector = getDynamicSelectorDescriptor(model, options.provider); + if (dynamicSelector) { + return { + pricing: DYNAMIC_SELECTOR_FALLBACK_PRICING, + source: DYNAMIC_SELECTOR_FALLBACK_PRICING_SOURCE, + tier: 'conservative', + accountingPolicy: 'dynamic_selector_fallback', + dynamicSelector: dynamicSelector.name, + usedFallbackPricing: true, + }; + } + // Speculative callers (e.g. filtering a fallback candidate pool) pass quiet:true // so that probing a model neither emits an operator-facing warning nor marks the // model as already-warned — which would suppress the warning if it is genuinely @@ -249,6 +300,7 @@ function calculateAiCredits(normalizedUsage, model, state = aiCreditsState, prov const pricingResolution = resolveModelPricing(model, state, provider, totalInputForTier); if (!pricingResolution) return null; const { pricing } = pricingResolution; + const dynamicSelector = getDynamicSelectorDescriptor(model, provider); // input_tokens semantics differ by provider: // - Anthropic and Copilot's precise copilot_usage report input_tokens as the @@ -283,6 +335,13 @@ function calculateAiCredits(normalizedUsage, model, state = aiCreditsState, prov pricingObservedAt: pricingResolution.observedAt, pricingApiVersion: pricingResolution.apiVersion, pricingDiscountPercent: pricingResolution.discountPercent, + accountingPolicy: pricingResolution.accountingPolicy || + (dynamicSelector ? 'dynamic_selector_runtime' : model === 'unknown' ? 'unknown_model_fallback' : 'concrete_model'), + usedFallbackPricing: pricingResolution.usedFallbackPricing === true || + pricingResolution.source === 'configured_default' || + pricingResolution.source === 'builtin_fallback' || + pricingResolution.source === 'dynamic_selector_fallback', + dynamicSelector: pricingResolution.dynamicSelector || dynamicSelector?.name || null, }; } @@ -301,6 +360,9 @@ function applyAiCreditsUsage(normalizedUsage, model, provider = undefined) { totalCredits: 0, pricingSource: calc.pricingSource, pricingTier: calc.pricingTier, + accountingPolicy: calc.accountingPolicy, + fallbackPricingUsed: calc.usedFallbackPricing, + dynamicSelector: calc.dynamicSelector, }; } @@ -312,6 +374,9 @@ function applyAiCreditsUsage(normalizedUsage, model, provider = undefined) { modelBucket.totalCredits += calc.totalCredits; modelBucket.pricingSource = calc.pricingSource; modelBucket.pricingTier = calc.pricingTier; + modelBucket.accountingPolicy = calc.accountingPolicy; + modelBucket.fallbackPricingUsed = calc.usedFallbackPricing; + modelBucket.dynamicSelector = calc.dynamicSelector; aiCreditsState.totalAiCredits += calc.totalCredits; process.env.AWF_AI_CREDITS_USED = String(roundCredits(aiCreditsState.totalAiCredits)); @@ -325,6 +390,9 @@ function applyAiCreditsUsage(normalizedUsage, model, provider = undefined) { totalAiCredits: roundCredits(aiCreditsState.totalAiCredits), pricingSource: calc.pricingSource, pricingTier: calc.pricingTier, + accountingPolicy: calc.accountingPolicy, + fallbackPricingUsed: calc.usedFallbackPricing, + dynamicSelector: calc.dynamicSelector, ...(calc.pricingObservedAt ? { pricingObservedAt: calc.pricingObservedAt } : {}), ...(calc.pricingApiVersion ? { pricingApiVersion: calc.pricingApiVersion } : {}), ...(calc.pricingDiscountPercent !== undefined @@ -344,6 +412,9 @@ function getAiCreditsReflectState() { total: roundCredits(usage.totalCredits), pricing_source: usage.pricingSource, pricing_tier: usage.pricingTier, + accounting_policy: usage.accountingPolicy || null, + fallback_pricing_used: usage.fallbackPricingUsed === true, + dynamic_selector: usage.dynamicSelector || null, }; } return { @@ -404,6 +475,7 @@ module.exports = { getAiCreditsBlockState, buildAiCreditsLimitError, checkUnknownModelRejection, + isRecognizedDynamicSelector, isModelPriceable, canonicalizeModel, resetAiCreditsGuardForTests, diff --git a/containers/api-proxy/guards/ai-credits-guard.test.js b/containers/api-proxy/guards/ai-credits-guard.test.js index 56867dd85..b06e2d0be 100644 --- a/containers/api-proxy/guards/ai-credits-guard.test.js +++ b/containers/api-proxy/guards/ai-credits-guard.test.js @@ -5,6 +5,7 @@ const { getAiCreditsBlockState, buildAiCreditsLimitError, checkUnknownModelRejection, + isRecognizedDynamicSelector, isModelPriceable, canonicalizeModel, resetAiCreditsGuardForTests, @@ -77,6 +78,9 @@ describe('ai-credits-guard', () => { total: 0.12275, pricing_source: 'curated', pricing_tier: 'default', + accounting_policy: 'concrete_model', + fallback_pricing_used: false, + dynamic_selector: null, }, }, }); @@ -665,15 +669,53 @@ describe('ai-credits-guard', () => { expect(sonnet5).toBeNull(); }); - it('rejects the auto selector when runtime pricing cannot be proven', () => { + it('allows the Copilot auto selector and rejects unknown-provider auto selectors', () => { process.env.AWF_MAX_AI_CREDITS = '10'; resetAiCreditsGuardForTests(); - expect(checkUnknownModelRejection('auto', PROVIDER_COPILOT)).not.toBeNull(); + expect(checkUnknownModelRejection('auto', PROVIDER_COPILOT)).toBeNull(); expect(checkUnknownModelRejection('auto', PROVIDER_OPENAI)).not.toBeNull(); }); }); + it('uses conservative dynamic-selector fallback pricing for Copilot auto', () => { + process.env.AWF_MAX_AI_CREDITS = '10'; + resetAiCreditsGuardForTests(); + + const usage = applyAiCreditsUsage({ + input_tokens: 1000, + output_tokens: 500, + }, 'auto', PROVIDER_COPILOT); + + expect(usage).toMatchObject({ + aiCreditsThisResponse: 3.5, + pricingSource: 'dynamic_selector_fallback', + pricingTier: 'conservative', + accountingPolicy: 'dynamic_selector_fallback', + fallbackPricingUsed: true, + dynamicSelector: 'copilot:auto', + }); + expect(isRecognizedDynamicSelector('auto', PROVIDER_COPILOT)).toBe(true); + expect(isRecognizedDynamicSelector('auto', PROVIDER_OPENAI)).toBe(false); + expect(getAiCreditsReflectState()).toEqual({ + total: 3.5, + by_model: { + auto: { + input_credits: 1, + cached_input_credits: 0, + cache_write_credits: 0, + output_credits: 2.5, + total: 3.5, + pricing_source: 'dynamic_selector_fallback', + pricing_tier: 'conservative', + accounting_policy: 'dynamic_selector_fallback', + fallback_pricing_used: true, + dynamic_selector: 'copilot:auto', + }, + }, + }); + }); + describe('isModelPriceable (side-effect-free)', () => { it('reports priced and unpriced models correctly', () => { expect(isModelPriceable('gpt-4-turbo', PROVIDER_OPENAI)).toBe(true); diff --git a/containers/api-proxy/guards/common-guard-checks.js b/containers/api-proxy/guards/common-guard-checks.js index 331357eaf..50b9df32e 100644 --- a/containers/api-proxy/guards/common-guard-checks.js +++ b/containers/api-proxy/guards/common-guard-checks.js @@ -131,7 +131,7 @@ function buildCommonGuardChecks(deps, model, provider = null) { // Model-specific guards — only active when a model was identified in the request. ...(model ? [ { - block: getModelMultiplierCapBlockState(model), + block: getModelMultiplierCapBlockState(model, provider), isBlocked: block => !!block, statusCode: 400, eventName: 'model_multiplier_cap_exceeded', diff --git a/containers/api-proxy/guards/max-model-multiplier-guard.js b/containers/api-proxy/guards/max-model-multiplier-guard.js index 49c9364d2..696d70082 100644 --- a/containers/api-proxy/guards/max-model-multiplier-guard.js +++ b/containers/api-proxy/guards/max-model-multiplier-guard.js @@ -2,6 +2,7 @@ const { sanitizeForLog } = require('../logging'); const { parseModelMultipliers, parsePositiveNumber } = require('./guard-utils'); +const { isRecognizedDynamicSelector } = require('./ai-credits-guard'); const maxModelMultiplierConfigCache = { rawCap: undefined, @@ -65,11 +66,11 @@ function resolveMultiplierForModel(model, config) { * @param {string|null} model - The model name from the request body (may be null) * @returns {{ model: string, multiplier: number, maxModelMultiplier: number } | null} */ -function getModelMultiplierCapBlockState(model) { +function getModelMultiplierCapBlockState(model, provider = undefined) { const config = getMaxModelMultiplierConfig(); if (!config.cap || !model) return null; - if (model.toLowerCase() === 'auto' && !Object.hasOwn(config.multipliers, model)) { + if (isRecognizedDynamicSelector(model, provider) && !Object.hasOwn(config.multipliers, model)) { return { model: sanitizeForLog(model), multiplier: null, @@ -99,7 +100,7 @@ function buildModelMultiplierCapError(state) { return { error: { type: 'model_multiplier_cap_unverifiable', - message: 'Model "auto" selects a concrete model at runtime, so its multiplier cannot be proven to be within the configured cap. Configure an explicit multiplier for "auto" to opt in.', + message: `Model "${state.model}" selects a concrete model at runtime, so its multiplier cannot be proven to be within the configured cap. Configure an explicit multiplier for "${state.model}" to opt in.`, model: state.model, model_multiplier: null, max_model_multiplier: state.maxModelMultiplier, diff --git a/containers/api-proxy/guards/max-model-multiplier-guard.test.js b/containers/api-proxy/guards/max-model-multiplier-guard.test.js index b3a375a1f..a60acb936 100644 --- a/containers/api-proxy/guards/max-model-multiplier-guard.test.js +++ b/containers/api-proxy/guards/max-model-multiplier-guard.test.js @@ -5,6 +5,7 @@ const { buildModelMultiplierCapError, resetMaxModelMultiplierGuardForTests, } = require('./max-model-multiplier-guard'); +const { PROVIDER_COPILOT, PROVIDER_OPENAI } = require('../provider-names'); describe('max-model-multiplier-guard', () => { beforeEach(() => { @@ -85,10 +86,10 @@ describe('max-model-multiplier-guard', () => { expect(getModelMultiplierCapBlockState('unknown-model')).toBeNull(); }); - it('fails closed for auto unless an explicit multiplier is configured', () => { + it('fails closed for Copilot auto unless an explicit multiplier is configured', () => { process.env.AWF_MAX_MODEL_MULTIPLIER = '5'; - const state = getModelMultiplierCapBlockState('auto'); + const state = getModelMultiplierCapBlockState('auto', PROVIDER_COPILOT); expect(state).toMatchObject({ model: 'auto', multiplier: null, @@ -102,7 +103,13 @@ describe('max-model-multiplier-guard', () => { process.env.AWF_MAX_MODEL_MULTIPLIER = '5'; process.env.AWF_EFFECTIVE_TOKEN_MODEL_MULTIPLIERS = JSON.stringify({ auto: 5 }); - expect(getModelMultiplierCapBlockState('auto')).toBeNull(); + expect(getModelMultiplierCapBlockState('auto', PROVIDER_COPILOT)).toBeNull(); + }); + + it('treats non-Copilot auto as a normal model name', () => { + process.env.AWF_MAX_MODEL_MULTIPLIER = '5'; + + expect(getModelMultiplierCapBlockState('auto', PROVIDER_OPENAI)).toBeNull(); }); it('blocks when configured default multiplier for unknown model exceeds cap', () => { diff --git a/containers/api-proxy/server.token-guards.test.js b/containers/api-proxy/server.token-guards.test.js index a6890aab2..7d6252296 100644 --- a/containers/api-proxy/server.token-guards.test.js +++ b/containers/api-proxy/server.token-guards.test.js @@ -341,7 +341,7 @@ describe('proxyRequest max-ai-credits guard', () => { expect(payload.error.total_ai_credits).toBeGreaterThanOrEqual(0.1); }); - it('rejects Copilot auto when its concrete runtime price cannot be proven', async () => { + it('allows Copilot auto and defers accounting to runtime usage tracking', async () => { const upstreamRequest = makeProxyReq(); const httpsRequestSpy = jest.spyOn(https, 'request').mockImplementation(() => upstreamRequest); @@ -351,11 +351,8 @@ describe('proxyRequest max-ai-credits guard', () => { req.emit('end'); await flushPromises(); - expect(httpsRequestSpy).not.toHaveBeenCalled(); - expect(res.writeHead).toHaveBeenCalledWith(400, expect.objectContaining({ - 'Content-Type': 'application/json', - })); - expect(JSON.parse(res.end.mock.calls[0][0]).type).toBe('unknown_model_ai_credits'); + expect(httpsRequestSpy).toHaveBeenCalledTimes(1); + expect(res.writeHead).not.toHaveBeenCalledWith(400, expect.anything()); }); }); diff --git a/containers/api-proxy/server.websocket.test.js b/containers/api-proxy/server.websocket.test.js index 94f78bce2..38ff9b14f 100644 --- a/containers/api-proxy/server.websocket.test.js +++ b/containers/api-proxy/server.websocket.test.js @@ -362,11 +362,9 @@ describe('proxyWebSocket', () => { // ── Security guard tests ────────────────────────────────────────────────────── // -// These tests verify that common (non-model-specific) security guards are +// These tests verify that common and query-model-specific security guards are // enforced on the WebSocket upgrade path using the shared buildCommonGuardChecks -// factory. Model-specific guards (model_multiplier_cap, retired_model, -// unknown_model_ai_credits) are intentionally skipped because WebSocket -// upgrades pass model=null (no JSON body to extract a model from). +// factory. // Guards are triggered by directly calling their apply functions (same // technique used in guards/*.test.js unit tests). @@ -377,6 +375,7 @@ describe('proxyWebSocket security guards', () => { let applyEffectiveTokenUsage, resetEffectiveTokenGuardForTests; let applyPermissionDenied, resetPermissionDeniedGuardForTests; let applyAiCreditsUsage, resetAiCreditsGuardForTests; + let resetMaxModelMultiplierGuardForTests; beforeAll(() => { jest.resetModules(); @@ -386,6 +385,7 @@ describe('proxyWebSocket security guards', () => { ({ applyEffectiveTokenUsage, resetEffectiveTokenGuardForTests } = require('./guards/effective-token-guard')); ({ applyPermissionDenied, resetPermissionDeniedGuardForTests } = require('./guards/max-permission-denied-guard')); ({ applyAiCreditsUsage, resetAiCreditsGuardForTests } = require('./guards/ai-credits-guard')); + ({ resetMaxModelMultiplierGuardForTests } = require('./guards/max-model-multiplier-guard')); }); afterAll(() => { @@ -398,11 +398,14 @@ describe('proxyWebSocket security guards', () => { delete process.env.AWF_MAX_EFFECTIVE_TOKENS; delete process.env.AWF_MAX_PERMISSION_DENIED; delete process.env.AWF_MAX_AI_CREDITS; + delete process.env.AWF_MAX_MODEL_MULTIPLIER; + delete process.env.AWF_EFFECTIVE_TOKEN_MODEL_MULTIPLIERS; resetMaxRunsGuardForTests(); resetMaxCacheMissesGuardForTests(); resetEffectiveTokenGuardForTests(); resetPermissionDeniedGuardForTests(); resetAiCreditsGuardForTests(); + resetMaxModelMultiplierGuardForTests(); jest.restoreAllMocks(); }); @@ -466,6 +469,28 @@ describe('proxyWebSocket security guards', () => { expect(socket.destroy).toHaveBeenCalled(); }); + it('rejects an unknown WebSocket query model when AI-credit accounting is enabled', () => { + process.env.AWF_MAX_AI_CREDITS = '10'; + + const socket = makeMockSocket(); + wsProxy(makeUpgradeReq({ url: '/v1/responses?model=bogus' }), socket, Buffer.alloc(0), 'api.openai.com', {}, 'copilot'); + + expect(socket.write).toHaveBeenCalledWith(expect.stringContaining('HTTP/1.1 400 Bad Request')); + expect(socket.write).toHaveBeenCalledWith(expect.stringContaining('"unknown_model_ai_credits"')); + expect(socket.destroy).toHaveBeenCalled(); + }); + + it('rejects an unverifiable Copilot auto WebSocket model multiplier', () => { + process.env.AWF_MAX_MODEL_MULTIPLIER = '1'; + + const socket = makeMockSocket(); + wsProxy(makeUpgradeReq({ url: '/v1/responses?model=auto' }), socket, Buffer.alloc(0), 'api.openai.com', {}, 'copilot'); + + expect(socket.write).toHaveBeenCalledWith(expect.stringContaining('HTTP/1.1 400 Bad Request')); + expect(socket.write).toHaveBeenCalledWith(expect.stringContaining('"model_multiplier_cap_unverifiable"')); + expect(socket.destroy).toHaveBeenCalled(); + }); + it('allows the upgrade when no guards are triggered', () => { // No guard env vars set and no usage applied — all guards pass. // Without HTTPS_PROXY the upgrade will fail with 502, but the key point is diff --git a/containers/api-proxy/token-budget-log.js b/containers/api-proxy/token-budget-log.js index 640c2f9bd..9596cc0ba 100644 --- a/containers/api-proxy/token-budget-log.js +++ b/containers/api-proxy/token-budget-log.js @@ -30,6 +30,9 @@ function computeTokenBudgetUsage({ logRequest, requestId, provider }, normalized ai_credits_total: aiCreditsUsage.totalAiCredits, pricing_source: aiCreditsUsage.pricingSource, pricing_tier: aiCreditsUsage.pricingTier, + accounting_policy: aiCreditsUsage.accountingPolicy, + fallback_pricing_used: aiCreditsUsage.fallbackPricingUsed, + dynamic_selector: aiCreditsUsage.dynamicSelector, }); } const budgetFields = {}; @@ -43,6 +46,9 @@ function computeTokenBudgetUsage({ logRequest, requestId, provider }, normalized budgetFields.ai_credits_total = aiCreditsUsage.totalAiCredits; budgetFields.ai_credits_pricing_source = aiCreditsUsage.pricingSource; budgetFields.ai_credits_pricing_tier = aiCreditsUsage.pricingTier; + budgetFields.ai_credits_accounting_policy = aiCreditsUsage.accountingPolicy; + budgetFields.ai_credits_fallback_pricing_used = aiCreditsUsage.fallbackPricingUsed; + budgetFields.ai_credits_dynamic_selector = aiCreditsUsage.dynamicSelector; if (aiCreditsUsage.pricingObservedAt) { budgetFields.ai_credits_pricing_observed_at = aiCreditsUsage.pricingObservedAt; } diff --git a/containers/api-proxy/token-budget-log.test.js b/containers/api-proxy/token-budget-log.test.js index 2b5bbcaa6..820c7b404 100644 --- a/containers/api-proxy/token-budget-log.test.js +++ b/containers/api-proxy/token-budget-log.test.js @@ -112,4 +112,30 @@ describe('computeTokenBudgetUsage', () => { }); expect(logRequest).toHaveBeenCalledTimes(1); }); + + it('reports dynamic selector accounting policy in token diagnostics', () => { + process.env.AWF_MAX_AI_CREDITS = '100'; + const result = computeTokenBudgetUsage( + { logRequest, requestId: 'req-dynamic', provider: 'copilot' }, + { input_tokens: 1000, output_tokens: 500 }, + 'auto', + ); + + expect(result).toMatchObject({ + ai_credits_pricing_source: 'dynamic_selector_fallback', + ai_credits_pricing_tier: 'conservative', + ai_credits_accounting_policy: 'dynamic_selector_fallback', + ai_credits_fallback_pricing_used: true, + ai_credits_dynamic_selector: 'copilot:auto', + }); + expect(logRequest).toHaveBeenCalledWith('info', 'token_budget_usage', expect.objectContaining({ + request_id: 'req-dynamic', + provider: 'copilot', + model: 'auto', + pricing_source: 'dynamic_selector_fallback', + accounting_policy: 'dynamic_selector_fallback', + fallback_pricing_used: true, + dynamic_selector: 'copilot:auto', + })); + }); }); diff --git a/containers/api-proxy/token-tracker-shared.js b/containers/api-proxy/token-tracker-shared.js index 4d56f6c37..bfb1b9250 100644 --- a/containers/api-proxy/token-tracker-shared.js +++ b/containers/api-proxy/token-tracker-shared.js @@ -45,6 +45,9 @@ function mergeBudgetFields(record, budgetResult) { for (const field of [ 'ai_credits_pricing_source', 'ai_credits_pricing_tier', + 'ai_credits_accounting_policy', + 'ai_credits_fallback_pricing_used', + 'ai_credits_dynamic_selector', 'ai_credits_pricing_observed_at', 'ai_credits_pricing_api_version', 'ai_credits_pricing_discount_percent', diff --git a/containers/api-proxy/token-tracker-ws.js b/containers/api-proxy/token-tracker-ws.js index 3b1648959..b9529df8e 100644 --- a/containers/api-proxy/token-tracker-ws.js +++ b/containers/api-proxy/token-tracker-ws.js @@ -131,10 +131,11 @@ function parseWebSocketFrames(buf, fragments) { * @param {string} opts.path - Request path * @param {number} opts.startTime - Request start time (Date.now()) * @param {object} opts.metrics - Metrics module reference + * @param {string|null} [opts.requestModel] - Model extracted from the request context, used when response metadata omits model * @param {(normalizedUsage: object, model: string|null) => Record|void} [opts.onUsage] - Optional callback invoked after normalized usage is extracted */ function trackWebSocketTokenUsage(upstreamSocket, opts) { - const { requestId, provider, path: reqPath, startTime, metrics: metricsRef, onUsage } = opts; + const { requestId, provider, path: reqPath, startTime, metrics: metricsRef, onUsage, requestModel } = opts; auditTrack('WS_TRACK_START', { rid: requestId, provider, path: reqPath }); logRequest('debug', 'ws_token_track_start', { @@ -238,10 +239,11 @@ function trackWebSocketTokenUsage(upstreamSocket, opts) { if (observedCacheReadTokens > 0 && normalized.cache_read_tokens === 0) { warnCacheReadRollupMismatch({ logRequest, diag, requestId, provider, model: streamingModel, observedCacheReadTokens, normalizedCacheReadTokens: normalized.cache_read_tokens, streaming: true, transport: 'websocket' }); } + const resolvedModel = streamingModel || requestModel || 'unknown'; let budgetResult; if (typeof onUsage === 'function') { try { - budgetResult = onUsage(normalized, streamingModel || 'unknown'); + budgetResult = onUsage(normalized, resolvedModel); } catch { // best-effort callback } @@ -252,7 +254,7 @@ function trackWebSocketTokenUsage(upstreamSocket, opts) { const record = buildTokenUsageRecord(normalized, { requestId, provider, - model: streamingModel, + model: streamingModel || requestModel, reqPath, status: 101, streaming: true, @@ -268,7 +270,7 @@ function trackWebSocketTokenUsage(upstreamSocket, opts) { logRequest('info', 'token_usage', { request_id: requestId, provider, - model: streamingModel || 'unknown', + model: resolvedModel, input_tokens: normalized.input_tokens, output_tokens: normalized.output_tokens, cache_read_tokens: normalized.cache_read_tokens, diff --git a/containers/api-proxy/token-tracker.websocket.test.js b/containers/api-proxy/token-tracker.websocket.test.js index 3e67f6642..dd94e57d3 100644 --- a/containers/api-proxy/token-tracker.websocket.test.js +++ b/containers/api-proxy/token-tracker.websocket.test.js @@ -400,4 +400,41 @@ describe('trackWebSocketTokenUsage', () => { } }, 10); }); + + test('uses request-model fallback when WebSocket response omits model metadata', (done) => { + const socket = new EventEmitter(); + const metricsRef = { increment: jest.fn() }; + const onUsage = jest.fn(() => undefined); + + trackWebSocketTokenUsage(socket, { + requestId: 'ws-request-model-fallback', + provider: 'copilot', + path: '/v1/chat/completions?model=auto', + startTime: Date.now(), + metrics: metricsRef, + requestModel: 'auto', + onUsage, + }); + + socket.emit('data', Buffer.from('HTTP/1.1 101 Switching Protocols\r\n\r\n')); + socket.emit('data', buildTextFrame(JSON.stringify({ + type: 'response.completed', + response: { + usage: { input_tokens: 120, output_tokens: 30, total_tokens: 150 }, + }, + }))); + socket.emit('close'); + + setTimeout(() => { + try { + expect(onUsage).toHaveBeenCalledWith( + expect.objectContaining({ input_tokens: 120, output_tokens: 30 }), + 'auto', + ); + done(); + } catch (err) { + done(err); + } + }, 10); + }); }); diff --git a/containers/api-proxy/websocket-guards.js b/containers/api-proxy/websocket-guards.js index 37607d951..f298ef478 100644 --- a/containers/api-proxy/websocket-guards.js +++ b/containers/api-proxy/websocket-guards.js @@ -14,10 +14,8 @@ const HTTP_STATUS_LINES = { * Writes a raw HTTP error response to the socket and destroys it when any * guard triggers, then returns true. Returns false when all guards pass. */ -function enforceWebSocketGuards({ socket, logRequest, requestId, provider }, guardDeps) { - // WebSocket upgrade requests have no JSON body, so model-specific guards - // receive null and are skipped (their getters return null for null models). - const guardChecks = buildCommonGuardChecks(guardDeps, null, provider); +function enforceWebSocketGuards({ socket, logRequest, requestId, provider, requestModel = null }, guardDeps) { + const guardChecks = buildCommonGuardChecks(guardDeps, requestModel, provider); for (const guard of guardChecks) { if (!guard.isBlocked(guard.block)) continue; diff --git a/containers/api-proxy/websocket-proxy.js b/containers/api-proxy/websocket-proxy.js index fe207f17a..27a78c9e7 100644 --- a/containers/api-proxy/websocket-proxy.js +++ b/containers/api-proxy/websocket-proxy.js @@ -1,7 +1,7 @@ 'use strict'; const { enforceWebSocketGuards, enforceWebSocketRateLimit } = require('./websocket-guards'); -const { createWebSocketTunnel } = require('./websocket-tunnel'); +const { createWebSocketTunnel, extractRequestModelFromUrl } = require('./websocket-tunnel'); function createProxyWebSocket({ limiter, @@ -97,9 +97,10 @@ function createProxyWebSocket({ return; } + const requestModel = extractRequestModelFromUrl(req.url); const upstreamPath = buildUpstreamPath(req.url, targetHost, basePath); - if (enforceWebSocketGuards({ socket, logRequest, requestId, provider }, guardDeps)) { + if (enforceWebSocketGuards({ socket, logRequest, requestId, provider, requestModel }, guardDeps)) { return; } @@ -122,6 +123,7 @@ function createProxyWebSocket({ requestId, startTime, upstreamPath, + requestModel, onSocketsReady: lifecycleHooks.onSocketsReady, }); }; diff --git a/containers/api-proxy/websocket-tunnel.js b/containers/api-proxy/websocket-tunnel.js index b0b53abf9..b8012f1fb 100644 --- a/containers/api-proxy/websocket-tunnel.js +++ b/containers/api-proxy/websocket-tunnel.js @@ -6,6 +6,17 @@ const { URL } = require('url'); const { computeTokenBudgetUsage } = require('./token-budget-log'); const { applyCopilotHostHeaders, mergeInjectedHeaders } = require('./request-headers'); +function extractRequestModelFromUrl(url) { + if (typeof url !== 'string' || !url.startsWith('/')) return null; + try { + const parsed = new URL(url, 'http://awf.local'); + const model = parsed.searchParams.get('model'); + return model && model.length > 0 ? model : null; + } catch { + return null; + } +} + function createProxyErrorResponder({ metrics, logRequest, @@ -69,6 +80,7 @@ function createWebSocketTunnel({ requestId, startTime, upstreamPath, + requestModel, onSocketsReady, }) { const { finalize, abort } = createProxyErrorResponder({ @@ -151,6 +163,7 @@ function createWebSocketTunnel({ path: sanitizeForLog(req.url), startTime, metrics, + requestModel: requestModel || extractRequestModelFromUrl(req.url), onUsage: (normalizedUsage, model) => computeTokenBudgetUsage({ logRequest, requestId, provider }, normalizedUsage, model), }); @@ -174,4 +187,5 @@ function createWebSocketTunnel({ module.exports = { createWebSocketTunnel, + extractRequestModelFromUrl, }; diff --git a/containers/api-proxy/websocket-tunnel.test.js b/containers/api-proxy/websocket-tunnel.test.js index 5067700c0..24d4c410c 100644 --- a/containers/api-proxy/websocket-tunnel.test.js +++ b/containers/api-proxy/websocket-tunnel.test.js @@ -1,5 +1,5 @@ const { EventEmitter } = require('events'); -const { createWebSocketTunnel } = require('./websocket-tunnel'); +const { createWebSocketTunnel, extractRequestModelFromUrl } = require('./websocket-tunnel'); function makeSocket() { const socket = new EventEmitter(); @@ -11,6 +11,11 @@ function makeSocket() { } describe('websocket-tunnel', () => { + it('extracts request model from websocket URL query', () => { + expect(extractRequestModelFromUrl('/v1/chat/completions?model=auto')).toBe('auto'); + expect(extractRequestModelFromUrl('/v1/chat/completions?foo=bar')).toBeNull(); + }); + it('returns 502 when HTTPS_PROXY is not configured', () => { const metrics = { gaugeDec: jest.fn(),