🪣 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:
Danny Avila 2026-06-23 08:43:09 -04:00 committed by GitHub
parent 2f800c5b52
commit 77854decdf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
8 changed files with 606 additions and 18 deletions

View 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();
});
});

View file

@ -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);