diff --git a/README.md b/README.md index 59d65e2..5b3b1c4 100644 --- a/README.md +++ b/README.md @@ -55,7 +55,7 @@ By default, ordinary clients may use only `minimax-m3`, `minimax-m2.7`, `minimax ALLOWED_MODELS=kimi-k3,gpt-4.1 ``` -Set `DATABASE_URL` in `.env` (or the process environment) to a Postgres connection string. The proxy creates its `requests` table and index automatically on startup. Only requests to the OpenAI Chat Completions and Responses endpoints and the Anthropic Messages endpoints are stored. Their complete responses—including streaming responses and errors—are stored with headers, status, model, client IP, and timestamp. On startup, stored requests for all other endpoints are deleted. Successful chat completion requests also store a normalized training conversation containing the full chat history and generated assistant output, including reasoning, tool calls, and tool results. +Set `DATABASE_URL` in `.env` (or the process environment) to a Postgres connection string. The proxy creates its `requests` table and index automatically on startup. Successful requests are stored for the OpenAI Chat Completions and Responses endpoints and the Anthropic Messages endpoints; failed requests from every endpoint are retained for diagnostics. Complete responses—including streaming responses and errors—are stored with headers, status, model, client IP, and timestamp. On startup, successful requests for all other endpoints are deleted. Successful chat completion requests also store a normalized training conversation containing the full chat history and generated assistant output, including reasoning, tool calls, and tool results. ```sh npm start diff --git a/src/db.js b/src/db.js index cdd4f36..cae7ae1 100644 --- a/src/db.js +++ b/src/db.js @@ -42,11 +42,34 @@ export async function ensureSchema(pool) { export async function removeNonChatCompletionRequests(pool) { const result = await pool.query(` DELETE FROM requests - WHERE split_part(endpoint, '?', 1) <> ALL($1::text[]) + WHERE response_status BETWEEN 200 AND 299 + AND split_part(endpoint, '?', 1) <> ALL($1::text[]) `, [CHAT_COMPLETION_ENDPOINTS]); return result.rowCount ?? 0; } +export async function selectTrainingRequests(pool, { cursor = null, limit = 1_000 } = {}) { + const params = cursor + ? [cursor.cursor_created_at, cursor.id, CHAT_COMPLETION_ENDPOINTS, limit] + : [CHAT_COMPLETION_ENDPOINTS, limit]; + const cursorClause = cursor ? 'AND (created_at, id) < ($1::timestamptz, $2::bigint)' : ''; + const endpointParameter = cursor ? '$3' : '$1'; + const limitParameter = cursor ? '$4' : '$2'; + + return pool.query( + `SELECT id, created_at, created_at::text AS cursor_created_at, endpoint, request_body, response_body, + response_status, duration_ms, stream, training_data + FROM requests + WHERE (method = 'POST' OR method IS NULL) + AND response_status BETWEEN 200 AND 299 + AND split_part(endpoint, '?', 1) = ANY(${endpointParameter}::text[]) + ${cursorClause} + ORDER BY created_at DESC, id DESC + LIMIT ${limitParameter}`, + params, + ); +} + export async function insertRequest(pool, entry) { await pool.query( `INSERT INTO requests ( diff --git a/src/export.js b/src/export.js index 74d8613..aebe76a 100644 --- a/src/export.js +++ b/src/export.js @@ -3,10 +3,9 @@ import { once } from 'node:events'; import { createWriteStream } from 'node:fs'; import { join, dirname } from 'node:path'; import { fileURLToPath } from 'node:url'; -import { createPool, ensureSchema } from './db.js'; +import { createPool, ensureSchema, selectTrainingRequests } from './db.js'; import { createTrainingExample } from './training-data.js'; import { createFullTraceSelector } from './traces.js'; -import { CHAT_COMPLETION_ENDPOINTS } from './chat-completion-endpoints.js'; const __dirname = dirname(fileURLToPath(import.meta.url)); @@ -65,24 +64,7 @@ try { while (limit === null || exported < limit) { const remaining = limit === null ? 1_000 : limit - exported; const batchSize = Math.min(1_000, Math.max(100, remaining * 2)); - const params = cursor - ? [cursor.cursor_created_at, cursor.id, CHAT_COMPLETION_ENDPOINTS, batchSize] - : [CHAT_COMPLETION_ENDPOINTS, batchSize]; - const cursorClause = cursor ? 'AND (created_at, id) < ($1::timestamptz, $2::bigint)' : ''; - const endpointParameter = cursor ? '$3' : '$1'; - const batchParameter = cursor ? '$4' : '$2'; - const result = await pool.query( - `SELECT id, created_at, created_at::text AS cursor_created_at, endpoint, request_body, response_body, - response_status, duration_ms, stream, training_data - FROM requests - WHERE (method = 'POST' OR method IS NULL) - AND response_status BETWEEN 200 AND 299 - AND split_part(endpoint, '?', 1) = ANY(${endpointParameter}::text[]) - ${cursorClause} - ORDER BY created_at DESC, id DESC - LIMIT ${batchParameter}`, - params, - ); + const result = await selectTrainingRequests(pool, { cursor, limit: batchSize }); for (const row of result.rows) { const example = row.training_data || createTrainingExample({ diff --git a/src/proxy.js b/src/proxy.js index 13f903d..c212d61 100644 --- a/src/proxy.js +++ b/src/proxy.js @@ -19,6 +19,12 @@ import { CHAT_COMPLETION_ENDPOINTS } from './chat-completion-endpoints.js'; export const DEFAULT_UPSTREAM_BASE_URL = 'https://opencode.ai/zen/go/v1'; +function shouldStoreRequest(method, endpoint, responseStatus) { + const pathname = endpoint.split('?')[0]; + const successful = responseStatus >= 200 && responseStatus < 300; + return !successful || (method === 'POST' && CHAT_COMPLETION_ENDPOINTS.includes(pathname)); +} + export function collectRequestData(requestLogger) { return (request, response, next) => { const startedAt = process.hrtime.bigint(); @@ -41,10 +47,12 @@ export function collectRequestData(requestLogger) { return originalEnd(...args); }; - const store = () => { + const store = (completed = true) => { if (stored) return; stored = true; const responseBody = Buffer.concat(chunks).toString('utf8'); + const responseStatus = completed ? response.statusCode : 499; + if (!shouldStoreRequest(request.method, request.originalUrl, responseStatus)) return; const entry = { method: request.method, endpoint: request.originalUrl, @@ -52,7 +60,7 @@ export function collectRequestData(requestLogger) { requestBody: request.body ?? null, responseHeaders: response.getHeaders(), responseBody, - responseStatus: response.statusCode, + responseStatus, durationMs: Number(process.hrtime.bigint() - startedAt) / 1_000_000, model: request.body?.model ?? null, clientIp: clientIp(request), @@ -74,8 +82,8 @@ export function collectRequestData(requestLogger) { } }; - response.once('finish', store); - response.once('close', store); + response.once('finish', () => store(true)); + response.once('close', () => store(response.writableFinished)); next(); }; } @@ -204,7 +212,7 @@ function validateAllowedModel(model, allowedModels) { export function createProxyApp({ keys, keyRotator = createKeyRotator(keys), fetchImpl = globalThis.fetch, upstreamBaseUrl = DEFAULT_UPSTREAM_BASE_URL, maxContextTokens = maxContextFromEnv(), rateLimiter = createRateLimiter(), unlimitedKeys = unlimitedKeysFromEnv(), allowedModels = allowedModelsFromEnv(), pagesDirectory = defaultPagesDirectory, requestLogger } = {}) { const rotator = keyRotator; const app = express(); - if (requestLogger) app.post(CHAT_COMPLETION_ENDPOINTS, collectRequestData(requestLogger)); + if (requestLogger) app.use(collectRequestData(requestLogger)); app.use(express.json({ limit: '10mb' })); app.get('/', (_request, response) => { diff --git a/test/db.test.js b/test/db.test.js index a23a0e3..af1b871 100644 --- a/test/db.test.js +++ b/test/db.test.js @@ -1,5 +1,11 @@ import { describe, expect, it, vi } from 'vitest'; -import { createRequestLogger, ensureSchema, insertRequest, removeNonChatCompletionRequests } from '../src/db.js'; +import { + createRequestLogger, + ensureSchema, + insertRequest, + removeNonChatCompletionRequests, + selectTrainingRequests, +} from '../src/db.js'; describe('request database', () => { it('creates the requests table, compatible columns, and recent-request index', async () => { @@ -37,12 +43,13 @@ describe('request database', () => { ]); }); - it('removes requests that are not for the chat completions endpoint', async () => { + it('removes successful non-training requests but retains failures', async () => { const pool = { query: vi.fn().mockResolvedValue({ rowCount: 4 }) }; await expect(removeNonChatCompletionRequests(pool)).resolves.toBe(4); expect(pool.query).toHaveBeenCalledOnce(); + expect(pool.query.mock.calls[0][0]).toContain('response_status BETWEEN 200 AND 299'); expect(pool.query.mock.calls[0][0]).toContain("split_part(endpoint, '?', 1) <> ALL($1::text[])"); expect(pool.query.mock.calls[0][1]).toEqual([[ '/oai/v1/chat/completions', '/oai/v1/responses', @@ -50,6 +57,20 @@ describe('request database', () => { ]]); }); + it('selects only successful training requests for export', async () => { + const pool = { query: vi.fn().mockResolvedValue({ rows: [] }) }; + + await selectTrainingRequests(pool, { limit: 250 }); + + const [sql, params] = pool.query.mock.calls[0]; + expect(sql).toContain('response_status BETWEEN 200 AND 299'); + expect(sql).toContain("split_part(endpoint, '?', 1) = ANY($1::text[])"); + expect(params).toEqual([[ + '/oai/v1/chat/completions', '/oai/v1/responses', + '/ant/v1/messages', '/ant/v1/v1/messages', + ], 250]); + }); + it('reports storage failures without rejecting request handling', async () => { const onError = vi.fn(); const logger = createRequestLogger({ query: vi.fn().mockRejectedValue(new Error('database offline')) }, { onError }); diff --git a/test/proxy.test.js b/test/proxy.test.js index ecae36f..d118ef5 100644 --- a/test/proxy.test.js +++ b/test/proxy.test.js @@ -1,6 +1,13 @@ import { afterEach, beforeEach, describe, expect, it, vi } from 'vitest'; +import { EventEmitter } from 'node:events'; import request from 'supertest'; -import { allowedModelsFromEnv, createProxyApp, DEFAULT_ALLOWED_MODELS, unlimitedKeysFromEnv } from '../src/proxy.js'; +import { + allowedModelsFromEnv, + collectRequestData, + createProxyApp, + DEFAULT_ALLOWED_MODELS, + unlimitedKeysFromEnv, +} from '../src/proxy.js'; const response = (body, options = {}) => new Response(body, { status: options.status ?? 200, @@ -254,7 +261,7 @@ describe('proxy', () => { }); }); - it('does not collect requests to non-chat-completion endpoints', async () => { + it('collects failed requests outside training endpoints but skips successful ones', async () => { const fetchImpl = vi.fn().mockResolvedValue(response('{"object":"list","data":[]}')); const requestLogger = { log: vi.fn() }; const app = createProxyApp({ keys: ['key'], fetchImpl, requestLogger }); @@ -262,7 +269,53 @@ describe('proxy', () => { await request(app).get('/missing').expect(404); await request(app).get('/oai/v1/models').expect(200); - expect(requestLogger.log).not.toHaveBeenCalled(); + expect(requestLogger.log).toHaveBeenCalledOnce(); + expect(requestLogger.log.mock.calls[0][0]).toMatchObject({ + method: 'GET', endpoint: '/missing', requestBody: null, responseStatus: 404, + trainingData: null, + }); + expect(requestLogger.log.mock.calls[0][0].responseBody).toContain('Cannot GET /missing'); + }); + + it('collects failed training requests without creating export data', async () => { + const fetchImpl = vi.fn().mockResolvedValue(response('{"error":"upstream failed"}', { status: 503 })); + const requestLogger = { log: vi.fn() }; + const app = createProxyApp({ keys: ['key'], fetchImpl, requestLogger }); + + await request(app).post('/oai/v1/chat/completions').send({ + model: 'm', messages: [{ role: 'user', content: 'Hi' }], + }).expect(503); + + expect(requestLogger.log).toHaveBeenCalledOnce(); + expect(requestLogger.log.mock.calls[0][0]).toMatchObject({ + endpoint: '/oai/v1/chat/completions', responseStatus: 503, + responseBody: '{"error":"upstream failed"}', trainingData: null, + }); + }); + + it('records an interrupted response as a failed request', () => { + const requestLogger = { log: vi.fn() }; + const middleware = collectRequestData(requestLogger); + const incoming = { + method: 'GET', originalUrl: '/oai/v1/models', headers: {}, body: undefined, + ip: '198.51.100.20', socket: {}, get: vi.fn(), + }; + const outgoing = Object.assign(new EventEmitter(), { + statusCode: 200, + writableFinished: false, + write: vi.fn(() => true), + end: vi.fn(), + getHeaders: vi.fn(() => ({})), + }); + + middleware(incoming, outgoing, vi.fn()); + outgoing.emit('close'); + outgoing.emit('finish'); + + expect(requestLogger.log).toHaveBeenCalledOnce(); + expect(requestLogger.log.mock.calls[0][0]).toMatchObject({ + endpoint: '/oai/v1/models', responseStatus: 499, trainingData: null, + }); }); it('collects OpenAI Responses and Anthropic Messages requests', async () => {