mirror of
https://github.com/danny-avila/LibreChat.git
synced 2026-08-04 14:57:42 +00:00
🪣 fix: Cap Context Projection Workload Before Tokenization (#13910)
* fix: bound context projection workload * fix: Address context projection CI failures * fix: Bound context projection database reads * fix: Sort projection spec imports * fix: Cap projection body reads with stats
This commit is contained in:
parent
2f800c5b52
commit
77854decdf
8 changed files with 606 additions and 18 deletions
206
packages/api/src/endpoints/projection.spec.ts
Normal file
206
packages/api/src/endpoints/projection.spec.ts
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
import { resolveContextProjection } from './projection';
|
||||
import { QUOTE_MAX_COUNT } from '~/utils/quotes';
|
||||
|
||||
jest.mock('@librechat/agents', () => ({
|
||||
Providers: { OPENAI: 'openai' },
|
||||
createTokenCounter: jest.fn(async () => jest.fn(() => 1)),
|
||||
projectAgentContextUsage: jest.fn(() => ({ tokenCount: 1, maxContextTokens: 1000 })),
|
||||
}));
|
||||
|
||||
const GRAPH_SELECT = 'messageId parentMessageId metadata.summaryUsedTokens';
|
||||
const BODY_SELECT = 'messageId parentMessageId tokenCount isCreatedByUser text quotes';
|
||||
|
||||
function textStats(messageId: string, textBytes = 5) {
|
||||
return {
|
||||
messageId,
|
||||
textBytes,
|
||||
quoteCount: 0,
|
||||
quoteBytes: 0,
|
||||
quoteLineCount: 0,
|
||||
nonStringQuoteCount: 0,
|
||||
};
|
||||
}
|
||||
|
||||
describe('resolveContextProjection', () => {
|
||||
const baseParams = {
|
||||
conversationId: 'conversation-1',
|
||||
messageId: 'message-1',
|
||||
endpoint: 'openai',
|
||||
maxContextTokens: 1000,
|
||||
model: 'gpt-4o',
|
||||
};
|
||||
|
||||
beforeEach(() => {
|
||||
jest.clearAllMocks();
|
||||
});
|
||||
|
||||
it('returns null before tokenization when the conversation is too large', async () => {
|
||||
const { createTokenCounter } = jest.requireMock('@librechat/agents');
|
||||
const messages = Array.from({ length: 513 }, (_, index) => ({
|
||||
messageId: `message-${index}`,
|
||||
parentMessageId: index === 0 ? null : `message-${index - 1}`,
|
||||
isCreatedByUser: true,
|
||||
text: 'hello',
|
||||
}));
|
||||
const getMessages = jest.fn(async () => messages);
|
||||
const getMessageTextStats = jest.fn();
|
||||
|
||||
const result = await resolveContextProjection(
|
||||
{ userId: 'user-1', getMessages, getMessageTextStats },
|
||||
{ ...baseParams, messageId: 'message-512' },
|
||||
);
|
||||
|
||||
expect(result).toBeNull();
|
||||
expect(getMessages).toHaveBeenCalledTimes(1);
|
||||
expect(getMessages).toHaveBeenCalledWith(
|
||||
{ conversationId: 'conversation-1', user: 'user-1' },
|
||||
GRAPH_SELECT,
|
||||
{ limit: 513, sort: false },
|
||||
);
|
||||
expect(getMessageTextStats).not.toHaveBeenCalled();
|
||||
expect(createTokenCounter).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('returns null before tokenization when the branch is too long', async () => {
|
||||
const { createTokenCounter } = jest.requireMock('@librechat/agents');
|
||||
const messages = Array.from({ length: 257 }, (_, index) => ({
|
||||
messageId: `message-${index}`,
|
||||
parentMessageId: index === 0 ? null : `message-${index - 1}`,
|
||||
isCreatedByUser: true,
|
||||
text: 'hello',
|
||||
}));
|
||||
const getMessages = jest.fn(async () => messages);
|
||||
const getMessageTextStats = jest.fn();
|
||||
|
||||
const result = await resolveContextProjection(
|
||||
{ userId: 'user-1', getMessages, getMessageTextStats },
|
||||
{ ...baseParams, messageId: 'message-256' },
|
||||
);
|
||||
|
||||
expect(result).toBeNull();
|
||||
expect(getMessages).toHaveBeenCalledTimes(1);
|
||||
expect(getMessageTextStats).not.toHaveBeenCalled();
|
||||
expect(createTokenCounter).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('returns null before loading bodies when the branch text is too large', async () => {
|
||||
const { createTokenCounter } = jest.requireMock('@librechat/agents');
|
||||
const getMessages = jest.fn(async () => [
|
||||
{
|
||||
messageId: 'message-1',
|
||||
parentMessageId: null,
|
||||
},
|
||||
]);
|
||||
const getMessageTextStats = jest.fn(async () => [textStats('message-1', 512 * 1024 + 1)]);
|
||||
const result = await resolveContextProjection(
|
||||
{
|
||||
userId: 'user-1',
|
||||
getMessages,
|
||||
getMessageTextStats,
|
||||
},
|
||||
baseParams,
|
||||
);
|
||||
|
||||
expect(result).toBeNull();
|
||||
expect(getMessages).toHaveBeenCalledTimes(1);
|
||||
expect(getMessageTextStats).toHaveBeenCalledWith(
|
||||
{
|
||||
conversationId: 'conversation-1',
|
||||
user: 'user-1',
|
||||
messageId: { $in: ['message-1'] },
|
||||
},
|
||||
{ limit: 1 },
|
||||
);
|
||||
expect(createTokenCounter).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it('loads only branch message bodies after resolving the graph', async () => {
|
||||
const graph = [
|
||||
{ messageId: 'message-1', parentMessageId: null },
|
||||
{ messageId: 'message-2', parentMessageId: 'message-1' },
|
||||
{ messageId: 'off-branch', parentMessageId: null },
|
||||
];
|
||||
const bodies = [
|
||||
{
|
||||
messageId: 'message-1',
|
||||
parentMessageId: null,
|
||||
isCreatedByUser: true,
|
||||
text: 'first',
|
||||
tokenCount: 5,
|
||||
},
|
||||
{
|
||||
messageId: 'message-2',
|
||||
parentMessageId: 'message-1',
|
||||
isCreatedByUser: false,
|
||||
text: 'second',
|
||||
tokenCount: 6,
|
||||
},
|
||||
];
|
||||
const getMessages = jest.fn(async (_filter: object, select?: string) =>
|
||||
select === GRAPH_SELECT ? graph : bodies,
|
||||
);
|
||||
const getMessageTextStats = jest.fn(async () => [
|
||||
textStats('message-1', 5),
|
||||
textStats('message-2', 6),
|
||||
]);
|
||||
|
||||
const result = await resolveContextProjection(
|
||||
{ userId: 'user-1', getMessages, getMessageTextStats },
|
||||
{ ...baseParams, messageId: 'message-2' },
|
||||
);
|
||||
|
||||
expect(result).toEqual({ tokenCount: 1, maxContextTokens: 1000 });
|
||||
expect(getMessages).toHaveBeenNthCalledWith(
|
||||
1,
|
||||
{ conversationId: 'conversation-1', user: 'user-1' },
|
||||
GRAPH_SELECT,
|
||||
{ limit: 513, sort: false },
|
||||
);
|
||||
expect(getMessageTextStats).toHaveBeenCalledWith(
|
||||
{
|
||||
conversationId: 'conversation-1',
|
||||
user: 'user-1',
|
||||
messageId: { $in: ['message-1', 'message-2'] },
|
||||
},
|
||||
{ limit: 2 },
|
||||
);
|
||||
expect(getMessages).toHaveBeenNthCalledWith(
|
||||
2,
|
||||
{
|
||||
conversationId: 'conversation-1',
|
||||
user: 'user-1',
|
||||
messageId: { $in: ['message-1', 'message-2'] },
|
||||
},
|
||||
BODY_SELECT,
|
||||
{ limit: 2, sort: false },
|
||||
);
|
||||
});
|
||||
|
||||
it('returns null before loading bodies when a branch message has too many quotes', async () => {
|
||||
const { createTokenCounter } = jest.requireMock('@librechat/agents');
|
||||
const getMessages = jest.fn(async () => [
|
||||
{
|
||||
messageId: 'message-1',
|
||||
parentMessageId: null,
|
||||
},
|
||||
]);
|
||||
const getMessageTextStats = jest.fn(async () => [
|
||||
{
|
||||
...textStats('message-1'),
|
||||
quoteCount: QUOTE_MAX_COUNT + 1,
|
||||
quoteBytes: 10,
|
||||
quoteLineCount: QUOTE_MAX_COUNT + 1,
|
||||
},
|
||||
]);
|
||||
|
||||
const result = await resolveContextProjection(
|
||||
{ userId: 'user-1', getMessages, getMessageTextStats },
|
||||
baseParams,
|
||||
);
|
||||
|
||||
expect(result).toBeNull();
|
||||
expect(getMessages).toHaveBeenCalledTimes(1);
|
||||
expect(getMessageTextStats).toHaveBeenCalledTimes(1);
|
||||
expect(createTokenCounter).not.toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
|
@ -2,7 +2,13 @@ import { HumanMessage, AIMessage } from '@langchain/core/messages';
|
|||
import { Providers, createTokenCounter, projectAgentContextUsage } from '@librechat/agents';
|
||||
import type { TContextProjectionRequest, TContextUsageEvent } from 'librechat-data-provider';
|
||||
import type { BaseMessage } from '@langchain/core/messages';
|
||||
import { mergeQuotedText } from '~/utils/quotes';
|
||||
import { QUOTE_MAX_COUNT, mergeQuotedText } from '~/utils/quotes';
|
||||
|
||||
const MAX_PROJECTION_MESSAGES = 512;
|
||||
const MAX_PROJECTION_BRANCH_MESSAGES = 256;
|
||||
const MAX_PROJECTION_BRANCH_TEXT_BYTES = 512 * 1024;
|
||||
const PROJECTION_GRAPH_SELECT = 'messageId parentMessageId metadata.summaryUsedTokens';
|
||||
const PROJECTION_BODY_SELECT = 'messageId parentMessageId tokenCount isCreatedByUser text quotes';
|
||||
|
||||
interface ProjectionMessage {
|
||||
messageId: string;
|
||||
|
|
@ -18,13 +24,42 @@ interface ProjectionMessage {
|
|||
metadata?: { summaryUsedTokens?: number };
|
||||
}
|
||||
|
||||
interface ProjectionMessageFilter {
|
||||
conversationId: string;
|
||||
user?: string;
|
||||
messageId?: string | { $in: string[] };
|
||||
}
|
||||
|
||||
interface ProjectionMessageQueryOptions {
|
||||
limit?: number;
|
||||
sort?: false;
|
||||
}
|
||||
|
||||
interface ProjectionMessageTextStats {
|
||||
messageId: string;
|
||||
textBytes: number;
|
||||
quoteCount: number;
|
||||
quoteBytes: number;
|
||||
quoteLineCount: number;
|
||||
nonStringQuoteCount: number;
|
||||
}
|
||||
|
||||
interface ProjectionMessageTextStatsOptions {
|
||||
limit?: number;
|
||||
}
|
||||
|
||||
export interface ContextProjectionDeps {
|
||||
/** Authenticated requester — branch lookups are scoped to this user. */
|
||||
userId?: string;
|
||||
getMessages: (
|
||||
filter: { conversationId: string; user?: string },
|
||||
filter: ProjectionMessageFilter,
|
||||
select?: string,
|
||||
options?: ProjectionMessageQueryOptions,
|
||||
) => Promise<ProjectionMessage[]>;
|
||||
getMessageTextStats: (
|
||||
filter: ProjectionMessageFilter,
|
||||
options?: ProjectionMessageTextStatsOptions,
|
||||
) => Promise<ProjectionMessageTextStats[]>;
|
||||
}
|
||||
|
||||
/**
|
||||
|
|
@ -51,6 +86,83 @@ function resolveBranch(messages: ProjectionMessage[], tailId: string): Projectio
|
|||
return branch.reverse();
|
||||
}
|
||||
|
||||
function hasValidProjectionIds(params: TContextProjectionRequest): boolean {
|
||||
return typeof params.conversationId === 'string' && typeof params.messageId === 'string';
|
||||
}
|
||||
|
||||
function getProjectionText(message: ProjectionMessage): string | null {
|
||||
const hasQuotes =
|
||||
message.isCreatedByUser === true && Array.isArray(message.quotes) && message.quotes.length > 0;
|
||||
if (!hasQuotes) {
|
||||
return message.text ?? '';
|
||||
}
|
||||
if (message.quotes == null || message.quotes.length > QUOTE_MAX_COUNT) {
|
||||
return null;
|
||||
}
|
||||
for (const quote of message.quotes) {
|
||||
if (typeof quote !== 'string') {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
return mergeQuotedText(message.text ?? '', message.quotes);
|
||||
}
|
||||
|
||||
function hasExceededBranchTextLimit(branch: ProjectionMessage[]): boolean {
|
||||
let bytes = 0;
|
||||
for (const message of branch) {
|
||||
const text = getProjectionText(message);
|
||||
if (text == null) {
|
||||
return true;
|
||||
}
|
||||
bytes += Buffer.byteLength(text, 'utf8');
|
||||
if (bytes > MAX_PROJECTION_BRANCH_TEXT_BYTES) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
function getEstimatedMergedTextBytes(stats: ProjectionMessageTextStats): number | null {
|
||||
if (
|
||||
stats.nonStringQuoteCount > 0 ||
|
||||
stats.quoteCount > QUOTE_MAX_COUNT ||
|
||||
stats.quoteLineCount < stats.quoteCount
|
||||
) {
|
||||
return null;
|
||||
}
|
||||
if (stats.quoteCount === 0) {
|
||||
return stats.textBytes;
|
||||
}
|
||||
|
||||
const quotePrefixBytes = stats.quoteLineCount * 2;
|
||||
const quoteLineBreakBytes = stats.quoteLineCount - stats.quoteCount;
|
||||
const quoteSeparatorBytes = (stats.quoteCount - 1) * 2;
|
||||
const bodySeparatorBytes = stats.textBytes > 0 ? 2 : 0;
|
||||
return (
|
||||
stats.textBytes +
|
||||
stats.quoteBytes +
|
||||
quotePrefixBytes +
|
||||
quoteLineBreakBytes +
|
||||
quoteSeparatorBytes +
|
||||
bodySeparatorBytes
|
||||
);
|
||||
}
|
||||
|
||||
function hasExceededBranchTextStatsLimit(stats: ProjectionMessageTextStats[]): boolean {
|
||||
let bytes = 0;
|
||||
for (const messageStats of stats) {
|
||||
const messageBytes = getEstimatedMergedTextBytes(messageStats);
|
||||
if (messageBytes == null) {
|
||||
return true;
|
||||
}
|
||||
bytes += messageBytes;
|
||||
if (bytes > MAX_PROJECTION_BRANCH_TEXT_BYTES) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
/** Maps an endpoint/provider string to the agents `Providers` enum. */
|
||||
function resolveProvider(value?: string): Providers {
|
||||
if (value == null || value === '') {
|
||||
|
|
@ -74,6 +186,43 @@ function resolveProvider(value?: string): Providers {
|
|||
return Providers.OPENAI;
|
||||
}
|
||||
|
||||
async function getBranchMessages(
|
||||
deps: ContextProjectionDeps,
|
||||
baseFilter: ProjectionMessageFilter,
|
||||
branch: ProjectionMessage[],
|
||||
): Promise<ProjectionMessage[] | null> {
|
||||
const branchIds = branch.map((message) => message.messageId);
|
||||
const stats = await deps.getMessageTextStats(
|
||||
{ ...baseFilter, messageId: { $in: branchIds } },
|
||||
{ limit: branchIds.length },
|
||||
);
|
||||
if (stats.length !== branchIds.length || hasExceededBranchTextStatsLimit(stats)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const stored = await deps.getMessages(
|
||||
{ ...baseFilter, messageId: { $in: branchIds } },
|
||||
PROJECTION_BODY_SELECT,
|
||||
{ limit: branchIds.length, sort: false },
|
||||
);
|
||||
if (stored.length !== branchIds.length) {
|
||||
return null;
|
||||
}
|
||||
const byId = new Map<string, ProjectionMessage>();
|
||||
for (const message of stored) {
|
||||
byId.set(message.messageId, message);
|
||||
}
|
||||
const ordered: ProjectionMessage[] = [];
|
||||
for (const messageId of branchIds) {
|
||||
const message = byId.get(messageId);
|
||||
if (message == null) {
|
||||
return null;
|
||||
}
|
||||
ordered.push(message);
|
||||
}
|
||||
return ordered;
|
||||
}
|
||||
|
||||
/**
|
||||
* Server-side context-usage projection: reconstructs the viewed branch and asks
|
||||
* the agents SDK what the next call's context would be, WITHOUT invoking the
|
||||
|
|
@ -90,19 +239,31 @@ export async function resolveContextProjection(
|
|||
deps: ContextProjectionDeps,
|
||||
params: TContextProjectionRequest,
|
||||
): Promise<TContextUsageEvent | null> {
|
||||
if (!hasValidProjectionIds(params)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const maxContextTokens = params.maxContextTokens;
|
||||
if (maxContextTokens == null || maxContextTokens <= 0) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const stored = await deps.getMessages(
|
||||
{ conversationId: params.conversationId, user: deps.userId },
|
||||
'messageId parentMessageId tokenCount isCreatedByUser text quotes metadata',
|
||||
);
|
||||
const baseFilter = { conversationId: params.conversationId, user: deps.userId };
|
||||
const stored = await deps.getMessages(baseFilter, PROJECTION_GRAPH_SELECT, {
|
||||
limit: MAX_PROJECTION_MESSAGES + 1,
|
||||
sort: false,
|
||||
});
|
||||
if (stored.length > MAX_PROJECTION_MESSAGES) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const branch = resolveBranch(stored, params.messageId);
|
||||
if (branch.length === 0) {
|
||||
return null;
|
||||
}
|
||||
if (branch.length > MAX_PROJECTION_BRANCH_MESSAGES) {
|
||||
return null;
|
||||
}
|
||||
|
||||
/** A summarized/compacted branch's next call sends the saved summary + the
|
||||
* post-summary tail, NOT this raw parent chain — projecting from the full
|
||||
|
|
@ -114,23 +275,29 @@ export async function resolveContextProjection(
|
|||
return null;
|
||||
}
|
||||
|
||||
const bodyBranch = await getBranchMessages(deps, baseFilter, branch);
|
||||
if (bodyBranch == null || hasExceededBranchTextLimit(bodyBranch)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const model = params.model;
|
||||
const encoding = (model ?? '').toLowerCase().includes('claude') ? 'claude' : 'o200k_base';
|
||||
const tokenCounter = await createTokenCounter(encoding);
|
||||
|
||||
const messages: BaseMessage[] = [];
|
||||
const indexTokenCountMap: Record<string, number> = {};
|
||||
for (let i = 0; i < branch.length; i++) {
|
||||
const message = branch[i];
|
||||
for (let i = 0; i < bodyBranch.length; i++) {
|
||||
const message = bodyBranch[i];
|
||||
/** Mirror the live path: prepend quoted excerpts into the user text the model
|
||||
* receives so the gauge counts the same prompt. */
|
||||
const hasQuotes =
|
||||
message.isCreatedByUser === true &&
|
||||
Array.isArray(message.quotes) &&
|
||||
message.quotes.length > 0;
|
||||
const text = hasQuotes
|
||||
? mergeQuotedText(message.text ?? '', message.quotes ?? [])
|
||||
: (message.text ?? '');
|
||||
const text = getProjectionText(message);
|
||||
if (text == null) {
|
||||
return null;
|
||||
}
|
||||
const lcMessage =
|
||||
message.isCreatedByUser === true ? new HumanMessage(text) : new AIMessage(text);
|
||||
messages.push(lcMessage);
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue