diff --git a/apps/mcp/src/client.ts b/apps/mcp/src/client.ts index 4e01a4796..ae02c9238 100644 --- a/apps/mcp/src/client.ts +++ b/apps/mcp/src/client.ts @@ -57,6 +57,27 @@ export interface ProfileResponse { searchResults?: SearchResult } +// Context passed to a governance hook alongside the retrieved memories. +export interface MemoryGovernanceContext { + containerTag?: string + query?: string +} + +export type SearchGovernanceHook = ( + result: SearchResult, + context: MemoryGovernanceContext, +) => SearchResult | Promise + +export type ProfileGovernanceHook = ( + result: ProfileResponse, + context: MemoryGovernanceContext, +) => ProfileResponse | Promise + +export interface MemoryGovernanceHooks { + onSearch?: SearchGovernanceHook + onProfile?: ProfileGovernanceHook +} + export interface Project { id: string name: string @@ -133,11 +154,13 @@ export class SupermemoryClient { private hasExplicitContainerTag: boolean private bearerToken: string private apiUrl: string + private governance?: MemoryGovernanceHooks constructor( bearerToken: string, containerTag?: string, apiUrl = "https://api.supermemory.ai", + governance?: MemoryGovernanceHooks, ) { this.bearerToken = bearerToken this.apiUrl = apiUrl @@ -148,6 +171,7 @@ export class SupermemoryClient { }) this.hasExplicitContainerTag = Boolean(containerTag) this.containerTag = containerTag || DEFAULT_PROJECT_ID + this.governance = governance } // Create memory using SDK @@ -290,11 +314,20 @@ export class SupermemoryClient { return { ...base, memory: text } }) - return { + const searchResult: SearchResult = { results, total: result.total, timing: result.timing, } + + if (this.governance?.onSearch) { + return await this.governance.onSearch(searchResult, { + containerTag, + query, + }) + } + + return searchResult } catch (error) { this.handleOperationError("Search request", error) } @@ -344,6 +377,13 @@ export class SupermemoryClient { } } + if (this.governance?.onProfile) { + return await this.governance.onProfile(response, { + containerTag: this.containerTag, + query, + }) + } + return response } catch (error) { this.handleOperationError("Profile request", error) diff --git a/packages/tools/src/mastra/processor.ts b/packages/tools/src/mastra/processor.ts index e7a39e96b..246fb22d3 100644 --- a/packages/tools/src/mastra/processor.ts +++ b/packages/tools/src/mastra/processor.ts @@ -22,6 +22,7 @@ import { type Logger, type MemoryMode, type PromptTemplate, + type MemoryGovernanceHook, } from "../shared" import { addConversation, @@ -51,6 +52,7 @@ interface ProcessorContext { logger: Logger promptTemplate?: PromptTemplate memoryCache: MemoryCache + governanceHook?: MemoryGovernanceHook } /** @@ -73,6 +75,7 @@ function createProcessorContext( logger, promptTemplate: options.promptTemplate, memoryCache: new MemoryCache(), + ...(options.governanceHook ? { governanceHook: options.governanceHook } : {}), } } @@ -181,6 +184,9 @@ export class SupermemoryInputProcessor implements Processor { apiKey: this.ctx.apiKey, logger: this.ctx.logger, promptTemplate: this.ctx.promptTemplate, + ...(this.ctx.governanceHook + ? { governanceHook: this.ctx.governanceHook } + : {}), }) if (memories) { diff --git a/packages/tools/src/mastra/types.ts b/packages/tools/src/mastra/types.ts index f1781a516..42cb56725 100644 --- a/packages/tools/src/mastra/types.ts +++ b/packages/tools/src/mastra/types.ts @@ -10,6 +10,7 @@ import type { MemoryMode, AddMemoryMode, MemoryPromptData, + MemoryGovernanceHook, } from "../shared" // Re-export Mastra core types for consumers @@ -51,6 +52,8 @@ export interface SupermemoryMastraOptions { verbose?: boolean /** Custom function to format memory data into the system prompt */ promptTemplate?: PromptTemplate + /** Governance hook invoked on raw retrieval results before dedup/formatting */ + governanceHook?: MemoryGovernanceHook } export type { PromptTemplate, MemoryMode, AddMemoryMode, MemoryPromptData } diff --git a/packages/tools/src/openai/middleware.ts b/packages/tools/src/openai/middleware.ts index c9b8b4b88..204ac3c0c 100644 --- a/packages/tools/src/openai/middleware.ts +++ b/packages/tools/src/openai/middleware.ts @@ -4,6 +4,7 @@ import { addConversation } from "../conversations-client" import { deduplicateMemoriesForMode } from "../tools-shared" import { createLogger, type Logger } from "../vercel/logger" import { convertProfileToMarkdown } from "../vercel/util" +import type { MemoryGovernanceHook } from "../shared" const normalizeBaseUrl = (url?: string): string => { const defaultUrl = "https://api.supermemory.ai" @@ -20,6 +21,8 @@ export interface OpenAIMiddlewareOptions { mode?: "profile" | "query" | "full" addMemory?: "always" | "never" baseUrl?: string + /** Governance hook invoked on raw retrieval results before dedup/formatting */ + governanceHook?: MemoryGovernanceHook } interface SupermemoryProfileSearch { @@ -161,17 +164,26 @@ const addSystemPrompt = async ( logger: Logger, mode: "profile" | "query" | "full", baseUrl: string, + governanceHook?: MemoryGovernanceHook, ) => { const systemPromptExists = messages.some((msg) => msg.role === "system") const queryText = mode !== "profile" ? getLastUserMessage(messages) : "" - const memoriesResponse = await supermemoryProfileSearch( + let memoriesResponse = await supermemoryProfileSearch( containerTag, queryText, baseUrl, ) + if (governanceHook) { + memoriesResponse = await governanceHook(memoriesResponse, { + containerTag, + queryText, + mode, + }) + } + const memoryCountStatic = memoriesResponse.profile.static?.length || 0 const memoryCountDynamic = memoriesResponse.profile.dynamic?.length || 0 @@ -429,6 +441,7 @@ export function createOpenAIMiddleware( const customId = options?.customId const mode = options?.mode ?? "profile" const addMemory = options?.addMemory ?? "always" + const governanceHook = options?.governanceHook const originalCreate = openaiClient.chat.completions.create const originalResponsesCreate = openaiClient.responses?.create @@ -453,12 +466,20 @@ export function createOpenAIMiddleware( mode: "profile" | "query" | "full", context: "chat" | "responses", ) => { - const memoriesResponse = await supermemoryProfileSearch( + let memoriesResponse = await supermemoryProfileSearch( containerTag, queryText, baseUrl, ) + if (governanceHook) { + memoriesResponse = await governanceHook(memoriesResponse, { + containerTag, + queryText, + mode, + }) + } + const memoryCountStatic = memoriesResponse.profile.static?.length || 0 const memoryCountDynamic = memoriesResponse.profile.dynamic?.length || 0 @@ -623,7 +644,14 @@ export function createOpenAIMiddleware( } operations.push( - addSystemPrompt(messages, containerTag, logger, mode, baseUrl), + addSystemPrompt( + messages, + containerTag, + logger, + mode, + baseUrl, + governanceHook, + ), ) const results = await Promise.all(operations) diff --git a/packages/tools/src/shared/index.ts b/packages/tools/src/shared/index.ts index 5a6e0f7ba..522ee14ed 100644 --- a/packages/tools/src/shared/index.ts +++ b/packages/tools/src/shared/index.ts @@ -8,6 +8,8 @@ export type { ProfileStructure, ProfileMarkdownData, SupermemoryBaseOptions, + MemoryGovernanceContext, + MemoryGovernanceHook, } from "./types" // Logger diff --git a/packages/tools/src/shared/memory-client.test.ts b/packages/tools/src/shared/memory-client.test.ts index 4b4edc0a1..fc01ab322 100644 --- a/packages/tools/src/shared/memory-client.test.ts +++ b/packages/tools/src/shared/memory-client.test.ts @@ -1,6 +1,7 @@ import { afterEach, describe, expect, it, vi } from "vitest" import { buildMemoriesText } from "./memory-client" import { createLogger } from "./logger" +import type { ProfileStructure } from "./types" const API_KEY = "sm_test_key" const BASE_URL = "https://api.supermemory.ai" @@ -75,4 +76,46 @@ describe("buildMemoriesText", () => { // Present once, under the profile — not duplicated into the search results. expect(memories.match(/User is allergic to peanuts/g)).toHaveLength(1) }) + + it("applies a governanceHook to the raw profile before formatting", async () => { + mockProfileResponse({ + profile: { + static: [{ memory: "User's SSN is 123-45-6789" }], + dynamic: [], + }, + searchResults: { results: [] }, + }) + + const governanceHook = vi.fn(async (profile: ProfileStructure) => ({ + ...profile, + profile: { + ...profile.profile, + static: profile.profile.static?.map((entry) => ({ + ...entry, + memory: entry.memory.replace(/\d{3}-\d{2}-\d{4}/, "[REDACTED]"), + })), + }, + })) + + const memories = await buildMemoriesText({ + containerTag: CONTAINER_TAG, + queryText: "", + mode: "profile", + baseUrl: BASE_URL, + apiKey: API_KEY, + logger, + governanceHook, + }) + + expect(governanceHook).toHaveBeenCalledWith( + expect.objectContaining({ + profile: expect.objectContaining({ + static: [{ memory: "User's SSN is 123-45-6789" }], + }), + }), + { containerTag: CONTAINER_TAG, queryText: "", mode: "profile" }, + ) + expect(memories).toContain("[REDACTED]") + expect(memories).not.toContain("123-45-6789") + }) }) diff --git a/packages/tools/src/shared/memory-client.ts b/packages/tools/src/shared/memory-client.ts index 9f2d73a7c..f6b732877 100644 --- a/packages/tools/src/shared/memory-client.ts +++ b/packages/tools/src/shared/memory-client.ts @@ -1,6 +1,7 @@ import { deduplicateMemoriesForMode } from "../tools-shared" import type { Logger, + MemoryGovernanceHook, MemoryMode, MemoryPromptData, ProfileStructure, @@ -76,6 +77,8 @@ export interface BuildMemoriesTextOptions { logger: Logger promptTemplate?: PromptTemplate signal?: AbortSignal + /** Governance hook invoked on raw retrieval results before dedup/formatting */ + governanceHook?: MemoryGovernanceHook } /** @@ -97,9 +100,10 @@ export const buildMemoriesText = async ( logger, promptTemplate = defaultPromptTemplate, signal, + governanceHook, } = options - const memoriesResponse = await supermemoryProfileSearch( + let memoriesResponse = await supermemoryProfileSearch( containerTag, queryText, baseUrl, @@ -107,6 +111,14 @@ export const buildMemoriesText = async ( signal, ) + if (governanceHook) { + memoriesResponse = await governanceHook(memoriesResponse, { + containerTag, + queryText, + mode, + }) + } + const memoryCountStatic = memoriesResponse.profile.static?.length || 0 const memoryCountDynamic = memoriesResponse.profile.dynamic?.length || 0 diff --git a/packages/tools/src/shared/types.ts b/packages/tools/src/shared/types.ts index 421785f52..6e0fbce91 100644 --- a/packages/tools/src/shared/types.ts +++ b/packages/tools/src/shared/types.ts @@ -123,4 +123,32 @@ export interface SupermemoryBaseOptions { verbose?: boolean /** Custom function to format memory data into the system prompt */ promptTemplate?: PromptTemplate + /** Governance hook invoked on raw retrieval results before formatting/injection */ + governanceHook?: MemoryGovernanceHook } + +/** + * Context passed to a governance hook alongside the retrieved memories. + */ +export interface MemoryGovernanceContext { + /** Container tag/user ID the retrieval was scoped to */ + containerTag: string + /** Query text used for the retrieval (empty string in "profile" mode) */ + queryText: string + /** Memory retrieval mode active for this call */ + mode: MemoryMode +} + +/** + * A hook invoked with the raw retrieval results before they are deduplicated, + * formatted, and injected into the LLM context. Lets a governance provider + * (PII redaction, prompt-injection detection, audit logging, etc.) inspect + * and/or rewrite `memory` strings, drop entries, or throw to abort retrieval. + * + * Runs at the retrieval boundary only — it does not scan content at ingestion + * time, and it is not implemented by Supermemory itself; providers plug in here. + */ +export type MemoryGovernanceHook = ( + profile: ProfileStructure, + context: MemoryGovernanceContext, +) => ProfileStructure | Promise diff --git a/packages/tools/src/vercel/middleware.ts b/packages/tools/src/vercel/middleware.ts index ac1227ab2..837f0ef06 100644 --- a/packages/tools/src/vercel/middleware.ts +++ b/packages/tools/src/vercel/middleware.ts @@ -12,6 +12,7 @@ import { type Logger, type PromptTemplate, type MemoryMode, + type MemoryGovernanceHook, } from "../shared" import { type LanguageModelCallOptions, getLastUserMessage } from "./util" import { extractQueryText, injectMemoriesIntoParams } from "./memory-prompt" @@ -224,6 +225,8 @@ interface SupermemoryMiddlewareOptions { promptTemplate?: PromptTemplate /** Max wait (ms) for the pre-LLM `/v4/profile` retrieval. Omit for no limit (e.g. tests). `withSupermemory` sets this internally. */ memoryRetrievalTimeoutMs?: number + /** Governance hook invoked on raw retrieval results before dedup/formatting */ + governanceHook?: MemoryGovernanceHook } interface SupermemoryMiddlewareContext { @@ -238,6 +241,7 @@ interface SupermemoryMiddlewareContext { apiKey: string promptTemplate?: PromptTemplate memoryRetrievalTimeoutMs?: number + governanceHook?: MemoryGovernanceHook /** * Per-turn memory cache. Stores the injected memories string for each * user turn (keyed by turnKey) to avoid redundant API calls during tool-call @@ -259,6 +263,7 @@ export const createSupermemoryContext = ( includeToolCalls = false, promptTemplate, memoryRetrievalTimeoutMs, + governanceHook, } = options const logger = createLogger(verbose) @@ -285,6 +290,7 @@ export const createSupermemoryContext = ( ...(memoryRetrievalTimeoutMs !== undefined ? { memoryRetrievalTimeoutMs } : {}), + ...(governanceHook ? { governanceHook } : {}), memoryCache: new MemoryCache(), } } @@ -368,6 +374,7 @@ export const transformParamsWithMemory = async ( logger: ctx.logger, promptTemplate: ctx.promptTemplate, ...(fetchSignal ? { signal: fetchSignal } : {}), + ...(ctx.governanceHook ? { governanceHook: ctx.governanceHook } : {}), }) } finally { if (timeoutId !== undefined) { diff --git a/packages/tools/src/voltagent/middleware.ts b/packages/tools/src/voltagent/middleware.ts index bf7717265..c958b4407 100644 --- a/packages/tools/src/voltagent/middleware.ts +++ b/packages/tools/src/voltagent/middleware.ts @@ -17,6 +17,7 @@ import { extractQueryText, type Logger, type MemoryMode, + type MemoryGovernanceHook, } from "../shared" import type { SupermemoryVoltAgent, VoltAgentMessage } from "./types" @@ -59,6 +60,7 @@ export interface SupermemoryMiddlewareContext { metadata?: Record searchMode?: "memories" | "documents" | "hybrid" entityContext?: string + governanceHook?: MemoryGovernanceHook } /** @@ -91,6 +93,7 @@ export const createSupermemoryContext = ( searchMode, entityContext, verbose = false, + governanceHook, } = options // Runtime validation: customId is required @@ -130,6 +133,7 @@ export const createSupermemoryContext = ( metadata, searchMode, entityContext, + ...(governanceHook ? { governanceHook } : {}), } } @@ -297,7 +301,25 @@ export const enhanceMessagesWithMemories = async ( chunk?: string metadata?: Record } - const formattedMemories = response.results + + let searchResults = response.results as SearchResult[] + if (ctx.governanceHook) { + const governed = await ctx.governanceHook( + { + profile: {}, + searchResults: { + results: searchResults.map((r) => ({ + memory: r.memory ?? r.chunk ?? "", + ...(r.metadata ? { metadata: r.metadata } : {}), + })), + }, + }, + { containerTag: ctx.containerTag, queryText, mode: ctx.mode }, + ) + searchResults = governed.searchResults.results + } + + const formattedMemories = searchResults .map((result: SearchResult) => { const text = result.memory || result.chunk return text ? `- ${text}` : null @@ -309,10 +331,10 @@ export const enhanceMessagesWithMemories = async ( ? ctx.promptTemplate({ userMemories: "", generalSearchMemories: formattedMemories, - searchResults: response.results as Array<{ - memory: string - metadata?: Record - }>, + searchResults: searchResults.map((r) => ({ + memory: r.memory || r.chunk || "", + ...(r.metadata ? { metadata: r.metadata } : {}), + })), }) : `The following are relevant memories and context about this user retrieved from previous interactions. Use these to personalize your response:\n\n${formattedMemories}` } else { @@ -324,6 +346,7 @@ export const enhanceMessagesWithMemories = async ( apiKey: ctx.apiKey, logger: ctx.logger, promptTemplate: ctx.promptTemplate, + ...(ctx.governanceHook ? { governanceHook: ctx.governanceHook } : {}), }) }