diff --git a/containers/api-proxy/server-factory.js b/containers/api-proxy/server-factory.js index c7f0e652c..15f3f6606 100644 --- a/containers/api-proxy/server-factory.js +++ b/containers/api-proxy/server-factory.js @@ -39,7 +39,48 @@ function createProxyHandler(adapter, checkRateLimit, proxyRequest) { } function createWebSocketUpgradeHandler(adapter, proxyWebSocket) { - return (req, socket, head) => { + const activeUpgradeSockets = new Set(); + + function trackUpgradeSocket(socket) { + if (!socket || typeof socket.destroy !== 'function') return; + activeUpgradeSockets.add(socket); + const cleanup = () => activeUpgradeSockets.delete(socket); + if (typeof socket.once === 'function') { + socket.once('close', cleanup); + socket.once('end', cleanup); + } + } + + async function shutdownConnections() { + const sockets = Array.from(activeUpgradeSockets); + await Promise.all(sockets.map((socket) => new Promise((resolve) => { + if (!activeUpgradeSockets.has(socket)) { + resolve(); + return; + } + let settled = false; + const finish = () => { + if (settled) return; + settled = true; + activeUpgradeSockets.delete(socket); + resolve(); + }; + if (typeof socket.once === 'function') { + socket.once('close', finish); + } + try { + if (socket.destroyed) { + finish(); + return; + } + socket.destroy(); + } catch { + finish(); + } + }))); + } + + const handler = (req, socket, head) => { if (!adapter.isEnabled()) { socket.write('HTTP/1.1 503 Service Unavailable\r\nConnection: close\r\n\r\n'); socket.destroy(); @@ -55,9 +96,17 @@ function createWebSocketUpgradeHandler(adapter, proxyWebSocket) { adapter.getTargetHost(req), adapter.getAuthHeaders(req), adapter.name, - adapter.getBasePath(req) + adapter.getBasePath(req), + { + onSocketsReady: (clientSocket, upstreamSocket) => { + trackUpgradeSocket(clientSocket); + trackUpgradeSocket(upstreamSocket); + }, + } ); }; + handler.shutdownConnections = shutdownConnections; + return handler; } function createProviderServer(adapter, deps) { @@ -71,6 +120,7 @@ function createProviderServer(adapter, deps) { const handleHealthCheck = createHealthCheckHandler(adapter); const handleProxy = createProxyHandler(adapter, checkRateLimit, proxyRequest); + const handleUpgrade = createWebSocketUpgradeHandler(adapter, proxyWebSocket); const server = http.createServer((req, res) => { if (adapter.isManagementPort && handleManagementEndpoint(req, res)) return; @@ -98,7 +148,8 @@ function createProviderServer(adapter, deps) { handleProxy(req, res); }); - server.on('upgrade', createWebSocketUpgradeHandler(adapter, proxyWebSocket)); + server.on('upgrade', handleUpgrade); + server.shutdownConnections = () => handleUpgrade.shutdownConnections(); return server; } diff --git a/containers/api-proxy/server-factory.test.js b/containers/api-proxy/server-factory.test.js new file mode 100644 index 000000000..5e20e589c --- /dev/null +++ b/containers/api-proxy/server-factory.test.js @@ -0,0 +1,48 @@ +'use strict'; + +const { EventEmitter } = require('events'); +const { createProviderServer } = require('./server-factory'); + +function makeTrackedSocket() { + const socket = new EventEmitter(); + socket.destroyed = false; + socket.destroy = jest.fn(() => { + if (socket.destroyed) return; + socket.destroyed = true; + socket.emit('close'); + }); + socket.write = jest.fn(); + return socket; +} + +describe('createProviderServer', () => { + test('shutdownConnections closes tracked upgraded sockets', async () => { + const clientSocket = makeTrackedSocket(); + const upstreamSocket = makeTrackedSocket(); + const proxyWebSocket = jest.fn((_req, socket, _head, _targetHost, _headers, _provider, _basePath, lifecycleHooks) => { + lifecycleHooks.onSocketsReady(socket, upstreamSocket); + }); + + const server = createProviderServer({ + name: 'anthropic', + isEnabled: () => true, + getTargetHost: () => 'api.anthropic.com', + getAuthHeaders: () => ({}), + getBasePath: () => '', + }, { + handleManagementEndpoint: () => false, + reflectEndpoints: () => [], + checkRateLimit: () => false, + proxyRequest: jest.fn(), + proxyWebSocket, + }); + + server.emit('upgrade', { url: '/v1/messages', headers: {} }, clientSocket, Buffer.alloc(0)); + + await server.shutdownConnections(); + + expect(proxyWebSocket).toHaveBeenCalledTimes(1); + expect(clientSocket.destroy).toHaveBeenCalledTimes(1); + expect(upstreamSocket.destroy).toHaveBeenCalledTimes(1); + }); +}); diff --git a/containers/api-proxy/startup.js b/containers/api-proxy/startup.js index 20fe3907c..5d7774f14 100644 --- a/containers/api-proxy/startup.js +++ b/containers/api-proxy/startup.js @@ -58,6 +58,7 @@ function bootPrimary({ } const adaptersToStart = registeredAdapters.filter(a => a.alwaysBind || a.isEnabled()); + const startedServers = []; const expectedListeners = adaptersToStart.filter(a => a.participatesInValidation).length; let readyListeners = 0; @@ -87,6 +88,7 @@ function bootPrimary({ for (const adapter of adaptersToStart) { const server = createProviderServer(adapter); + startedServers.push(server); server.listen(adapter.port, '0.0.0.0', () => { logRequest('info', 'server_start', { message: `${adapter.name} proxy listening on port ${adapter.port}`, @@ -98,8 +100,32 @@ function bootPrimary({ }); } + let shuttingDown = false; async function shutdownGracefully(signal) { + if (shuttingDown) return; + shuttingDown = true; logRequest('info', 'shutdown', { message: `Received ${signal}, shutting down gracefully` }); + const forceExitMs = Number.parseInt(process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS || '8000', 10); + const forceExitTimer = setTimeout(() => { + logRequest('warn', 'shutdown_force_exit', { + message: `Forced process exit after ${forceExitMs}ms shutdown timeout`, + }); + process.exit(0); + }, Number.isFinite(forceExitMs) && forceExitMs > 0 ? forceExitMs : 8000); + forceExitTimer.unref(); + await Promise.all(startedServers.map((server) => { + if (typeof server.shutdownConnections === 'function') { + return server.shutdownConnections(); + } + return Promise.resolve(); + })); + await Promise.all(startedServers.map((server) => new Promise((resolve) => { + try { + server.close(() => resolve()); + } catch { + resolve(); + } + }))); for (const adapter of registeredAdapters) { if (typeof adapter.getOidcProvider === 'function') { adapter.getOidcProvider()?.shutdown(); @@ -110,6 +136,7 @@ function bootPrimary({ } await closeLogStream(); await otelShutdown(); + clearTimeout(forceExitTimer); process.exit(0); } diff --git a/containers/api-proxy/startup.test.js b/containers/api-proxy/startup.test.js new file mode 100644 index 000000000..962229154 --- /dev/null +++ b/containers/api-proxy/startup.test.js @@ -0,0 +1,101 @@ +'use strict'; + +const { bootPrimary } = require('./startup'); + +describe('bootPrimary shutdown', () => { + let handlers; + let processOnSpy; + let processExitSpy; + let originalShutdownTimeout; + + beforeEach(() => { + handlers = {}; + processOnSpy = jest.spyOn(process, 'on').mockImplementation((event, handler) => { + handlers[event] = handler; + return process; + }); + processExitSpy = jest.spyOn(process, 'exit').mockImplementation(() => undefined); + originalShutdownTimeout = process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS; + process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS = '1000'; + }); + + afterEach(() => { + processOnSpy.mockRestore(); + processExitSpy.mockRestore(); + if (originalShutdownTimeout === undefined) { + delete process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS; + } else { + process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS = originalShutdownTimeout; + } + }); + + test('closes servers and flushes logs on SIGTERM before exiting', async () => { + const callOrder = []; + const oidcProvider = { initialize: jest.fn().mockResolvedValue(undefined), shutdown: jest.fn() }; + const awsOidcProvider = { initialize: jest.fn().mockResolvedValue(undefined), shutdown: jest.fn() }; + const server = { + listen: jest.fn((port, host, cb) => cb()), + shutdownConnections: jest.fn().mockImplementation(async () => { + callOrder.push('shutdownConnections'); + }), + close: jest.fn((cb) => { + callOrder.push('close'); + cb(); + }), + }; + oidcProvider.shutdown.mockImplementation(() => { + callOrder.push('oidcShutdown'); + }); + awsOidcProvider.shutdown.mockImplementation(() => { + callOrder.push('awsOidcShutdown'); + }); + const closeLogStream = jest.fn().mockImplementation(async () => { + callOrder.push('closeLogStream'); + }); + const otelShutdown = jest.fn().mockImplementation(async () => { + callOrder.push('otelShutdown'); + }); + processExitSpy.mockImplementation((code) => { + callOrder.push(`exit:${code}`); + return undefined; + }); + + bootPrimary({ + registeredAdapters: [{ + name: 'openai', + port: 10000, + alwaysBind: true, + participatesInValidation: false, + isEnabled: () => true, + getTargetHost: () => 'api.openai.com', + getOidcProvider: () => oidcProvider, + getAwsOidcProvider: () => awsOidcProvider, + }], + createProviderServer: () => server, + validateApiKeys: jest.fn(), + fetchStartupModels: jest.fn().mockResolvedValue(undefined), + writeModelsJson: jest.fn(), + validateRequestedModel: jest.fn(), + setKeyValidationComplete: jest.fn(), + setModelFetchComplete: jest.fn(), + closeLogStream, + otelShutdown, + logRequest: jest.fn(), + HTTPS_PROXY: 'http://proxy:3128', + }); + + await handlers.SIGTERM(); + + expect(server.close).toHaveBeenCalledTimes(1); + expect(server.shutdownConnections).toHaveBeenCalledTimes(1); + expect(oidcProvider.shutdown).toHaveBeenCalledTimes(1); + expect(awsOidcProvider.shutdown).toHaveBeenCalledTimes(1); + expect(closeLogStream).toHaveBeenCalledTimes(1); + expect(otelShutdown).toHaveBeenCalledTimes(1); + expect(processExitSpy).toHaveBeenCalledWith(0); + expect(callOrder.indexOf('shutdownConnections')).toBeLessThan(callOrder.indexOf('close')); + expect(callOrder.indexOf('close')).toBeLessThan(callOrder.indexOf('closeLogStream')); + expect(callOrder.indexOf('closeLogStream')).toBeLessThan(callOrder.indexOf('otelShutdown')); + expect(callOrder.indexOf('otelShutdown')).toBeLessThan(callOrder.indexOf('exit:0')); + }); +}); diff --git a/containers/api-proxy/token-persistence.js b/containers/api-proxy/token-persistence.js index bb4658db3..43f0f5e02 100644 --- a/containers/api-proxy/token-persistence.js +++ b/containers/api-proxy/token-persistence.js @@ -36,6 +36,18 @@ let logStream = null; let diagStream = null; let auditStream = null; +function ensureTokenUsageFileExists() { + try { + fs.mkdirSync(TOKEN_LOG_DIR, { recursive: true }); + const fd = fs.openSync(TOKEN_LOG_FILE, 'a', 0o644); + fs.closeSync(fd); + return true; + } catch (err) { + logRequest('warn', 'token_log_init_error', { error: err.message }); + return false; + } +} + /** * Write a diagnostic line to the diagnostics log file. * Only active when AWF_DEBUG_TOKENS=1 environment variable is set. @@ -99,9 +111,8 @@ function auditTrack(event, data) { */ function getLogStream() { if (logStream) return logStream; + if (!ensureTokenUsageFileExists()) return null; try { - // Ensure directory exists - fs.mkdirSync(TOKEN_LOG_DIR, { recursive: true }); logStream = fs.createWriteStream(TOKEN_LOG_FILE, { flags: 'a', mode: 0o644 }); logStream.on('error', (err) => { logRequest('warn', 'token_log_error', { error: err.message }); @@ -246,7 +257,15 @@ function writeTokenUsage(record) { const stream = getLogStream(); if (stream && !stream.writableEnded) { - const ok = stream.write(JSON.stringify(record) + '\n'); + const ok = stream.write(JSON.stringify(record) + '\n', () => { + try { + if (typeof stream.fd === 'number' && Number.isInteger(stream.fd)) { + fs.fdatasyncSync(stream.fd); + } + } catch { + // best-effort durability + } + }); if (!ok) { // Backpressure — stream buffer full. Drop this write rather than // accumulating unbounded memory. The 'drain' event will unblock @@ -303,3 +322,5 @@ module.exports = { writeTokenUsage, closeLogStream, }; + +ensureTokenUsageFileExists(); diff --git a/containers/api-proxy/token-tracker.http.test.js b/containers/api-proxy/token-tracker.http.test.js index cf9c9add9..1c417f62e 100644 --- a/containers/api-proxy/token-tracker.http.test.js +++ b/containers/api-proxy/token-tracker.http.test.js @@ -4,11 +4,13 @@ require('./test-helpers/token-tracker-setup'); +const fs = require('fs'); const { isStreamingResponse, isCompressedResponse, trackTokenUsage, closeLogStream, + TOKEN_LOG_FILE, } = require('./token-tracker'); const { EventEmitter } = require('events'); const zlib = require('zlib'); @@ -81,6 +83,40 @@ describe('trackTokenUsage', () => { }, 10); }); + test('writes token-usage.jsonl incrementally before shutdown', (done) => { + const proxyRes = new EventEmitter(); + proxyRes.headers = { 'content-type': 'application/json' }; + proxyRes.statusCode = 200; + + const metricsRef = { + increment: jest.fn(), + }; + + trackTokenUsage(proxyRes, { + requestId: 'test-incremental-write', + provider: 'openai', + path: '/v1/chat/completions', + startTime: Date.now(), + metrics: metricsRef, + }); + + proxyRes.emit('data', Buffer.from(JSON.stringify({ + model: 'gpt-5.4', + usage: { prompt_tokens: 12, completion_tokens: 7, total_tokens: 19 }, + }))); + proxyRes.emit('end'); + + setTimeout(() => { + expect(fs.existsSync(TOKEN_LOG_FILE)).toBe(true); + const lines = fs.readFileSync(TOKEN_LOG_FILE, 'utf8') + .split('\n') + .filter(Boolean); + const matchingLine = lines.find((line) => line.includes('"request_id":"test-incremental-write"')); + expect(matchingLine).toBeTruthy(); + done(); + }, 20); + }); + test('extracts usage from streaming SSE response', (done) => { const proxyRes = new EventEmitter(); proxyRes.headers = { 'content-type': 'text/event-stream' }; diff --git a/containers/api-proxy/token-tracker.schema.test.js b/containers/api-proxy/token-tracker.schema.test.js index 3d4ac36f8..74703260a 100644 --- a/containers/api-proxy/token-tracker.schema.test.js +++ b/containers/api-proxy/token-tracker.schema.test.js @@ -12,6 +12,7 @@ const { validateTokenUsageRecord, writeTokenUsage, closeLogStream, + TOKEN_LOG_FILE, } = require('./token-tracker'); const { buildTokenUsageRecord, @@ -187,7 +188,12 @@ function makeMockStream() { const chunks = []; const stream = { writableEnded: false, - write: jest.fn((chunk) => { chunks.push(chunk); return true; }), + fd: 123, + write: jest.fn((chunk, cb) => { + chunks.push(chunk); + if (typeof cb === 'function') cb(); + return true; + }), end: jest.fn((cb) => { stream.writableEnded = true; if (cb) cb(); }), on: jest.fn(), get writtenRecords() { @@ -246,6 +252,62 @@ describe('token-usage JSONL record schema field', () => { expect(parsed.request_id).toBe('direct-write-test'); }); + test('writeTokenUsage fdatasyncs the stream fd after a successful write callback', () => { + const fdatasyncSyncSpy = jest.spyOn(fs, 'fdatasyncSync').mockImplementation(() => undefined); + const record = { + _schema: 'token-usage/v0.0.0', + timestamp: new Date().toISOString(), + event: 'token_usage', + request_id: 'fdatasync-test', + provider: 'openai', + model: 'gpt-4o', + path: '/v1/chat/completions', + status: 200, + streaming: false, + input_tokens: 1, + output_tokens: 1, + cache_read_tokens: 0, + cache_write_tokens: 0, + duration_ms: 10, + }; + + try { + writeTokenUsage(record); + expect(fdatasyncSyncSpy).toHaveBeenCalledWith(123); + } finally { + fdatasyncSyncSpy.mockRestore(); + } + }); + + test('writeTokenUsage swallows fdatasync failures after the write callback', () => { + const fdatasyncSyncSpy = jest.spyOn(fs, 'fdatasyncSync').mockImplementation(() => { + throw new Error('disk flush failed'); + }); + const record = { + _schema: 'token-usage/v0.0.0', + timestamp: new Date().toISOString(), + event: 'token_usage', + request_id: 'fdatasync-error-test', + provider: 'openai', + model: 'gpt-4o', + path: '/v1/chat/completions', + status: 200, + streaming: false, + input_tokens: 1, + output_tokens: 1, + cache_read_tokens: 0, + cache_write_tokens: 0, + duration_ms: 10, + }; + + try { + expect(() => writeTokenUsage(record)).not.toThrow(); + expect(fdatasyncSyncSpy).toHaveBeenCalledWith(123); + } finally { + fdatasyncSyncSpy.mockRestore(); + } + }); + test('trackTokenUsage HTTP path writes versioned _schema to the stream', (done) => { const proxyRes = new EventEmitter(); proxyRes.headers = { 'content-type': 'application/json' }; @@ -440,6 +502,26 @@ describe('token-usage JSONL record schema field', () => { }); }); +describe('token-usage file sentinel', () => { + test('creates token-usage.jsonl even before first usage record', async () => { + await closeLogStream(); + if (fs.existsSync(TOKEN_LOG_FILE)) { + fs.unlinkSync(TOKEN_LOG_FILE); + } + + let isolated; + jest.isolateModules(() => { + isolated = require('./token-persistence'); + }); + + try { + expect(fs.existsSync(isolated.TOKEN_LOG_FILE)).toBe(true); + } finally { + await isolated.closeLogStream(); + } + }); +}); + // ── AWF_VERSION env var propagated as exact _schema value ───────────── // // Uses jest.isolateModules() to load a fresh token-tracker instance with a diff --git a/containers/api-proxy/websocket-proxy.js b/containers/api-proxy/websocket-proxy.js index bc91ba41b..fe207f17a 100644 --- a/containers/api-proxy/websocket-proxy.js +++ b/containers/api-proxy/websocket-proxy.js @@ -70,7 +70,7 @@ function createProxyWebSocket({ * @param {string} provider - Provider name for logging and metrics * @param {string} [basePath=''] - Optional base-path prefix */ - return function proxyWebSocket(req, socket, head, targetHost, injectHeaders, provider, basePath = '') { + return function proxyWebSocket(req, socket, head, targetHost, injectHeaders, provider, basePath = '', lifecycleHooks = {}) { const startTime = Date.now(); const clientRequestId = req.headers['x-request-id']; const requestId = isValidRequestId(clientRequestId) ? clientRequestId : generateRequestId(); @@ -122,6 +122,7 @@ function createProxyWebSocket({ requestId, startTime, upstreamPath, + onSocketsReady: lifecycleHooks.onSocketsReady, }); }; } diff --git a/containers/api-proxy/websocket-tunnel.js b/containers/api-proxy/websocket-tunnel.js index 40f82c663..3d488d14b 100644 --- a/containers/api-proxy/websocket-tunnel.js +++ b/containers/api-proxy/websocket-tunnel.js @@ -68,6 +68,7 @@ function createWebSocketTunnel({ requestId, startTime, upstreamPath, + onSocketsReady, }) { const { finalize, abort } = createProxyErrorResponder({ metrics, @@ -135,6 +136,10 @@ function createWebSocketTunnel({ if (head && head.length > 0) tlsSocket.write(head); + if (typeof onSocketsReady === 'function') { + onSocketsReady(socket, tlsSocket); + } + tlsSocket.pipe(socket); socket.pipe(tlsSocket); diff --git a/src/services/api-proxy-env-config.test.ts b/src/services/api-proxy-env-config.test.ts index 48c1db0f6..c57aedfb3 100644 --- a/src/services/api-proxy-env-config.test.ts +++ b/src/services/api-proxy-env-config.test.ts @@ -15,6 +15,7 @@ const { buildRateLimitEnv, buildModelPolicyEnv, buildOidcEnv, + resolveApiProxyShutdownTimeoutMs, } = testHelpers; const networkConfig = { @@ -128,6 +129,29 @@ describe('buildProviderRoutingEnv', () => { }); expect(env.COPILOT_INTEGRATION_ID).toBeUndefined(); }); + + it('forwards the default api-proxy shutdown timeout', () => { + const env = buildProviderRoutingEnv({ ...baseConfig, workDir: '/tmp/awf-test' }); + expect(env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS).toBe('8000'); + }); +}); + +describe('resolveApiProxyShutdownTimeoutMs', () => { + it('prefers trimmed additionalEnv values', () => { + expect(resolveApiProxyShutdownTimeoutMs({ + ...baseConfig, + workDir: '/tmp/awf-test', + additionalEnv: { AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS: ' 15000 ' }, + })).toBe(15000); + }); + + it('falls back to the default for invalid values', () => { + expect(resolveApiProxyShutdownTimeoutMs({ + ...baseConfig, + workDir: '/tmp/awf-test', + additionalEnv: { AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS: '0' }, + })).toBe(8000); + }); }); describe('buildProxyRoutingEnv', () => { diff --git a/src/services/api-proxy-env-config.ts b/src/services/api-proxy-env-config.ts index e518750ea..c35461abf 100644 --- a/src/services/api-proxy-env-config.ts +++ b/src/services/api-proxy-env-config.ts @@ -5,6 +5,8 @@ import { getConfigEnvValue, getLowerCaseProcessEnvValue, pickEnvVars } from '../ import { OPENAI_ENV, ANTHROPIC_ENV, GEMINI_ENV, COPILOT_ENV, VERTEX_ENV, OIDC_AUTH_ENV_VARS, OIDC_AUTH_ENV_MAPPING } from '../api-proxy-env-constants'; import { NetworkConfig } from './squid-service'; +const DEFAULT_API_PROXY_SHUTDOWN_TIMEOUT_MS = 8000; + /** * Builds provider API target/basePath environment variables for the api-proxy container. * Centralizes the repetitive per-provider target/basePath conditional env generation. @@ -67,6 +69,16 @@ function resolveProviderSessionId(config: WrapperConfig): string | undefined { return normalizedValue || undefined; } +export function resolveApiProxyShutdownTimeoutMs(config: WrapperConfig): number { + const rawValue = getConfigEnvValue(config, 'AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS') + ?? process.env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS; + const parsedValue = Number.parseInt(rawValue || '', 10); + if (!Number.isFinite(parsedValue) || parsedValue <= 0) { + return DEFAULT_API_PROXY_SHUTDOWN_TIMEOUT_MS; + } + return parsedValue; +} + /** * Builds API credential environment variables for the api-proxy sidecar. * These keys are passed securely to the sidecar and are NOT visible to the agent container. @@ -114,6 +126,7 @@ function buildProviderRoutingEnv(config: WrapperConfig): Record // token-usage.jsonl _schema field reflects the api-proxy image version rather than // the CLI version. This ensures correct versioning when --image-tag pins the proxy // to a different release. + AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS: String(resolveApiProxyShutdownTimeoutMs(config)), }; } @@ -306,4 +319,5 @@ export const testHelpers = { buildRateLimitEnv, buildModelPolicyEnv, buildOidcEnv, + resolveApiProxyShutdownTimeoutMs, }; diff --git a/src/services/api-proxy-service-config.test.ts b/src/services/api-proxy-service-config.test.ts index 4311f21d8..de1822bab 100644 --- a/src/services/api-proxy-service-config.test.ts +++ b/src/services/api-proxy-service-config.test.ts @@ -124,7 +124,19 @@ describe('API proxy sidecar: service configuration', () => { const configWithProxy = { ...mockConfig, enableApiProxy: true, openaiApiKey: 'sk-test-key' }; const result = generateDockerCompose(configWithProxy, mockNetworkConfigWithProxy); const proxy = result.services['api-proxy'] as any; - expect(proxy.stop_grace_period).toBe('2s'); + expect(proxy.stop_grace_period).toBe('10s'); + }); + + it('should size stop_grace_period from AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS', () => { + const configWithProxy = { + ...mockConfig, + enableApiProxy: true, + openaiApiKey: 'sk-test-key', + additionalEnv: { AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS: '15000' }, + }; + const result = generateDockerCompose(configWithProxy, mockNetworkConfigWithProxy); + const proxy = result.services['api-proxy'] as any; + expect(proxy.stop_grace_period).toBe('17s'); }); it('should set resource limits', () => { @@ -304,4 +316,17 @@ describe('API proxy sidecar: service configuration', () => { else delete process.env.AWF_PROVIDER_SESSION_ID; } }); + + it('should forward AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS from additionalEnv to api-proxy', () => { + const configWithProxy = { + ...mockConfig, + enableApiProxy: true, + openaiApiKey: 'sk-test-key', + additionalEnv: { AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS: '15000' }, + }; + const result = generateDockerCompose(configWithProxy, mockNetworkConfigWithProxy); + const proxy = result.services['api-proxy']; + const env = proxy.environment as Record; + expect(env.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS).toBe('15000'); + }); }); diff --git a/src/services/api-proxy-service-config.ts b/src/services/api-proxy-service-config.ts index 061088299..bd759125f 100644 --- a/src/services/api-proxy-service-config.ts +++ b/src/services/api-proxy-service-config.ts @@ -7,7 +7,7 @@ import { getSafeHostGid, getSafeHostUid } from '../host-identity'; import { NetworkConfig, ImageBuildConfig } from './squid-service'; import { applyHostPathPrefixToVolumes } from './host-path-prefix'; import { buildContainerSecurityHardening } from './service-security'; -import { buildApiProxyBaseEnv } from './api-proxy-env-config'; +import { buildApiProxyBaseEnv, resolveApiProxyShutdownTimeoutMs } from './api-proxy-env-config'; import { buildApiProxyLifecycleConfig } from './api-proxy-lifecycle-config'; interface ApiProxyServiceConfigParams { @@ -23,6 +23,8 @@ export function buildApiProxyServiceConfig(params: ApiProxyServiceConfigParams): throw new Error('buildApiProxyServiceConfig: networkConfig.proxyIp is required'); } const { useGHCR, registry, parsedTag, projectRoot } = imageConfig; + const shutdownTimeoutMs = resolveApiProxyShutdownTimeoutMs(config); + const stopGracePeriodSeconds = Math.ceil((shutdownTimeoutMs + 2000) / 1000); const proxyService: any = { container_name: API_PROXY_CONTAINER_NAME, @@ -38,7 +40,7 @@ export function buildApiProxyServiceConfig(params: ApiProxyServiceConfigParams): environment: buildApiProxyBaseEnv(config, networkConfig), // Security hardening and resource limits to prevent DoS attacks ...buildContainerSecurityHardening({ memLimit: '512m', pidsLimit: 100, cpuShares: 512 }), - stop_grace_period: '2s', + stop_grace_period: `${stopGracePeriodSeconds}s`, }; // Use GHCR image or build locally diff --git a/src/services/api-proxy-service-split.test.ts b/src/services/api-proxy-service-split.test.ts index 90caa34ca..c94a6322d 100644 --- a/src/services/api-proxy-service-split.test.ts +++ b/src/services/api-proxy-service-split.test.ts @@ -41,7 +41,9 @@ describe('API proxy split builders', () => { expect(service.container_name).toBe('awf-api-proxy'); expect(service.environment.OPENAI_API_KEY).toBe('sk-test-openai-key'); expect(service.environment.HTTP_PROXY).toBe('http://172.30.0.10:3128'); + expect(service.environment.AWF_API_PROXY_SHUTDOWN_TIMEOUT_MS).toBe('8000'); expect(service.image).toBe('ghcr.io/github/gh-aw-firewall/api-proxy:latest'); + expect(service.stop_grace_period).toBe('10s'); }); it('buildApiProxyBaseEnv builds proxy routing and key env', () => {