You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 
 

255 lines
8.0 KiB

import { dbGlobal } from "drizzle-pkg/lib/db";
import { agentSessions, agentMessages, agentMessageFeedback } from "drizzle-pkg/lib/schema/agent";
import { eq, desc, asc, and, lt, isNull, sql, max } from "drizzle-orm";
import type { AgentSessionRow, AgentMessageRow, AgentMessageWithFeedback } from "./types";
export async function createSession(params: {
userId?: number | null;
tempToken?: string | null;
modelId?: number | null;
expiresAt?: Date | null;
}): Promise<AgentSessionRow> {
const id = `as_${Date.now().toString(36)}_${Math.random().toString(36).slice(2, 10)}`;
const [row] = await dbGlobal
.insert(agentSessions)
.values({
id,
userId: params.userId ?? null,
tempToken: params.tempToken ?? null,
modelId: params.modelId ?? null,
expiresAt: params.expiresAt ?? null,
})
.returning();
return row!;
}
export async function getSessionById(id: string): Promise<AgentSessionRow | null> {
const [row] = await dbGlobal
.select()
.from(agentSessions)
.where(and(eq(agentSessions.id, id), isNull(agentSessions.deletedAt)))
.limit(1);
return row ?? null;
}
export async function getSessionByIdAndUser(
id: string,
userId: number | null,
tempToken: string | null,
): Promise<AgentSessionRow | null> {
const conditions = [eq(agentSessions.id, id), isNull(agentSessions.deletedAt)];
if (userId) {
conditions.push(eq(agentSessions.userId, userId));
} else if (tempToken) {
conditions.push(eq(agentSessions.tempToken, tempToken));
} else {
return null;
}
const [row] = await dbGlobal
.select()
.from(agentSessions)
.where(and(...conditions))
.limit(1);
return row ?? null;
}
export async function listSessions(params: {
userId?: number | null;
tempToken?: string | null;
page?: number;
pageSize?: number;
}): Promise<{ list: AgentSessionRow[]; total: number; page: number; pageSize: number; totalPages: number }> {
const { page = 1, pageSize = 30 } = params;
const conditions = [isNull(agentSessions.deletedAt)];
if (params.userId) {
conditions.push(eq(agentSessions.userId, params.userId));
} else if (params.tempToken) {
conditions.push(eq(agentSessions.tempToken, params.tempToken));
} else {
return { list: [], total: 0, page, pageSize, totalPages: 0 };
}
const where = and(...conditions);
const [countResult] = await dbGlobal
.select({ count: sql<number>`count(*)` })
.from(agentSessions)
.where(where);
const total = countResult?.count ?? 0;
const totalPages = Math.ceil(total / pageSize);
const list = await dbGlobal
.select()
.from(agentSessions)
.where(where)
.orderBy(desc(agentSessions.lastActiveAt))
.limit(pageSize)
.offset((page - 1) * pageSize);
return { list, total, page, pageSize, totalPages };
}
export async function updateSession(
id: string,
data: {
title?: string;
modelId?: number | null;
enableThinking?: number;
enableTools?: number;
lastActiveAt?: Date;
},
): Promise<void> {
const updates: Record<string, any> = {};
if (data.title !== undefined) updates.title = data.title;
if (data.modelId !== undefined) updates.modelId = data.modelId;
if (data.enableThinking !== undefined) updates.enableThinking = data.enableThinking;
if (data.enableTools !== undefined) updates.enableTools = data.enableTools;
if (data.lastActiveAt !== undefined) updates.lastActiveAt = data.lastActiveAt;
if (Object.keys(updates).length === 0) return;
await dbGlobal.update(agentSessions).set(updates).where(eq(agentSessions.id, id));
}
export async function softDeleteSession(id: string): Promise<void> {
await dbGlobal
.update(agentSessions)
.set({ deletedAt: new Date() })
.where(eq(agentSessions.id, id));
}
export async function getMessagesBySession(
sessionId: string,
params?: { before?: number; limit?: number; latest?: boolean; userId?: number | null },
): Promise<AgentMessageWithFeedback[]> {
const { before, limit = 50, latest = false, userId = null } = params ?? {};
const conditions = [eq(agentMessages.sessionId, sessionId)];
if (before !== undefined) {
conditions.push(lt(agentMessages.sortOrder, before));
}
const baseQuery = () =>
dbGlobal
.select({
message: agentMessages,
feedback: agentMessageFeedback.feedback,
})
.from(agentMessages)
.leftJoin(
agentMessageFeedback,
and(
eq(agentMessageFeedback.messageId, agentMessages.id),
userId ? eq(agentMessageFeedback.userId, userId) : sql`false`,
),
)
.where(and(...conditions));
if (latest) {
const rows = await baseQuery()
.orderBy(desc(agentMessages.sortOrder))
.limit(limit);
const reversed = rows.reverse();
return reversed.map((r) => ({
...r.message,
feedback: (r.feedback ?? null) as "like" | "dislike" | null,
}));
}
const rows = await baseQuery()
.orderBy(asc(agentMessages.sortOrder))
.limit(limit);
return rows.map((r) => ({
...r.message,
feedback: (r.feedback ?? null) as "like" | "dislike" | null,
}));
}
export async function getMaxSortOrder(sessionId: string): Promise<number> {
const [result] = await dbGlobal
.select({ maxSort: max(agentMessages.sortOrder) })
.from(agentMessages)
.where(eq(agentMessages.sessionId, sessionId));
return result?.maxSort ?? 0;
}
export async function saveMessage(params: {
sessionId: string;
role: string;
content: string;
parts?: string | null;
modelId?: number | null;
inputTokens?: number | null;
outputTokens?: number | null;
sortOrder: number;
}): Promise<AgentMessageRow> {
const id = `am_${Date.now().toString(36)}_${Math.random().toString(36).slice(2, 10)}`;
const [row] = await dbGlobal
.insert(agentMessages)
.values({
id,
sessionId: params.sessionId,
role: params.role,
content: params.content,
parts: params.parts ?? null,
modelId: params.modelId ?? null,
inputTokens: params.inputTokens ?? null,
outputTokens: params.outputTokens ?? null,
sortOrder: params.sortOrder,
})
.returning();
return row!;
}
export async function truncateMessagesAfter(sessionId: string, sortOrder: number): Promise<void> {
await dbGlobal
.delete(agentMessages)
.where(and(eq(agentMessages.sessionId, sessionId), sql`${agentMessages.sortOrder} > ${sortOrder}`));
}
export async function deleteMessage(messageId: string): Promise<void> {
await dbGlobal.delete(agentMessages).where(eq(agentMessages.id, messageId));
}
export async function getMessageById(messageId: string): Promise<AgentMessageRow | null> {
const [row] = await dbGlobal
.select()
.from(agentMessages)
.where(eq(agentMessages.id, messageId))
.limit(1);
return row ?? null;
}
export async function updateMessageContent(messageId: string, content: string): Promise<void> {
await dbGlobal
.update(agentMessages)
.set({ content })
.where(eq(agentMessages.id, messageId));
}
export async function updateMessageParts(messageId: string, parts: string): Promise<void> {
await dbGlobal
.update(agentMessages)
.set({ parts })
.where(eq(agentMessages.id, messageId));
}
export async function countAssistantMessages(sessionId: string): Promise<number> {
const [result] = await dbGlobal
.select({ cnt: sql<number>`cast(count(*) as integer)` })
.from(agentMessages)
.where(and(eq(agentMessages.sessionId, sessionId), eq(agentMessages.role, "assistant")));
return result?.cnt ?? 0;
}
export async function touchSession(sessionId: string): Promise<void> {
await dbGlobal
.update(agentSessions)
.set({ lastActiveAt: new Date() })
.where(eq(agentSessions.id, sessionId));
}
export async function migrateTempSessions(tempToken: string, userId: number): Promise<number> {
const result = await dbGlobal
.update(agentSessions)
.set({ userId, tempToken: null, expiresAt: null })
.where(and(eq(agentSessions.tempToken, tempToken), isNull(agentSessions.deletedAt)))
.returning({ id: agentSessions.id });
return result.length;
}