From 2163f592b7469fccab0ae8aca3068c0b6c3ef3e4 Mon Sep 17 00:00:00 2001 From: Sandipan kundu <69718573+skundu42@users.noreply.github.com> Date: Sat, 11 Jul 2026 07:24:56 +0530 Subject: [PATCH] fix(tools,validation): persist tool-call turns in conversation memory & tidy search thresholds (#1211) Co-authored-by: Dhravya Shah Co-authored-by: Claude Fable 5 --- apps/docs/integrations/ai-sdk.mdx | 12 + packages/tools/src/vercel/index.ts | 9 + packages/tools/src/vercel/middleware.ts | 140 +++++-- .../conversation-conversion.test.ts | 358 ++++++++++++++++++ packages/validation/api.test.ts | 82 ++++ packages/validation/api.ts | 9 +- 6 files changed, 576 insertions(+), 34 deletions(-) create mode 100644 packages/tools/test/with-supermemory/conversation-conversion.test.ts create mode 100644 packages/validation/api.test.ts diff --git a/apps/docs/integrations/ai-sdk.mdx b/apps/docs/integrations/ai-sdk.mdx index bdf73fd3..eede8429 100644 --- a/apps/docs/integrations/ai-sdk.mdx +++ b/apps/docs/integrations/ai-sdk.mdx @@ -140,6 +140,18 @@ const model = withSupermemory(openai("gpt-5"), { }) ``` +### Persisting Tool Calls (default: off) + +By default, saved conversations include only user and assistant text — tool calls and tool results are dropped, since tool payloads are often large and low-signal and would pollute memory extraction. To persist the full tool round trip (tool calls with their arguments, plus tool results, in their original order), set `includeToolCalls: true`: + +```typescript +const model = withSupermemory(openai("gpt-5"), { + containerTag: "user-123", + customId: "conv-1", + includeToolCalls: true, +}) +``` + --- ## Memory Tools diff --git a/packages/tools/src/vercel/index.ts b/packages/tools/src/vercel/index.ts index 3dbdcccd..7726ba2e 100644 --- a/packages/tools/src/vercel/index.ts +++ b/packages/tools/src/vercel/index.ts @@ -38,6 +38,12 @@ interface WrapVercelLanguageModelOptions { apiKey?: string /** Custom Supermemory API base URL */ baseUrl?: string + /** + * Persist assistant tool calls and tool results as part of the saved + * conversation. Off by default: tool payloads are often large and + * low-signal, and would pollute memory extraction. + */ + includeToolCalls?: boolean /** * Custom function to format memory data into the system prompt. * If not provided, uses the default "User Supermemories:" format. @@ -134,6 +140,7 @@ const wrapVercelLanguageModel = ( mode: options.mode ?? "profile", addMemory: options.addMemory ?? "always", baseUrl: options.baseUrl, + includeToolCalls: options.includeToolCalls ?? false, promptTemplate: options.promptTemplate, memoryRetrievalTimeoutMs: DEFAULT_MEMORY_RETRIEVAL_TIMEOUT_MS, }) @@ -193,6 +200,7 @@ const wrapVercelLanguageModel = ( ctx.logger, ctx.apiKey, ctx.normalizedBaseUrl, + ctx.includeToolCalls, ) } @@ -268,6 +276,7 @@ const wrapVercelLanguageModel = ( ctx.logger, ctx.apiKey, ctx.normalizedBaseUrl, + ctx.includeToolCalls, ) } }, diff --git a/packages/tools/src/vercel/middleware.ts b/packages/tools/src/vercel/middleware.ts index 144bd7df..ac1227ab 100644 --- a/packages/tools/src/vercel/middleware.ts +++ b/packages/tools/src/vercel/middleware.ts @@ -1,6 +1,7 @@ import Supermemory from "supermemory" import { addConversation, + type ContentPart, type ConversationMessage, } from "../conversations-client" import { @@ -15,16 +16,57 @@ import { import { type LanguageModelCallOptions, getLastUserMessage } from "./util" import { extractQueryText, injectMemoriesIntoParams } from "./memory-prompt" -const convertToConversationMessages = ( +const safeJsonStringify = (value: unknown): string => { + try { + return JSON.stringify(value) ?? "" + } catch { + return "" + } +} + +const serializeToolOutput = (output: unknown): string => { + if (typeof output === "string") return output + if (typeof output !== "object" || output === null) { + return safeJsonStringify(output) + } + + const wrapper = output as { + type?: unknown + value?: unknown + reason?: unknown + } + + if ( + (wrapper.type === "text" || wrapper.type === "error-text") && + typeof wrapper.value === "string" + ) { + return wrapper.value + } + if ( + wrapper.type === "json" || + wrapper.type === "error-json" || + wrapper.type === "content" + ) { + return safeJsonStringify(wrapper.value) + } + if (wrapper.type === "execution-denied") { + return typeof wrapper.reason === "string" && wrapper.reason + ? wrapper.reason + : "Tool execution denied" + } + + return safeJsonStringify(output) +} + +export const convertToConversationMessages = ( params: LanguageModelCallOptions, assistantResponseText: string, + includeToolCalls = false, ): ConversationMessage[] => { const messages: ConversationMessage[] = [] for (const msg of params.prompt) { - if (msg.role === "system") { - continue - } + if (msg.role === "system") continue if (typeof msg.content === "string") { if (msg.content) { @@ -33,36 +75,67 @@ const convertToConversationMessages = ( content: msg.content, }) } - } else { - const contentParts = msg.content - .map((c) => { - if (c.type === "text" && c.text) { - return { - type: "text" as const, - text: c.text, - } - } - if ( - c.type === "file" && - typeof c.data === "string" && - c.mediaType.startsWith("image/") - ) { - return { - type: "image_url" as const, - image_url: { url: c.data }, - } - } - return null - }) - .filter((part) => part !== null) + continue + } - if (contentParts.length > 0) { + let contentParts: ContentPart[] = [] + let toolCalls: NonNullable = [] + + // Flush any pending assistant/user content accumulated so far. Called + // before each tool-result so a tool result never jumps ahead of the + // text/tool-calls that preceded it (or behind text that follows it), + // preserving the original chronology for memory extraction. + const flushContent = () => { + if (contentParts.length > 0 || toolCalls.length > 0) { messages.push({ role: msg.role as "user" | "assistant" | "tool", - content: contentParts, + content: contentParts.length > 0 ? contentParts : "", + ...(toolCalls.length > 0 ? { tool_calls: toolCalls } : {}), + }) + contentParts = [] + toolCalls = [] + } + } + + for (const content of msg.content) { + if (content.type === "text" && content.text) { + contentParts.push({ + type: "text", + text: content.text, + }) + } else if ( + content.type === "file" && + typeof content.data === "string" && + content.mediaType.startsWith("image/") + ) { + contentParts.push({ + type: "image_url", + image_url: { url: content.data }, + }) + } else if ( + includeToolCalls && + content.type === "tool-call" && + msg.role === "assistant" + ) { + toolCalls.push({ + id: content.toolCallId, + type: "function", + function: { + name: content.toolName, + arguments: safeJsonStringify(content.input) || "{}", + }, + }) + } else if (includeToolCalls && content.type === "tool-result") { + flushContent() + messages.push({ + role: "tool", + content: serializeToolOutput(content.output), + tool_call_id: content.toolCallId, }) } } + + flushContent() } if (assistantResponseText) { @@ -84,11 +157,13 @@ export const saveMemoryAfterResponse = async ( logger: Logger, apiKey: string, baseUrl: string, + includeToolCalls = false, ): Promise => { try { const conversationMessages = convertToConversationMessages( params, assistantResponseText, + includeToolCalls, ) const response = await addConversation({ @@ -139,6 +214,12 @@ interface SupermemoryMiddlewareOptions { addMemory?: "always" | "never" /** Custom Supermemory API base URL */ baseUrl?: string + /** + * Persist assistant tool calls and tool results as part of the saved + * conversation. Off by default: tool payloads are often large and + * low-signal, and would pollute memory extraction. + */ + includeToolCalls?: boolean /** Custom function to format memory data into the system prompt */ promptTemplate?: PromptTemplate /** Max wait (ms) for the pre-LLM `/v4/profile` retrieval. Omit for no limit (e.g. tests). `withSupermemory` sets this internally. */ @@ -152,6 +233,7 @@ interface SupermemoryMiddlewareContext { customId: string mode: MemoryMode addMemory: "always" | "never" + includeToolCalls: boolean normalizedBaseUrl: string apiKey: string promptTemplate?: PromptTemplate @@ -174,6 +256,7 @@ export const createSupermemoryContext = ( mode = "profile", addMemory = "always", baseUrl, + includeToolCalls = false, promptTemplate, memoryRetrievalTimeoutMs, } = options @@ -195,6 +278,7 @@ export const createSupermemoryContext = ( customId, mode, addMemory, + includeToolCalls, normalizedBaseUrl, apiKey, promptTemplate, diff --git a/packages/tools/test/with-supermemory/conversation-conversion.test.ts b/packages/tools/test/with-supermemory/conversation-conversion.test.ts new file mode 100644 index 00000000..8497e334 --- /dev/null +++ b/packages/tools/test/with-supermemory/conversation-conversion.test.ts @@ -0,0 +1,358 @@ +import type Supermemory from "supermemory" +import type { + LanguageModelV2CallOptions, + LanguageModelV2Message, +} from "@ai-sdk/provider" +import { afterEach, describe, expect, it } from "vitest" +import { createLogger } from "../../src/shared" +import { saveMemoryAfterResponse } from "../../src/vercel/middleware" + +const originalFetch = globalThis.fetch + +const persistMessages = async ( + params: LanguageModelV2CallOptions, + assistantResponseText: string, + includeToolCalls = true, +) => { + let messages: unknown[] | undefined + + const fetchStub: typeof fetch = Object.assign( + async (input: RequestInfo | URL, init?: RequestInit) => { + const url = typeof input === "string" ? input : input.toString() + expect(url).toContain("/v4/conversations") + + const body = typeof init?.body === "string" ? init.body : "" + messages = (JSON.parse(body) as { messages: unknown[] }).messages + + return new Response( + JSON.stringify({ + id: "document-id", + conversationId: "conversation-id", + status: "done", + }), + { status: 200 }, + ) + }, + { preconnect: originalFetch.preconnect }, + ) + + globalThis.fetch = fetchStub + + await saveMemoryAfterResponse( + {} as Supermemory, + "user-id", + "conversation-id", + assistantResponseText, + params, + createLogger(false), + "test-api-key", + "https://api.example.com", + includeToolCalls, + ) + + return messages +} + +afterEach(() => { + globalThis.fetch = originalFetch +}) + +describe("convertToConversationMessages", () => { + it("preserves a tool-call and tool-result round trip", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "user", + content: [{ type: "text", text: "Search my memories" }], + }, + { + role: "assistant", + content: [ + { + type: "tool-call", + toolCallId: "call-1", + toolName: "search", + input: { query: "project" }, + }, + ], + } as unknown as LanguageModelV2Message, + { + role: "tool", + content: [ + { + type: "tool-result", + toolCallId: "call-1", + toolName: "search", + output: { + type: "json", + value: { memory: "Project memory" }, + }, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "Found it")).toEqual([ + { + role: "user", + content: [{ type: "text", text: "Search my memories" }], + }, + { + role: "assistant", + content: "", + tool_calls: [ + { + id: "call-1", + type: "function", + function: { + name: "search", + arguments: '{"query":"project"}', + }, + }, + ], + }, + { + role: "tool", + content: '{"memory":"Project memory"}', + tool_call_id: "call-1", + }, + { role: "assistant", content: "Found it" }, + ]) + }) + + it("keeps assistant text alongside tool calls", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "assistant", + content: [ + { type: "text", text: "I will search." }, + { + type: "tool-call", + toolCallId: "call-2", + toolName: "search", + input: {}, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "")).toEqual([ + { + role: "assistant", + content: [{ type: "text", text: "I will search." }], + tool_calls: [ + { + id: "call-2", + type: "function", + function: { name: "search", arguments: "{}" }, + }, + ], + }, + ]) + }) + + it("serializes a tool call with no input as empty JSON object", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "assistant", + content: [ + { + type: "tool-call", + toolCallId: "call-4", + toolName: "now", + input: undefined, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "")).toEqual([ + { + role: "assistant", + content: "", + tool_calls: [ + { + id: "call-4", + type: "function", + function: { name: "now", arguments: "{}" }, + }, + ], + }, + ]) + }) + + it("does not abort the save when tool-call input is not JSON-serializable", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "user", + content: [{ type: "text", text: "hi" }], + }, + { + role: "assistant", + content: [ + { + type: "tool-call", + toolCallId: "call-5", + toolName: "search", + input: { cursor: BigInt(1) }, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "done")).toEqual([ + { + role: "user", + content: [{ type: "text", text: "hi" }], + }, + { + role: "assistant", + content: "", + tool_calls: [ + { + id: "call-5", + type: "function", + function: { name: "search", arguments: "{}" }, + }, + ], + }, + { role: "assistant", content: "done" }, + ]) + }) + + it("preserves order when tool result is followed by assistant text in one message", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "assistant", + content: [ + { + type: "tool-call", + toolCallId: "call-6", + toolName: "search", + input: { query: "project" }, + }, + { + type: "tool-result", + toolCallId: "call-6", + toolName: "search", + output: { + type: "json", + value: { memory: "Project memory" }, + }, + }, + { type: "text", text: "Here is what I found." }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "")).toEqual([ + { + role: "assistant", + content: "", + tool_calls: [ + { + id: "call-6", + type: "function", + function: { + name: "search", + arguments: '{"query":"project"}', + }, + }, + ], + }, + { + role: "tool", + content: '{"memory":"Project memory"}', + tool_call_id: "call-6", + }, + { + role: "assistant", + content: [{ type: "text", text: "Here is what I found." }], + }, + ]) + }) + + it("unwraps text tool output", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "tool", + content: [ + { + type: "tool-result", + toolCallId: "call-3", + toolName: "search", + output: { + type: "text", + value: "No memories found", + }, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "")).toEqual([ + { + role: "tool", + content: "No memories found", + tool_call_id: "call-3", + }, + ]) + }) + + it("drops tool calls and tool results by default", async () => { + const params: LanguageModelV2CallOptions = { + prompt: [ + { + role: "user", + content: [{ type: "text", text: "Search my memories" }], + }, + { + role: "assistant", + content: [ + { type: "text", text: "I will search." }, + { + type: "tool-call", + toolCallId: "call-7", + toolName: "search", + input: { query: "project" }, + }, + ], + } as unknown as LanguageModelV2Message, + { + role: "tool", + content: [ + { + type: "tool-result", + toolCallId: "call-7", + toolName: "search", + output: { type: "json", value: { memory: "Project memory" } }, + }, + ], + } as unknown as LanguageModelV2Message, + ], + } + + expect(await persistMessages(params, "Found it", false)).toEqual([ + { + role: "user", + content: [{ type: "text", text: "Search my memories" }], + }, + { + role: "assistant", + content: [{ type: "text", text: "I will search." }], + }, + { role: "assistant", content: "Found it" }, + ]) + }) +}) diff --git a/packages/validation/api.test.ts b/packages/validation/api.test.ts new file mode 100644 index 00000000..05c18934 --- /dev/null +++ b/packages/validation/api.test.ts @@ -0,0 +1,82 @@ +import { describe, expect, it } from "bun:test" +import { readFileSync } from "node:fs" +import { SearchRequestSchema, Searchv4RequestSchema } from "./api" + +describe("search threshold schemas", () => { + it("do not contain redundant number transforms or unreachable range guards", () => { + const source = readFileSync(new URL("./api.ts", import.meta.url), "utf8") + const searchSchemas = source.slice( + source.indexOf("export const SearchRequestSchema"), + source.indexOf("export const SearchResultSchema"), + ) + + expect(searchSchemas).not.toContain(".transform(Number)") + expect(searchSchemas).not.toContain("v === undefined || (v >= 0 && v <= 1)") + }) + + it("preserves threshold defaults", () => { + const search = SearchRequestSchema.parse({ q: "memory" }) + const searchV4 = Searchv4RequestSchema.parse({ q: "memory" }) + + expect(search.chunkThreshold).toBe(0) + expect(search.documentThreshold).toBe(0) + expect(searchV4.threshold).toBe(0.6) + }) + + it.each([0, 0.5, 1])("accepts inclusive threshold value %p", (threshold) => { + expect( + SearchRequestSchema.parse({ + q: "memory", + chunkThreshold: threshold, + documentThreshold: threshold, + }), + ).toMatchObject({ + chunkThreshold: threshold, + documentThreshold: threshold, + }) + expect( + Searchv4RequestSchema.parse({ q: "memory", threshold }).threshold, + ).toBe(threshold) + }) + + it.each([ + -0.1, 1.1, + ])("rejects out-of-range threshold value %p", (threshold) => { + expect( + SearchRequestSchema.safeParse({ + q: "memory", + chunkThreshold: threshold, + }).success, + ).toBe(false) + expect( + SearchRequestSchema.safeParse({ + q: "memory", + documentThreshold: threshold, + }).success, + ).toBe(false) + expect( + Searchv4RequestSchema.safeParse({ q: "memory", threshold }).success, + ).toBe(false) + }) + + it("does not coerce threshold strings", () => { + expect( + SearchRequestSchema.safeParse({ + q: "memory", + chunkThreshold: "0.5", + }).success, + ).toBe(false) + expect( + SearchRequestSchema.safeParse({ + q: "memory", + documentThreshold: "0.5", + }).success, + ).toBe(false) + expect( + Searchv4RequestSchema.safeParse({ + q: "memory", + threshold: "0.5", + }).success, + ).toBe(false) + }) +}) diff --git a/packages/validation/api.ts b/packages/validation/api.ts index 23ccff2a..8579ca40 100644 --- a/packages/validation/api.ts +++ b/packages/validation/api.ts @@ -346,14 +346,13 @@ export const SearchRequestSchema = z.object({ .number() .optional() .default(0) - .refine((v) => v === undefined || (v >= 0 && v <= 1), { + .refine((v) => v >= 0 && v <= 1, { message: "chunkThreshold must be between 0 and 1", params: { max: 1, min: 0, }, }) - .transform(Number) .openapi({ description: "Threshold / sensitivity for chunk selection. 0 is least sensitive (returns most chunks, more results), 1 is most sensitive (returns lesser chunks, accurate results)", @@ -378,14 +377,13 @@ export const SearchRequestSchema = z.object({ .number() .optional() .default(0) - .refine((v) => v === undefined || (v >= 0 && v <= 1), { + .refine((v) => v >= 0 && v <= 1, { message: "documentThreshold must be between 0 and 1", params: { max: 1, min: 0, }, }) - .transform(Number) .openapi({ description: "Threshold / sensitivity for document selection. 0 is least sensitive (returns most documents, more results), 1 is most sensitive (returns lesser documents, accurate results)", @@ -473,14 +471,13 @@ export const Searchv4RequestSchema = z.object({ .number() .optional() .default(0.6) - .refine((v) => v === undefined || (v >= 0 && v <= 1), { + .refine((v) => v >= 0 && v <= 1, { message: "documentThreshold must be between 0 and 1", params: { max: 1, min: 0, }, }) - .transform(Number) .openapi({ description: "Threshold / sensitivity for memories selection. 0 is least sensitive (returns most memories, more results), 1 is most sensitive (returns lesser memories, accurate results)",