From 5aceaa8ea3d11b80b36d28aafa6b0d276f2c1ead Mon Sep 17 00:00:00 2001 From: jiang Date: Thu, 3 Sep 2026 19:18:59 +0800 Subject: [PATCH] fix(memory): guard summary model context --- Memory/readme.md | 6 + Memory/src/model/llm.ts | 117 ++++++++++++++++-- Memory/tests/llm-json-retry.test.ts | 177 ++++++++++++++++++++++++++++ 3 files changed, 288 insertions(+), 12 deletions(-) diff --git a/Memory/readme.md b/Memory/readme.md index 7f2405745..624f19346 100644 --- a/Memory/readme.md +++ b/Memory/readme.md @@ -78,6 +78,12 @@ a conservative 7,500-token per-input budget. Set `MEMMY_EMBEDDING_MAX_INPUT_TOKENS` to use a smaller budget for a provider with a shorter context window. +All remote summary-model calls (capture, reflection, long-turn splitting, +reward scoring, retrieval filtering, and turn routing) are clipped before +requesting when their estimated input would exceed an 8,192-token context; +the requested output budget and a 512-token provider overhead margin are +reserved first. + When `storage.token`, `MEMMY_MEMORY_TOKEN`, or `MEMORY_SERVICE_TOKEN` is set, all HTTP routes except `GET /api/v1/health` require that token as a bearer token or `x-api-key`. diff --git a/Memory/src/model/llm.ts b/Memory/src/model/llm.ts index 1492f6f95..1eec75c89 100644 --- a/Memory/src/model/llm.ts +++ b/Memory/src/model/llm.ts @@ -1,4 +1,5 @@ -import type { LlmConfig } from "../config/index.js"; +import { get_encoding } from "tiktoken"; +import { MEMORY_SUMMARY_MAX_TOKENS, type LlmConfig } from "../config/index.js"; import { createMemoryLogger, memoryErrorFields } from "../logging/logger.js"; import { resolveMemoryAgentRegion } from "./agent-region.js"; import { bearer, postJsonWithRetry, trimTrailingSlash } from "./http.js"; @@ -62,12 +63,21 @@ const ANTHROPIC_THINKING_BUDGET_TOKENS = 4096; const ANTHROPIC_MIN_THINKING_OUTPUT_TOKENS = ANTHROPIC_THINKING_BUDGET_TOKENS + 4096; const GEMINI_THINKING_BUDGET_ENABLED = -1; const GEMINI_THINKING_BUDGET_DISABLED = 0; +const SUMMARY_CONTEXT_TOKEN_LIMIT = 8_192; +const SUMMARY_INPUT_TOKEN_BUDGET = 7_000; +const SUMMARY_CONTEXT_SAFETY_MARGIN = 512; +// The largest summary-role operation is span.big_turn; retries must leave +// room for its evidence rather than requesting the entire 8K window as output. +const SUMMARY_OUTPUT_TOKEN_LIMIT = 4_096; +const SUMMARY_INPUT_TRUNCATION_MARKER = "\n[... input truncated ...]\n"; interface ThinkingControl { enabled: boolean; fields: Record; } +let summaryEncoder: ReturnType | undefined; + export interface CreateLlmClientOptions { modelRole?: MemoryLlmModelRole; } @@ -105,12 +115,17 @@ class HttpLlmClient implements LlmClient { options: LlmCompletionOptions ): Promise { const startedAt = Date.now(); + const callOptions = { + ...options, + temperature: options.temperature ?? this.config.temperature, + maxTokens: this.outputTokenBudget(options.maxTokens) + }; const fields = { role: this.options.modelRole ?? "unspecified", operation: options.operation, provider: this.config.provider, model: this.config.model, - maxTokens: options.maxTokens ?? this.config.maxTokens, + maxTokens: callOptions.maxTokens, timeoutMs: options.timeoutMs ?? this.config.timeoutMs, maxRetries: options.maxRetries ?? this.config.maxRetries, jsonMode: options.jsonMode ?? false @@ -127,14 +142,12 @@ class HttpLlmClient implements LlmClient { logger.error("request.rejected", { ...fields, ...memoryErrorFields(error) }); throw error; } - const callOptions = { - ...options, - temperature: options.temperature ?? this.config.temperature, - maxTokens: options.maxTokens ?? this.config.maxTokens - }; logger.debug("request.started", fields); try { - const result = await this.completeOnce(messages, callOptions); + const requestMessages = this.options.modelRole === "memory_summary" + ? constrainSummaryMessages(messages, callOptions.maxTokens!, fields) + : messages; + const result = await this.completeOnce(requestMessages, callOptions); this.lastOkAt = new Date().toISOString(); this.lastError = undefined; logger.info("request.succeeded", { @@ -163,7 +176,7 @@ class HttpLlmClient implements LlmClient { let malformedRetriesRemaining = Math.max(0, this.config.malformedRetries); let lengthRetryUsed = false; let previousWasTruncated = false; - let maxTokens = options.maxTokens ?? this.config.maxTokens; + let maxTokens = this.outputTokenBudget(options.maxTokens); let jsonAttempt = 0; while (true) { @@ -186,7 +199,7 @@ class HttpLlmClient implements LlmClient { parseError !== undefined && looksLikeTruncatedJson(result.text, parseError) ); const expandedMaxTokens = truncated && !lengthRetryUsed - ? doubleMaxTokens(maxTokens) + ? doubleMaxTokens(maxTokens, this.options.modelRole === "memory_summary" ? SUMMARY_OUTPUT_TOKEN_LIMIT : undefined) : undefined; if (expandedMaxTokens !== undefined) { logger.warn("json.truncated_retry", { @@ -267,6 +280,13 @@ class HttpLlmClient implements LlmClient { }; } + private outputTokenBudget(requested: number | undefined): number | undefined { + const maxTokens = requested ?? this.config.maxTokens; + return this.options.modelRole === "memory_summary" + ? Math.min(maxTokens ?? MEMORY_SUMMARY_MAX_TOKENS, SUMMARY_OUTPUT_TOKEN_LIMIT) + : maxTokens; + } + private completeOnce(messages: LlmMessage[], options: Required> & LlmCompletionOptions): Promise { switch (this.config.provider) { case "openai_compatible": @@ -561,6 +581,78 @@ class HttpLlmClient implements LlmClient { } } +function constrainSummaryMessages( + messages: LlmMessage[], + outputTokens: number, + fields: Record +): LlmMessage[] { + summaryEncoder ??= get_encoding("cl100k_base"); + const encoded = messages.map((message) => summaryEncoder!.encode(message.content, [], [])); + const inputTokens = encoded.reduce((total, tokens) => total + tokens.length, 0); + // cl100k_base is an estimate for compatible providers. The margin also + // reserves space for their chat template and special tokens. + const budget = Math.min( + SUMMARY_INPUT_TOKEN_BUDGET, + SUMMARY_CONTEXT_TOKEN_LIMIT - outputTokens - SUMMARY_CONTEXT_SAFETY_MARGIN + ); + if (inputTokens <= budget) return messages; + + const systemTokens = encoded.reduce((total, tokens, index) => + total + (messages[index]!.role === "system" ? tokens.length : 0), 0); + const contentCount = messages.filter((message) => message.role !== "system").length; + const markerTokens = summaryEncoder.encode(SUMMARY_INPUT_TRUNCATION_MARKER, [], []).length + 2; + if (systemTokens + contentCount * markerTokens > budget) { + throw new Error("Summary system instructions exceed the input token budget"); + } + // Keep system instructions intact regardless of message order. Short + // messages keep their contents; long evidence messages share the remainder. + let remaining = budget - systemTokens; + let remainingCount = contentCount; + const result = [...messages]; + const contentIndices = messages.map((_message, index) => index) + .filter((index) => messages[index]!.role !== "system") + .sort((a, b) => encoded[a]!.length - encoded[b]!.length); + for (const index of contentIndices) { + const tokens = encoded[index]!; + const limit = Math.floor(remaining / remainingCount); + const content = tokens.length <= limit + ? messages[index]!.content + : truncateSummaryContent(tokens, limit, markerTokens); + remaining -= summaryEncoder.encode(content, [], []).length; + remainingCount -= 1; + result[index] = { ...messages[index]!, content }; + } + logger.warn("request.input_truncated", { + ...fields, + inputTokens, + inputTokenBudget: budget, + outputTokensReserved: outputTokens + }); + return result; +} + +function truncateSummaryContent(tokens: Uint32Array, budget: number, markerTokens: number): string { + let keep = Math.max(0, budget - markerTokens); + while (true) { + const head = Math.ceil(keep / 2); + const tail = keep - head; + const content = decodeSummaryTokens(tokens.slice(0, head)) + SUMMARY_INPUT_TRUNCATION_MARKER + + decodeSummaryTokens(tail > 0 ? tokens.slice(-tail) : tokens.slice(0, 0)); + const length = summaryEncoder!.encode(content, [], []).length; + if (length <= budget) return content; + keep = Math.max(0, keep - Math.max(1, length - budget)); + } +} + +function decodeSummaryTokens(tokens: Uint32Array): string { + const bytes = summaryEncoder!.decode(tokens); + let start = 0; + while (start < bytes.length && (bytes[start]! & 0xc0) === 0x80) start += 1; + // A token can contain only part of a UTF-8 character. Streaming decoding + // drops an incomplete trailing character without inserting U+FFFD. + return new TextDecoder("utf-8", { ignoreBOM: true }).decode(bytes.subarray(start), { stream: true }); +} + function openAiCompatibleThinkingControl(input: { vendor: string; endpoint: string; @@ -942,9 +1034,10 @@ function parseJsonObject(text: string): Record { return parsed as Record; } -function doubleMaxTokens(value: number | undefined): number | undefined { +function doubleMaxTokens(value: number | undefined, limit = Number.MAX_SAFE_INTEGER): number | undefined { if (value === undefined || !Number.isFinite(value) || value <= 0) return undefined; - return Math.min(Number.MAX_SAFE_INTEGER, Math.max(1, Math.floor(value)) * 2); + const expanded = Math.min(limit, Math.max(1, Math.floor(value)) * 2); + return expanded > value ? expanded : undefined; } function normalizeFinishReason(value: string | undefined): LlmCallResult["finishReason"] { diff --git a/Memory/tests/llm-json-retry.test.ts b/Memory/tests/llm-json-retry.test.ts index b85decc53..aebb960d6 100644 --- a/Memory/tests/llm-json-retry.test.ts +++ b/Memory/tests/llm-json-retry.test.ts @@ -1,8 +1,11 @@ import { afterEach, describe, expect, it, vi } from "vitest"; +import { get_encoding } from "tiktoken"; import type { LlmConfig } from "../src/config/index.js"; import { createLlmClient } from "../src/model/llm.js"; +import type { LlmMessage } from "../src/model/types.js"; const originalLogLevel = process.env.MEMMY_LOG_LEVEL; +const encoding = get_encoding("cl100k_base"); afterEach(() => { vi.unstubAllGlobals(); @@ -14,6 +17,175 @@ afterEach(() => { } }); +describe("memory summary input budget", () => { + it.each([ + ["capture.summarize", 512], + ["capture.reflection.synth", 500], + ["capture.alpha.reflection_score.v1", 700], + ["span.big_turn.v1", 4096], + ["reward.r_human.v1", 700], + ["retrieval.query_extract.v1", 320], + ["retrieval.filter.v1", 512], + ["relation.classify.v1", 512], + ["relation.arbitration.v1", 512] + ])("budgets the summary role for %s", async (operation, maxTokens) => { + const fetchMock = sequenceFetch([openAiResponse("ok")]); + vi.stubGlobal("fetch", fetchMock); + + const messages: LlmMessage[] = [ + { role: "system", content: "Keep the system instructions." }, + { role: "user", content: "BEGIN " + "tool output ".repeat(6_000) + "END" } + ]; + const original = structuredClone(messages); + const client = createLlmClient(llmConfig(), { modelRole: "memory_summary" }); + await client.complete(messages, { operation, maxTokens }); + + const body = requestBodies(fetchMock)[0]!; + const sent = body.messages as LlmMessage[]; + expect(body.max_tokens).toBe(maxTokens); + expect(sent[0]).toEqual(messages[0]); + expect(sent[1]?.content).toMatch(/^BEGIN /); + expect(sent[1]?.content).toContain("[... input truncated ...]"); + expect(sent[1]?.content).toMatch(/END$/); + expect(inputTokenCount(body)).toBeLessThanOrEqual(Math.min(7_000, 8_192 - maxTokens - 512)); + expect(messages).toEqual(original); + }); + + it("keeps short input unchanged and sends the default output reservation", async () => { + const fetchMock = sequenceFetch([openAiResponse("ok")]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig({ maxTokens: undefined }), { modelRole: "memory_summary" }); + const messages: LlmMessage[] = [{ role: "user", content: "Summarize this." }]; + + await client.complete(messages, { operation: "capture.summarize" }); + + expect(requestBodies(fetchMock)[0]).toMatchObject({ messages, max_tokens: 512 }); + }); + + it("preserves system messages after long evidence and shares space among other messages", async () => { + const fetchMock = sequenceFetch([openAiResponse("ok")]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig(), { modelRole: "memory_summary" }); + const messages: LlmMessage[] = [ + { role: "user", content: "user evidence ".repeat(6_000) }, + { role: "system", content: "Do not drop these instructions." }, + { role: "assistant", content: "assistant evidence ".repeat(6_000) }, + { role: "user", content: "Keep this short request." } + ]; + + await client.complete(messages, { operation: "capture.summarize" }); + + const body = requestBodies(fetchMock)[0]!; + const sent = body.messages as LlmMessage[]; + expect(sent.map((message) => message.role)).toEqual(messages.map((message) => message.role)); + expect(sent[1]).toEqual(messages[1]); + expect(sent[3]).toEqual(messages[3]); + expect(sent[0]?.content).toContain("[... input truncated ...]"); + expect(sent[2]?.content).toContain("[... input truncated ...]"); + expect(inputTokenCount(body) + Number(body.max_tokens) + 512).toBeLessThanOrEqual(8_192); + }); + + it("cuts Chinese and emoji only at valid Unicode boundaries", async () => { + const fetchMock = sequenceFetch([openAiResponse("ok")]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig(), { modelRole: "memory_summary" }); + const source = "开头" + "中文🧠检索𠮷野家🧪".repeat(2_000) + "结尾"; + + await client.complete([{ role: "user", content: source }], { operation: "span.big_turn.v1" }); + + const body = requestBodies(fetchMock)[0]!; + const content = (body.messages as LlmMessage[])[0]!.content; + const [head, tail] = content.split("\n[... input truncated ...]\n"); + expect(content).not.toContain("\ufffd"); + expect(head?.startsWith("开头")).toBe(true); + expect(tail?.endsWith("结尾")).toBe(true); + expect(source.startsWith(head!)).toBe(true); + expect(source.endsWith(tail!)).toBe(true); + expect(inputTokenCount(body)).toBeLessThanOrEqual(3_584); + }); + + it("rejects oversized system instructions without dropping them or making a request", async () => { + const fetchMock = sequenceFetch([]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig(), { modelRole: "memory_summary" }); + + await expect(client.complete([ + { role: "system", content: "system instruction ".repeat(6_000) }, + { role: "user", content: "evidence" } + ], { operation: "capture.summarize" })).rejects.toThrow("Summary system instructions exceed"); + + expect(fetchMock).not.toHaveBeenCalled(); + expect(client.status().lastError).toContain("Summary system instructions exceed"); + }); + + it("caps an oversized summary output reservation", async () => { + const fetchMock = sequenceFetch([openAiResponse("ok")]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig({ maxTokens: 8_192 }), { modelRole: "memory_summary" }); + + await client.complete([{ role: "user", content: "evidence ".repeat(10_000) }], { + operation: "span.big_turn.v1" + }); + + const body = requestBodies(fetchMock)[0]!; + expect(body.max_tokens).toBe(4_096); + expect(inputTokenCount(body) + Number(body.max_tokens) + 512).toBeLessThanOrEqual(8_192); + }); + + it("recomputes the input budget including JSON hints when output doubles", async () => { + const fetchMock = sequenceFetch([ + openAiResponse('{"ok":', "length"), + openAiResponse('{"ok":true}', "stop") + ]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig({ maxTokens: 512 }), { modelRole: "memory_summary" }); + + await expect(client.completeJson([ + { role: "system", content: "Keep the system instructions." }, + { role: "user", content: "evidence ".repeat(10_000) } + ], { operation: "retrieval.filter.v1" })).resolves.toEqual({ ok: true }); + + const bodies = requestBodies(fetchMock); + expect(bodies.map((body) => body.max_tokens)).toEqual([512, 1_024]); + for (const body of bodies) { + const messages = body.messages as LlmMessage[]; + expect(messages[0]?.content).toContain("Return exactly one valid JSON object."); + expect(messages[0]?.content).toContain("Keep the system instructions."); + expect(inputTokenCount(body) + Number(body.max_tokens) + 512).toBeLessThanOrEqual(8_192); + } + expect((bodies[1]!.messages as LlmMessage[])[0]?.content).toContain("previous output was truncated"); + }); + + it("does not expand a summary JSON retry to an impossible 8192-token output", async () => { + const fetchMock = sequenceFetch([openAiResponse('{"spans":[', "length")]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig(), { modelRole: "memory_summary" }); + + await expect(client.completeJson([{ role: "user", content: "segment this turn" }], { + operation: "span.big_turn.v1" + })).rejects.toThrow(); + + expect(fetchMock).toHaveBeenCalledTimes(1); + expect(requestBodies(fetchMock)[0]?.max_tokens).toBe(4_096); + }); + + it("does not constrain evolution requests even when they handle a summary fallback", async () => { + const fetchMock = sequenceFetch([ + openAiResponse('{"ok":', "length"), + openAiResponse('{"ok":true}', "stop") + ]); + vi.stubGlobal("fetch", fetchMock); + const client = createLlmClient(llmConfig(), { modelRole: "memory_evolution" }); + const content = "evidence ".repeat(10_000); + + await client.completeJson([{ role: "user", content }], { operation: "capture.summarize" }); + + const bodies = requestBodies(fetchMock); + expect(bodies.map((body) => body.max_tokens)).toEqual([4_096, 8_192]); + expect(bodies.every((body) => (body.messages as LlmMessage[])[1]?.content === content)).toBe(true); + }); +}); + describe("memory LLM JSON length retry", () => { it("merges existing system messages into one leading system message", async () => { const fetchMock = sequenceFetch([openAiResponse('{"ok":true}', "stop")]); @@ -242,3 +414,8 @@ function sequenceFetch(responses: Response[]): ReturnType>): Array> { return fetchMock.mock.calls.map(([, init]) => JSON.parse(String(init?.body)) as Record); } + +function inputTokenCount(body: Record): number { + return (body.messages as LlmMessage[]) + .reduce((total, message) => total + encoding.encode(message.content, [], []).length, 0); +}