From 48c4b44f7289be80827225da92da0bb0f518222d Mon Sep 17 00:00:00 2001 From: Owen Date: Tue, 11 Aug 2026 12:19:17 -0400 Subject: [PATCH] log chat sessions to database --- messages/en-US.json | 2 + server/cleanup.ts | 2 + server/db/pg/schema/schema.ts | 70 ++++++ server/db/sqlite/schema/schema.ts | 74 ++++++ server/lib/cleanupLogs.ts | 21 +- server/private/cleanup.ts | 2 + .../routers/billing/featureLifecycle.ts | 11 + server/routers/aiGateway/logAiSession.ts | 238 ++++++++++++++++++ server/routers/aiGateway/pipeline.ts | 142 +++++------ .../aiGateway/streamAiGatewayResponse.ts | 86 +++++++ server/routers/aiGateway/targetRouting.ts | 88 ++++--- server/routers/org/updateOrg.ts | 20 ++ .../settings/general/security/page.tsx | 109 +++++++- 13 files changed, 749 insertions(+), 116 deletions(-) create mode 100644 server/routers/aiGateway/logAiSession.ts create mode 100644 server/routers/aiGateway/streamAiGatewayResponse.ts diff --git a/messages/en-US.json b/messages/en-US.json index b2594efcb..c7db842a7 100644 --- a/messages/en-US.json +++ b/messages/en-US.json @@ -3397,6 +3397,8 @@ "logRetentionActionDescription": "How long to retain action logs", "logRetentionConnectionLabel": "Network Log Retention", "logRetentionConnectionDescription": "How long to retain connection logs", + "logRetentionAISessionsLabel": "AI Gateway Session Log Retention", + "logRetentionAISessionsDescription": "How long to retain AI gateway prompt/response session logs", "logRetentionDisabled": "Disabled", "logRetention3Days": "3 days", "logRetention7Days": "7 days", diff --git a/server/cleanup.ts b/server/cleanup.ts index 3e7a335b4..3ff4ccf19 100644 --- a/server/cleanup.ts +++ b/server/cleanup.ts @@ -4,6 +4,7 @@ import { flushSiteBandwidthToDb } from "@server/routers/gerbil/receiveBandwidth" import { stopPingAccumulator } from "@server/routers/newt/pingAccumulator"; import { cleanup as wsCleanup } from "#dynamic/routers/ws"; import { shutdownUsageRecorder } from "@server/lib/aiBudgetEnforcement"; +import { shutdownAiSessionLogger } from "@server/routers/aiGateway/logAiSession"; async function cleanup() { await stopPingAccumulator(); @@ -11,6 +12,7 @@ async function cleanup() { await flushConnectionLogToDb(); await flushSiteBandwidthToDb(); await shutdownUsageRecorder(); + await shutdownAiSessionLogger(); await wsCleanup(); process.exit(0); diff --git a/server/db/pg/schema/schema.ts b/server/db/pg/schema/schema.ts index 39dfb889e..7895b0baf 100644 --- a/server/db/pg/schema/schema.ts +++ b/server/db/pg/schema/schema.ts @@ -65,6 +65,11 @@ export const orgs = pgTable("orgs", { ) // where 0 = dont keep logs and -1 = keep forever and 9001 = end of the following year .notNull() .default(0), + settingsLogRetentionDaysAISessions: integer( + "settingsLogRetentionDaysAISessions" + ) // where 0 = dont keep logs and -1 = keep forever and 9001 = end of the following year + .notNull() + .default(0), sshCaPrivateKey: text("sshCaPrivateKey"), // Encrypted SSH CA private key (PEM format) sshCaPublicKey: text("sshCaPublicKey"), // SSH CA public key (OpenSSH format) isBillingOrg: boolean("isBillingOrg"), @@ -1899,6 +1904,70 @@ export const aiBudgetBreachEvents = pgTable( ] ); +// Logs the aggregated prompt + response for a single AI gateway request, for +// session replay. One row per request (not per streaming chunk). `sessionId` +// is a fresh random id per row for now - no cross-request correlation yet, +// but the column exists so a future pass can link multiple rows into a real +// multi-turn session. +export const aiSessionLog = pgTable( + "aiSessionLog", + { + id: serial("id").primaryKey(), + sessionId: varchar("sessionId").notNull(), + orgId: varchar("orgId").references(() => orgs.orgId, { + onDelete: "cascade" + }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + capability: varchar("capability").notNull(), + resourceId: integer("resourceId").references( + () => resources.resourceId, + { onDelete: "cascade" } + ), + siteResourceId: integer("siteResourceId").references( + () => siteResources.siteResourceId, + { onDelete: "cascade" } + ), + userId: varchar("userId").references(() => users.userId, { + onDelete: "set null" + }), + requestedModel: varchar("requestedModel"), + isStream: boolean("isStream").notNull().default(false), + requestBody: text("requestBody"), + responseBody: text("responseBody"), + // True if requestBody/responseBody were cut short at + // AI_SESSION_LOG_MAX_BODY_CHARS before storage. + truncated: boolean("truncated").notNull().default(false), + statusCode: integer("statusCode"), + createdAt: bigint("createdAt", { mode: "number" }).notNull() // epoch ms + }, + (t) => [ + index("idx_ai_session_log_org_created").on(t.orgId, t.createdAt), + index("idx_ai_session_log_org_provider_created").on( + t.orgId, + t.providerId, + t.createdAt + ), + index("idx_ai_session_log_org_resource_created").on( + t.orgId, + t.resourceId, + t.createdAt + ), + index("idx_ai_session_log_org_site_resource_created").on( + t.orgId, + t.siteResourceId, + t.createdAt + ), + index("idx_ai_session_log_org_user_created").on( + t.orgId, + t.userId, + t.createdAt + ), + index("idx_ai_session_log_session").on(t.sessionId) + ] +); + export type Org = InferSelectModel; export type User = InferSelectModel; export type Site = InferSelectModel; @@ -1992,6 +2061,7 @@ export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; export type AiUsageRecord = InferSelectModel; export type AiBudgetBreachEvent = InferSelectModel; +export type AiSessionLog = InferSelectModel; export type ResourceAiProvider = InferSelectModel; export type SiteResourceAiProvider = InferSelectModel< typeof siteResourceAiProviders diff --git a/server/db/sqlite/schema/schema.ts b/server/db/sqlite/schema/schema.ts index 1c2236dfe..3fb15c0d6 100644 --- a/server/db/sqlite/schema/schema.ts +++ b/server/db/sqlite/schema/schema.ts @@ -64,6 +64,11 @@ export const orgs = sqliteTable("orgs", { ) // where 0 = dont keep logs and -1 = keep forever and 9001 = end of the following year .notNull() .default(0), + settingsLogRetentionDaysAISessions: integer( + "settingsLogRetentionDaysAISessions" + ) // where 0 = dont keep logs and -1 = keep forever and 9001 = end of the following year + .notNull() + .default(0), sshCaPrivateKey: text("sshCaPrivateKey"), // Encrypted SSH CA private key (PEM format) sshCaPublicKey: text("sshCaPublicKey"), // SSH CA public key (OpenSSH format) isBillingOrg: integer("isBillingOrg", { mode: "boolean" }), @@ -1889,6 +1894,74 @@ export const aiBudgetBreachEvents = sqliteTable( ] ); +// Logs the aggregated prompt + response for a single AI gateway request, for +// session replay. One row per request (not per streaming chunk). `sessionId` +// is a fresh random id per row for now - no cross-request correlation yet, +// but the column exists so a future pass can link multiple rows into a real +// multi-turn session. +export const aiSessionLog = sqliteTable( + "aiSessionLog", + { + id: integer("id").primaryKey({ autoIncrement: true }), + sessionId: text("sessionId").notNull(), + orgId: text("orgId").references(() => orgs.orgId, { + onDelete: "cascade" + }), + providerId: integer("providerId") + .notNull() + .references(() => aiProviders.providerId, { onDelete: "cascade" }), + capability: text("capability").notNull(), + resourceId: integer("resourceId").references( + () => resources.resourceId, + { onDelete: "cascade" } + ), + siteResourceId: integer("siteResourceId").references( + () => siteResources.siteResourceId, + { onDelete: "cascade" } + ), + userId: text("userId").references(() => users.userId, { + onDelete: "set null" + }), + requestedModel: text("requestedModel"), + isStream: integer("isStream", { mode: "boolean" }) + .notNull() + .default(false), + requestBody: text("requestBody"), + responseBody: text("responseBody"), + // True if requestBody/responseBody were cut short at + // AI_SESSION_LOG_MAX_BODY_CHARS before storage. + truncated: integer("truncated", { mode: "boolean" }) + .notNull() + .default(false), + statusCode: integer("statusCode"), + createdAt: integer("createdAt").notNull() // epoch ms + }, + (t) => [ + index("idx_ai_session_log_org_created").on(t.orgId, t.createdAt), + index("idx_ai_session_log_org_provider_created").on( + t.orgId, + t.providerId, + t.createdAt + ), + index("idx_ai_session_log_org_resource_created").on( + t.orgId, + t.resourceId, + t.createdAt + ), + index("idx_ai_session_log_org_site_resource_created").on( + t.orgId, + t.siteResourceId, + t.createdAt + ), + index("idx_ai_session_log_org_user_created").on( + t.orgId, + t.userId, + t.createdAt + ), + index("idx_ai_session_log_session").on(t.sessionId) + ] +); + export type Org = InferSelectModel; export type User = InferSelectModel; export type Site = InferSelectModel; @@ -1980,6 +2053,7 @@ export type AiModel = InferSelectModel; export type AiBudget = InferSelectModel; export type AiUsageRecord = InferSelectModel; export type AiBudgetBreachEvent = InferSelectModel; +export type AiSessionLog = InferSelectModel; export type ResourceAiProvider = InferSelectModel; export type SiteResourceAiProvider = InferSelectModel< typeof siteResourceAiProviders diff --git a/server/lib/cleanupLogs.ts b/server/lib/cleanupLogs.ts index f5b6d8b2f..ae2e0cdb7 100644 --- a/server/lib/cleanupLogs.ts +++ b/server/lib/cleanupLogs.ts @@ -3,12 +3,14 @@ import { cleanUpOldLogs as cleanUpOldAccessLogs } from "#dynamic/lib/logAccessAu import { cleanUpOldLogs as cleanUpOldActionLogs } from "#dynamic/middlewares/logActionAudit"; import { cleanUpOldLogs as cleanUpOldRequestLogs } from "@server/routers/badger/logRequestAudit"; import { cleanUpOldLogs as cleanUpOldConnectionLogs } from "#dynamic/routers/newt"; +import { cleanUpOldLogs as cleanUpOldAiSessionLogs } from "@server/routers/aiGateway/logAiSession"; import { gt, or } from "drizzle-orm"; import { cleanUpOldFingerprintSnapshots } from "@server/routers/olm/fingerprintingUtils"; import { build } from "@server/build"; export function initLogCleanupInterval() { - if (build == "saas") { // skip log cleanup for saas builds + if (build == "saas") { + // skip log cleanup for saas builds return null; } return setInterval( @@ -23,7 +25,9 @@ export function initLogCleanupInterval() { settingsLogRetentionDaysRequest: orgs.settingsLogRetentionDaysRequest, settingsLogRetentionDaysConnection: - orgs.settingsLogRetentionDaysConnection + orgs.settingsLogRetentionDaysConnection, + settingsLogRetentionDaysAISessions: + orgs.settingsLogRetentionDaysAISessions }) .from(orgs) .where( @@ -31,7 +35,8 @@ export function initLogCleanupInterval() { gt(orgs.settingsLogRetentionDaysAction, 0), gt(orgs.settingsLogRetentionDaysAccess, 0), gt(orgs.settingsLogRetentionDaysRequest, 0), - gt(orgs.settingsLogRetentionDaysConnection, 0) + gt(orgs.settingsLogRetentionDaysConnection, 0), + gt(orgs.settingsLogRetentionDaysAISessions, 0) ) ); @@ -42,7 +47,8 @@ export function initLogCleanupInterval() { settingsLogRetentionDaysAction, settingsLogRetentionDaysAccess, settingsLogRetentionDaysRequest, - settingsLogRetentionDaysConnection + settingsLogRetentionDaysConnection, + settingsLogRetentionDaysAISessions } = org; if (settingsLogRetentionDaysAction > 0) { @@ -72,6 +78,13 @@ export function initLogCleanupInterval() { settingsLogRetentionDaysConnection ); } + + if (settingsLogRetentionDaysAISessions > 0) { + await cleanUpOldAiSessionLogs( + orgId, + settingsLogRetentionDaysAISessions + ); + } } await cleanUpOldFingerprintSnapshots(365); diff --git a/server/private/cleanup.ts b/server/private/cleanup.ts index a1ad24976..c949e4e0c 100644 --- a/server/private/cleanup.ts +++ b/server/private/cleanup.ts @@ -19,6 +19,7 @@ import { flushConnectionLogToDb } from "#private/routers/newt"; import { flushSiteBandwidthToDb } from "@server/routers/gerbil/receiveBandwidth"; import { stopPingAccumulator } from "@server/routers/newt/pingAccumulator"; import { shutdownUsageRecorder } from "@server/lib/aiBudgetEnforcement"; +import { shutdownAiSessionLogger } from "@server/routers/aiGateway/logAiSession"; async function cleanup() { await stopPingAccumulator(); @@ -26,6 +27,7 @@ async function cleanup() { await flushConnectionLogToDb(); await flushSiteBandwidthToDb(); await shutdownUsageRecorder(); + await shutdownAiSessionLogger(); await rateLimitService.cleanup(); await wsCleanup(); await logStreamingManager.shutdown(); diff --git a/server/private/routers/billing/featureLifecycle.ts b/server/private/routers/billing/featureLifecycle.ts index 75d11c756..84a7b4f5a 100644 --- a/server/private/routers/billing/featureLifecycle.ts +++ b/server/private/routers/billing/featureLifecycle.ts @@ -134,6 +134,17 @@ async function capRetentionDays( ); } + if ( + org.settingsLogRetentionDaysAISessions !== null && + org.settingsLogRetentionDaysAISessions > maxRetentionDays + ) { + updates.settingsLogRetentionDaysAISessions = maxRetentionDays; + needsUpdate = true; + logger.info( + `Capping AI session log retention from ${org.settingsLogRetentionDaysAISessions} to ${maxRetentionDays} days for org ${orgId}` + ); + } + // Apply updates if needed if (needsUpdate) { await db.update(orgs).set(updates).where(eq(orgs.orgId, orgId)); diff --git a/server/routers/aiGateway/logAiSession.ts b/server/routers/aiGateway/logAiSession.ts new file mode 100644 index 000000000..4c7384497 --- /dev/null +++ b/server/routers/aiGateway/logAiSession.ts @@ -0,0 +1,238 @@ +import { randomUUID } from "crypto"; +import { logsDb, db, orgs, aiSessionLog, type AiProvider } from "@server/db"; +import type { InferInsertModel } from "drizzle-orm"; +import logger from "@server/logger"; +import { and, eq, lt } from "drizzle-orm"; +import cache from "#dynamic/lib/cache"; +import { calculateCutoffTimestamp } from "@server/lib/cleanupLogs"; +import { sanitizeString } from "@server/lib/sanitize"; +import type { AiCapability } from "@server/lib/aiCapabilities"; + +// Caps how much of the request/response body we keep per row, so a single +// huge multimodal payload can't blow up buffer memory or storage. +const AI_SESSION_LOG_MAX_BODY_CHARS = 200_000; + +type AiSessionLogInsert = InferInsertModel; + +// In-memory buffer for batching AI session log inserts, mirroring the +// approach in server/routers/badger/logRequestAudit.ts. +const sessionLogBuffer: AiSessionLogInsert[] = []; + +const BATCH_SIZE = 100; // Write to DB every 100 logs +const BATCH_INTERVAL_MS = 5000; // Or every 5 seconds, whichever comes first +const MAX_BUFFER_SIZE = 10000; // Prevent unbounded memory growth +let flushTimer: NodeJS.Timeout | null = null; +let isFlushInProgress = false; + +/** + * Flush buffered logs to database + */ +async function flushSessionLogs() { + if (sessionLogBuffer.length === 0 || isFlushInProgress) { + return; + } + + isFlushInProgress = true; + + // Take all current logs and clear buffer + const logsToWrite = sessionLogBuffer.splice(0, sessionLogBuffer.length); + + try { + // Use a transaction to ensure all inserts succeed or fail together + await logsDb.transaction(async (tx) => { + // Batch insert logs in groups of 25 to avoid overwhelming the database + const BATCH_DB_SIZE = 25; + for (let i = 0; i < logsToWrite.length; i += BATCH_DB_SIZE) { + const batch = logsToWrite.slice(i, i + BATCH_DB_SIZE); + await tx.insert(aiSessionLog).values(batch); + } + }); + logger.debug( + `Flushed ${logsToWrite.length} AI session logs to database` + ); + } catch (error) { + logger.error("Error flushing AI session logs:", error); + // On transaction error, put logs back at the front of the buffer to retry + // but only if buffer isn't too large + if (sessionLogBuffer.length < MAX_BUFFER_SIZE - logsToWrite.length) { + sessionLogBuffer.unshift(...logsToWrite); + logger.info( + `Re-queued ${logsToWrite.length} AI session logs for retry` + ); + } else { + logger.error( + `Buffer full, dropped ${logsToWrite.length} AI session logs` + ); + } + } finally { + isFlushInProgress = false; + // If buffer filled up while we were flushing, flush again + if (sessionLogBuffer.length >= BATCH_SIZE) { + flushSessionLogs().catch((err) => + logger.error("Error in follow-up AI session log flush:", err) + ); + } + } +} + +/** + * Schedule a flush if not already scheduled + */ +function scheduleFlush() { + if (flushTimer === null) { + flushTimer = setTimeout(() => { + flushTimer = null; + flushSessionLogs().catch((err) => + logger.error("Error in scheduled AI session log flush:", err) + ); + }, BATCH_INTERVAL_MS); + } +} + +/** + * Gracefully flush all pending logs (call this on shutdown) + */ +export async function shutdownAiSessionLogger() { + if (flushTimer) { + clearTimeout(flushTimer); + flushTimer = null; + } + // Force flush even if one is in progress by waiting and retrying + while (isFlushInProgress) { + await new Promise((resolve) => setTimeout(resolve, 100)); + } + await flushSessionLogs(); +} + +async function getRetentionDays(orgId: string): Promise { + // check cache first + const cached = await cache.get(`org_${orgId}_aiSessionsDays`); + if (cached !== undefined) { + return cached; + } + + const [org] = await db + .select({ + settingsLogRetentionDaysAISessions: + orgs.settingsLogRetentionDaysAISessions + }) + .from(orgs) + .where(eq(orgs.orgId, orgId)) + .limit(1); + + if (!org) { + return 0; + } + + // store the result in cache + await cache.set( + `org_${orgId}_aiSessionsDays`, + org.settingsLogRetentionDaysAISessions, + 300 + ); + + return org.settingsLogRetentionDaysAISessions; +} + +export async function cleanUpOldLogs(orgId: string, retentionDays: number) { + // calculateCutoffTimestamp returns a seconds-epoch cutoff (built for + // requestAuditLog.timestamp), but aiSessionLog.createdAt is ms-epoch to + // match aiUsageRecords - convert before comparing. + const cutoffTimestampMs = calculateCutoffTimestamp(retentionDays) * 1000; + + try { + await logsDb + .delete(aiSessionLog) + .where( + and( + lt(aiSessionLog.createdAt, cutoffTimestampMs), + eq(aiSessionLog.orgId, orgId) + ) + ); + } catch (error) { + logger.error("Error cleaning up old AI session logs:", error); + } +} + +function truncateBody(value: string): { value: string; truncated: boolean } { + if (value.length <= AI_SESSION_LOG_MAX_BODY_CHARS) { + return { value, truncated: false }; + } + return { + value: value.slice(0, AI_SESSION_LOG_MAX_BODY_CHARS), + truncated: true + }; +} + +export function logAiSession(data: { + capability: AiCapability; + provider: AiProvider; + requestedModel: string | undefined; + requestBody: unknown; + responseText: string; + isStream: boolean; + statusCode: number; + orgId: string | null; + resourceId: number | null; + siteResourceId: number | null; + requestUserId: string | null; +}): void { + (async () => { + try { + // Check retention before buffering any logs + if (data.orgId) { + const retentionDays = await getRetentionDays(data.orgId); + if (retentionDays === 0) { + // do not log + return; + } + } else { + // No org resolved for this request - nothing to govern + // retention with, so don't log it. + return; + } + + const requestBodyText = truncateBody( + JSON.stringify(data.requestBody ?? "") + ); + const responseBodyText = truncateBody(data.responseText ?? ""); + + // Prevent unbounded buffer growth - drop oldest entries if buffer is too large + if (sessionLogBuffer.length >= MAX_BUFFER_SIZE) { + const dropped = sessionLogBuffer.splice(0, BATCH_SIZE); + logger.warn( + `AI session log buffer exceeded max size (${MAX_BUFFER_SIZE}), dropped ${dropped.length} oldest entries` + ); + } + + sessionLogBuffer.push({ + sessionId: randomUUID(), + orgId: sanitizeString(data.orgId), + providerId: data.provider.providerId, + capability: data.capability, + resourceId: data.resourceId ?? undefined, + siteResourceId: data.siteResourceId ?? undefined, + userId: sanitizeString(data.requestUserId ?? undefined), + requestedModel: sanitizeString(data.requestedModel), + isStream: data.isStream, + requestBody: sanitizeString(requestBodyText.value), + responseBody: sanitizeString(responseBodyText.value), + truncated: + requestBodyText.truncated || responseBodyText.truncated, + statusCode: data.statusCode, + createdAt: Date.now() + }); + + // Flush immediately if buffer is full, otherwise schedule a flush + if (sessionLogBuffer.length >= BATCH_SIZE) { + flushSessionLogs().catch((err) => + logger.error("Error flushing AI session logs:", err) + ); + } else { + scheduleFlush(); + } + } catch (error) { + logger.error("Failed to log AI session", { error }); + } + })(); +} diff --git a/server/routers/aiGateway/pipeline.ts b/server/routers/aiGateway/pipeline.ts index 22274e049..9e1433832 100644 --- a/server/routers/aiGateway/pipeline.ts +++ b/server/routers/aiGateway/pipeline.ts @@ -63,10 +63,11 @@ import { isUsageEmpty, needsStreamUsageInjection, withStreamUsageOption, - stripInjectedUsageFrame, extractResponseModel, type AiUsage } from "@server/lib/aiUsageExtraction"; +import { streamAiGatewayResponse } from "@server/routers/aiGateway/streamAiGatewayResponse"; +import { logAiSession } from "@server/routers/aiGateway/logAiSession"; const EXIT_NODE_RANGES_CACHE_KEY = "aiGateway:exitNodeRanges"; const EXIT_NODE_RANGES_TTL_SEC = 6000; @@ -526,13 +527,20 @@ async function selectProvider( }; } -function logAiUsageAndCost(args: { +// Extracts usage/cost from a completed AI gateway request, records it for +// budget enforcement, and logs the aggregated prompt/response for session +// replay. Shared by both the direct-upstream path (below) and the +// "custom"/target routing-mode path (targetRouting.ts) so both get identical +// usage/cost tracking and session logging instead of only the direct-upstream +// path having it. +export function recordAiGatewayCompletion(args: { capability: AiCapability; provider: AiProvider; requestedModel: string | undefined; requestBody: unknown; responseText: string; isStream: boolean; + statusCode: number; headers: Headers; orgId: string | null; resourceId: number | null; @@ -547,6 +555,7 @@ function logAiUsageAndCost(args: { requestBody, responseText, isStream, + statusCode, headers, orgId, resourceId, @@ -608,6 +617,20 @@ function logAiUsageAndCost(args: { }); } } + + logAiSession({ + capability, + provider, + requestedModel, + requestBody, + responseText, + isStream, + statusCode, + orgId, + resourceId, + siteResourceId, + requestUserId + }); } export async function handleAiGatewayProxy( @@ -722,7 +745,14 @@ export async function handleAiGatewayProxy( res, provider, requestUser, - capability + capability, + { + orgId, + resourceId, + siteResourceId, + requestedModel, + budgets: appliedBudgets + } ); } @@ -843,84 +873,38 @@ export async function handleAiGatewayProxy( throw fetchError; } - const contentType = upstreamRes.headers.get("content-type") || ""; - const isStream = def.isStreaming(req, contentType); + const isStream = def.isStreaming( + req, + upstreamRes.headers.get("content-type") || "" + ); - res.status(upstreamRes.status); - res.setHeader("Content-Type", contentType || "application/json"); - - if (isStream && upstreamRes.body) { - res.flushHeaders(); - const reader = upstreamRes.body.getReader(); - const decoder = new TextDecoder(); - let fullText = ""; - // Frame-boundary buffer, only used when we need to filter the - // usage-only frame we injected out of what reaches the client. - let sseCarry = ""; - try { - while (!abortController.signal.aborted) { - const { done, value } = await reader.read(); - if (done) break; - const chunkText = decoder.decode(value, { stream: true }); - fullText += chunkText; - if (injectedUsageOurselves) { - sseCarry += chunkText; - const lastBoundary = sseCarry.lastIndexOf("\n\n"); - if (lastBoundary !== -1) { - const toEmit = sseCarry.slice(0, lastBoundary + 2); - sseCarry = sseCarry.slice(lastBoundary + 2); - res.write(stripInjectedUsageFrame(toEmit)); - } - } else { - res.write(value); - } - } - if (injectedUsageOurselves && sseCarry) { - res.write(stripInjectedUsageFrame(sseCarry)); - } - } finally { - await reader.cancel().catch(() => {}); - res.off("close", onClientClose); - } - if (!res.writableEnded) { - res.end(); - } - if (!abortController.signal.aborted) { - logAiUsageAndCost({ - capability, - provider, - requestedModel, - requestBody: req.body, - responseText: fullText, - isStream: true, - headers: upstreamRes.headers, - orgId, - resourceId, - siteResourceId, - requestUserId: requestUser?.userId ?? null, - budgets: appliedBudgets - }); - } - return; - } - - res.off("close", onClientClose); - const text = await upstreamRes.text(); - logAiUsageAndCost({ - capability, - provider, - requestedModel, - requestBody: req.body, - responseText: text, - isStream: false, - headers: upstreamRes.headers, - orgId, - resourceId, - siteResourceId, - requestUserId: requestUser?.userId ?? null, - budgets: appliedBudgets + const { fullText, aborted } = await streamAiGatewayResponse({ + res, + upstreamRes, + isStream, + injectedUsageOurselves, + abortController, + onClientClose }); - return res.send(text); + + if (!aborted) { + recordAiGatewayCompletion({ + capability, + provider, + requestedModel, + requestBody: outboundBody, + responseText: fullText, + isStream, + statusCode: upstreamRes.status, + headers: upstreamRes.headers, + orgId, + resourceId, + siteResourceId, + requestUserId: requestUser?.userId ?? null, + budgets: appliedBudgets + }); + } + return; } catch (error) { logger.error(error); return res.status(HttpCode.INTERNAL_SERVER_ERROR).json({ diff --git a/server/routers/aiGateway/streamAiGatewayResponse.ts b/server/routers/aiGateway/streamAiGatewayResponse.ts new file mode 100644 index 000000000..8cd369350 --- /dev/null +++ b/server/routers/aiGateway/streamAiGatewayResponse.ts @@ -0,0 +1,86 @@ +import { Response } from "express"; +import { stripInjectedUsageFrame } from "@server/lib/aiUsageExtraction"; + +/** + * Reads an upstream AI provider response, writes it through to the client + * (streaming or buffered), and returns the full response text once done, so + * the caller can extract usage/cost and log the completed session. Shared by + * both the direct-upstream path (pipeline.ts) and the "custom"/target + * routing-mode path (targetRouting.ts) so usage/cost tracking and session + * logging apply identically to both instead of each maintaining its own copy + * of this loop. + * + * Callers own fetching the upstream response and the AbortController/ + * `res.on("close", onClientClose)` wiring, since those differ meaningfully + * between the two transports (direct upstream fetch with TLS-skip support vs + * a plain fetch to gerbil) - only the "read the stream, write to the client, + * accumulate the full text" part is actually identical logic between them. + */ +export async function streamAiGatewayResponse(args: { + res: Response; + upstreamRes: globalThis.Response; + isStream: boolean; + // True when we injected stream_options.include_usage ourselves (the + // caller didn't ask for it) and need to strip the extra usage-only frame + // back out of what's forwarded to the client. + injectedUsageOurselves: boolean; + abortController: AbortController; + onClientClose: () => void; +}): Promise<{ fullText: string; aborted: boolean }> { + const { + res, + upstreamRes, + isStream, + injectedUsageOurselves, + abortController, + onClientClose + } = args; + + const contentType = upstreamRes.headers.get("content-type") || ""; + res.status(upstreamRes.status); + res.setHeader("Content-Type", contentType || "application/json"); + + if (isStream && upstreamRes.body) { + res.flushHeaders(); + const reader = upstreamRes.body.getReader(); + const decoder = new TextDecoder(); + let fullText = ""; + // Frame-boundary buffer, only used when we need to filter the + // usage-only frame we injected out of what reaches the client. + let sseCarry = ""; + try { + while (!abortController.signal.aborted) { + const { done, value } = await reader.read(); + if (done) break; + const chunkText = decoder.decode(value, { stream: true }); + fullText += chunkText; + if (injectedUsageOurselves) { + sseCarry += chunkText; + const lastBoundary = sseCarry.lastIndexOf("\n\n"); + if (lastBoundary !== -1) { + const toEmit = sseCarry.slice(0, lastBoundary + 2); + sseCarry = sseCarry.slice(lastBoundary + 2); + res.write(stripInjectedUsageFrame(toEmit)); + } + } else { + res.write(value); + } + } + if (injectedUsageOurselves && sseCarry) { + res.write(stripInjectedUsageFrame(sseCarry)); + } + } finally { + await reader.cancel().catch(() => {}); + res.off("close", onClientClose); + } + if (!res.writableEnded) { + res.end(); + } + return { fullText, aborted: abortController.signal.aborted }; + } + + res.off("close", onClientClose); + const text = await upstreamRes.text(); + res.send(text); + return { fullText: text, aborted: abortController.signal.aborted }; +} diff --git a/server/routers/aiGateway/targetRouting.ts b/server/routers/aiGateway/targetRouting.ts index a6d08d925..00ef39d12 100644 --- a/server/routers/aiGateway/targetRouting.ts +++ b/server/routers/aiGateway/targetRouting.ts @@ -1,6 +1,13 @@ import { Request, Response } from "express"; import { and, eq } from "drizzle-orm"; -import { AiProvider, db, exitNodes, sites, targets } from "@server/db"; +import { + AiBudget, + AiProvider, + db, + exitNodes, + sites, + targets +} from "@server/db"; import config from "@server/lib/config"; import { decrypt } from "@server/lib/crypto"; import { localCache } from "@server/lib/cache"; @@ -14,12 +21,18 @@ import { AI_CAPABILITY_DEFS, type AiCapability } from "@server/lib/aiCapabilities"; +import { + needsStreamUsageInjection, + withStreamUsageOption +} from "@server/lib/aiUsageExtraction"; import logger from "@server/logger"; import HttpCode from "@server/types/HttpCode"; import { applyRequestUserHeaders, + recordAiGatewayCompletion, type RequestUser } from "@server/routers/aiGateway/pipeline"; +import { streamAiGatewayResponse } from "@server/routers/aiGateway/streamAiGatewayResponse"; // Short TTL: long enough to spare the DB on a burst of requests, short // enough that target/site changes (added, removed, exit node moved) show up @@ -157,7 +170,14 @@ export async function proxyAiGatewayToSiteTarget( res: Response, provider: AiProvider, requestUser: RequestUser | null, - capability: AiCapability + capability: AiCapability, + ctx: { + orgId: string | null; + resourceId: number | null; + siteResourceId: number | null; + requestedModel: string | undefined; + budgets: AiBudget[]; + } ): Promise { const providerTargets = await getProviderTargets(provider.providerId); if (providerTargets.length === 0) { @@ -205,7 +225,17 @@ export async function proxyAiGatewayToSiteTarget( headers[PANGOLIN_DEST_HEADER] = target.destination; headers[PANGOLIN_HOST_HEADER] = target.hostHeader; - const body = JSON.stringify(req.body); + // Same OpenAI stream_options.include_usage injection direct-upstream + // requests get (pipeline.ts) - needed here too now that target-routed + // requests get usage/cost tracking and session logging as well. + const injectedUsageOurselves = needsStreamUsageInjection( + capability, + req.body + ); + const outboundBody = injectedUsageOurselves + ? withStreamUsageOption(req.body) + : req.body; + const body = JSON.stringify(outboundBody); logger.debug("AI gateway target-routed request", { providerId: provider.providerId, @@ -214,7 +244,7 @@ export async function proxyAiGatewayToSiteTarget( hostHeader: target.hostHeader, url: gerbilUrl, headers, - body: req.body + body: outboundBody }); // Cancel the request to gerbil (which cascades to gerbil cancelling its @@ -259,35 +289,35 @@ export async function proxyAiGatewayToSiteTarget( return; } - const contentType = upstreamRes.headers.get("content-type") || ""; const isStream = AI_CAPABILITY_DEFS[capability].isStreaming( req, - contentType + upstreamRes.headers.get("content-type") || "" ); - res.status(upstreamRes.status); - res.setHeader("Content-Type", contentType || "application/json"); + const { fullText, aborted } = await streamAiGatewayResponse({ + res, + upstreamRes, + isStream, + injectedUsageOurselves, + abortController, + onClientClose + }); - if (isStream && upstreamRes.body) { - res.flushHeaders(); - const reader = upstreamRes.body.getReader(); - try { - while (!abortController.signal.aborted) { - const { done, value } = await reader.read(); - if (done) break; - res.write(value); - } - } finally { - await reader.cancel().catch(() => {}); - res.off("close", onClientClose); - } - if (!res.writableEnded) { - res.end(); - } - return; + if (!aborted) { + recordAiGatewayCompletion({ + capability, + provider, + requestedModel: ctx.requestedModel, + requestBody: outboundBody, + responseText: fullText, + isStream, + statusCode: upstreamRes.status, + headers: upstreamRes.headers, + orgId: ctx.orgId, + resourceId: ctx.resourceId, + siteResourceId: ctx.siteResourceId, + requestUserId: requestUser?.userId ?? null, + budgets: ctx.budgets + }); } - - res.off("close", onClientClose); - const text = await upstreamRes.text(); - res.send(text); } diff --git a/server/routers/org/updateOrg.ts b/server/routers/org/updateOrg.ts index f98bdec27..021448375 100644 --- a/server/routers/org/updateOrg.ts +++ b/server/routers/org/updateOrg.ts @@ -41,6 +41,10 @@ const updateOrgBodySchema = z .number() .min(build === "saas" ? 0 : -1) .optional(), + settingsLogRetentionDaysAISessions: z + .number() + .min(build === "saas" ? 0 : -1) + .optional(), settingsEnableGlobalNewtAutoUpdate: z.boolean().optional() }) .refine((data) => Object.keys(data).length > 0, { @@ -212,6 +216,19 @@ export async function updateOrg( ) ); } + if ( + parsedBody.data.settingsLogRetentionDaysAISessions !== + undefined && + parsedBody.data.settingsLogRetentionDaysAISessions > + maxRetentionDays + ) { + return next( + createHttpError( + HttpCode.FORBIDDEN, + `You are not allowed to set log retention days greater than ${maxRetentionDays} with your current subscription` + ) + ); + } } } @@ -230,6 +247,8 @@ export async function updateOrg( parsedBody.data.settingsLogRetentionDaysAction, settingsLogRetentionDaysConnection: parsedBody.data.settingsLogRetentionDaysConnection, + settingsLogRetentionDaysAISessions: + parsedBody.data.settingsLogRetentionDaysAISessions, settingsEnableGlobalNewtAutoUpdate: parsedBody.data.settingsEnableGlobalNewtAutoUpdate }) @@ -250,6 +269,7 @@ export async function updateOrg( await cache.del(`org_${orgId}_actionDays`); await cache.del(`org_${orgId}_accessDays`); await cache.del(`org_${orgId}_connectionDays`); + await cache.del(`org_${orgId}_aiSessionsDays`); return response(res, { data: updatedOrg[0], diff --git a/src/app/[orgId]/settings/general/security/page.tsx b/src/app/[orgId]/settings/general/security/page.tsx index 51afa6077..7f83711b5 100644 --- a/src/app/[orgId]/settings/general/security/page.tsx +++ b/src/app/[orgId]/settings/general/security/page.tsx @@ -80,7 +80,8 @@ const SecurityFormSchema = z.object({ settingsLogRetentionDaysRequest: z.number(), settingsLogRetentionDaysAccess: z.number(), settingsLogRetentionDaysAction: z.number(), - settingsLogRetentionDaysConnection: z.number() + settingsLogRetentionDaysConnection: z.number(), + settingsLogRetentionDaysAISessions: z.number() }); const LOG_RETENTION_OPTIONS = [ @@ -122,7 +123,8 @@ function LogRetentionSectionForm({ org }: SectionFormProps) { settingsLogRetentionDaysRequest: true, settingsLogRetentionDaysAccess: true, settingsLogRetentionDaysAction: true, - settingsLogRetentionDaysConnection: true + settingsLogRetentionDaysConnection: true, + settingsLogRetentionDaysAISessions: true }) ), defaultValues: { @@ -133,7 +135,9 @@ function LogRetentionSectionForm({ org }: SectionFormProps) { settingsLogRetentionDaysAction: org.settingsLogRetentionDaysAction ?? 15, settingsLogRetentionDaysConnection: - org.settingsLogRetentionDaysConnection ?? 15 + org.settingsLogRetentionDaysConnection ?? 15, + settingsLogRetentionDaysAISessions: + org.settingsLogRetentionDaysAISessions ?? 15 }, mode: "onChange" }); @@ -161,7 +165,9 @@ function LogRetentionSectionForm({ org }: SectionFormProps) { settingsLogRetentionDaysAction: data.settingsLogRetentionDaysAction, settingsLogRetentionDaysConnection: - data.settingsLogRetentionDaysConnection + data.settingsLogRetentionDaysConnection, + settingsLogRetentionDaysAISessions: + data.settingsLogRetentionDaysAISessions } as any; // Update organization @@ -292,6 +298,101 @@ function LogRetentionSectionForm({ org }: SectionFormProps) { )} /> + ( + + + {t("logRetentionAISessionsLabel")} + + + + + + + )} + /> + {!env.flags.disableEnterpriseFeatures && ( <>