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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 54 additions & 3 deletions containers/api-proxy/server-factory.js
Original file line number Diff line number Diff line change
Expand Up @@ -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();
Expand All @@ -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) {
Expand All @@ -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;
Expand Down Expand Up @@ -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;
}
Expand Down
48 changes: 48 additions & 0 deletions containers/api-proxy/server-factory.test.js
Original file line number Diff line number Diff line change
@@ -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);
});
});
27 changes: 27 additions & 0 deletions containers/api-proxy/startup.js
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -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}`,
Expand All @@ -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();
Expand All @@ -110,6 +136,7 @@ function bootPrimary({
}
await closeLogStream();
await otelShutdown();
clearTimeout(forceExitTimer);
process.exit(0);
}

Expand Down
101 changes: 101 additions & 0 deletions containers/api-proxy/startup.test.js
Original file line number Diff line number Diff line change
@@ -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'));
});
});
27 changes: 24 additions & 3 deletions containers/api-proxy/token-persistence.js
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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 });
Expand Down Expand Up @@ -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);
Comment on lines +260 to +263
}
} catch {
// best-effort durability
}
});
if (!ok) {
// Backpressure — stream buffer full. Drop this write rather than
// accumulating unbounded memory. The 'drain' event will unblock
Expand Down Expand Up @@ -303,3 +322,5 @@ module.exports = {
writeTokenUsage,
closeLogStream,
};

ensureTokenUsageFileExists();
Loading
Loading