refactor: add entity ID and search param validation, rename textLemmatized field, update tests (#4843)
This commit is contained in:
@@ -82,6 +82,58 @@ function rejectTopLevelEntityParams(
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Validates and normalizes an entity ID.
|
||||
* - Trims leading/trailing whitespace
|
||||
* - Rejects empty or whitespace-only strings
|
||||
* - Rejects strings containing internal whitespace
|
||||
* @returns The trimmed entity ID, or undefined if input is undefined
|
||||
* @throws Error if entity ID is invalid
|
||||
*/
|
||||
function validateAndTrimEntityId(
|
||||
value: string | undefined,
|
||||
name: string,
|
||||
): string | undefined {
|
||||
if (value === undefined) return undefined;
|
||||
const trimmed = value.trim();
|
||||
if (trimmed === "") {
|
||||
throw new Error(
|
||||
`Invalid ${name}: cannot be empty or whitespace-only. Provide a valid identifier.`,
|
||||
);
|
||||
}
|
||||
if (/\s/.test(trimmed)) {
|
||||
throw new Error(
|
||||
`Invalid ${name}: cannot contain whitespace. Provide a valid identifier without spaces.`,
|
||||
);
|
||||
}
|
||||
return trimmed;
|
||||
}
|
||||
|
||||
/**
|
||||
* Validates search parameters.
|
||||
* @throws Error if threshold or topK are invalid
|
||||
*/
|
||||
function validateSearchParams(threshold?: number, topK?: number): void {
|
||||
if (threshold !== undefined) {
|
||||
if (typeof threshold !== "number" || isNaN(threshold)) {
|
||||
throw new Error("threshold must be a valid number");
|
||||
}
|
||||
if (threshold < 0 || threshold > 1) {
|
||||
throw new Error(
|
||||
`Invalid threshold: ${threshold}. Must be between 0 and 1 (inclusive).`,
|
||||
);
|
||||
}
|
||||
}
|
||||
if (topK !== undefined) {
|
||||
if (typeof topK !== "number" || isNaN(topK) || !Number.isInteger(topK)) {
|
||||
throw new Error("topK must be a valid integer");
|
||||
}
|
||||
if (topK < 0) {
|
||||
throw new Error(`Invalid topK: ${topK}. Must be a non-negative integer.`);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
export class Memory {
|
||||
private config: MemoryConfig;
|
||||
private customInstructions: string | undefined;
|
||||
@@ -276,6 +328,13 @@ export class Memory {
|
||||
messages: string | Message[],
|
||||
config: AddMemoryOptions,
|
||||
): Promise<SearchResult> {
|
||||
// Validate messages input
|
||||
if (messages === undefined || messages === null) {
|
||||
throw new Error(
|
||||
"messages is required and cannot be undefined or null. Provide a string or array of messages.",
|
||||
);
|
||||
}
|
||||
|
||||
await this._ensureInitialized();
|
||||
await this._captureEvent("add", {
|
||||
message_count: Array.isArray(messages) ? messages.length : 1,
|
||||
@@ -283,14 +342,12 @@ export class Memory {
|
||||
has_filters: !!config.filters,
|
||||
infer: config.infer,
|
||||
});
|
||||
const {
|
||||
userId,
|
||||
agentId,
|
||||
runId,
|
||||
metadata = {},
|
||||
filters = {},
|
||||
infer = true,
|
||||
} = config;
|
||||
const { metadata = {}, filters = {}, infer = true } = config;
|
||||
|
||||
// Validate and trim entity IDs
|
||||
const userId = validateAndTrimEntityId(config.userId, "userId");
|
||||
const agentId = validateAndTrimEntityId(config.agentId, "agentId");
|
||||
const runId = validateAndTrimEntityId(config.runId, "runId");
|
||||
|
||||
// Convert camelCase entity params to snake_case for storage (matches API and search/getAll filters)
|
||||
if (userId) filters.user_id = metadata.user_id = userId;
|
||||
@@ -804,14 +861,32 @@ export class Memory {
|
||||
// Reject top-level entity params - must use filters instead
|
||||
rejectTopLevelEntityParams(config as Record<string, any>, "search");
|
||||
|
||||
// Validate search parameters (before applying defaults)
|
||||
validateSearchParams(config.threshold, config.topK);
|
||||
|
||||
// Validate and trim entity IDs in filters
|
||||
const normalizedFilters = config.filters
|
||||
? {
|
||||
...config.filters,
|
||||
user_id: validateAndTrimEntityId(config.filters.user_id, "user_id"),
|
||||
agent_id: validateAndTrimEntityId(
|
||||
config.filters.agent_id,
|
||||
"agent_id",
|
||||
),
|
||||
run_id: validateAndTrimEntityId(config.filters.run_id, "run_id"),
|
||||
}
|
||||
: {};
|
||||
|
||||
await this._ensureInitialized();
|
||||
const { topK = 20, threshold = 0.1 } = config;
|
||||
|
||||
await this._captureEvent("search", {
|
||||
query_length: query.length,
|
||||
topK: config.topK,
|
||||
topK,
|
||||
has_filters: !!config.filters,
|
||||
});
|
||||
const { topK = 100, threshold = 0.1 } = config;
|
||||
let effectiveFilters: Record<string, any> = { ...(config.filters || {}) };
|
||||
|
||||
let effectiveFilters: Record<string, any> = { ...normalizedFilters };
|
||||
|
||||
// Apply enhanced metadata filtering if advanced operators are detected
|
||||
if (this._hasAdvancedOperators(effectiveFilters)) {
|
||||
@@ -1113,11 +1188,23 @@ export class Memory {
|
||||
// Reject top-level entity params - must use filters instead
|
||||
rejectTopLevelEntityParams(config as Record<string, any>, "getAll");
|
||||
|
||||
// Validate topK if provided (before applying defaults)
|
||||
validateSearchParams(undefined, config.topK);
|
||||
|
||||
await this._ensureInitialized();
|
||||
const { topK = 100, filters = {} } = config;
|
||||
|
||||
const { topK = 20 } = config;
|
||||
|
||||
// Validate and trim entity IDs in filters
|
||||
const filters = {
|
||||
...(config.filters || {}),
|
||||
user_id: validateAndTrimEntityId(config.filters?.user_id, "user_id"),
|
||||
agent_id: validateAndTrimEntityId(config.filters?.agent_id, "agent_id"),
|
||||
run_id: validateAndTrimEntityId(config.filters?.run_id, "run_id"),
|
||||
};
|
||||
|
||||
await this._captureEvent("get_all", {
|
||||
topK: topK,
|
||||
topK,
|
||||
has_user_id: !!filters.user_id,
|
||||
has_agent_id: !!filters.agent_id,
|
||||
has_run_id: !!filters.run_id,
|
||||
|
||||
@@ -260,7 +260,7 @@ export class MemoryVectorStore implements VectorStore {
|
||||
};
|
||||
|
||||
if (this.filterVector(memoryVector, filters)) {
|
||||
const text = payload.text_lemmatized || payload.data || "";
|
||||
const text = payload.textLemmatized || payload.data || "";
|
||||
candidates.push({ id: row.id, payload, tokens: this.tokenize(text) });
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,9 +186,9 @@ export class PGVector implements VectorStore {
|
||||
: "";
|
||||
|
||||
const searchQuery = `
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'text_lemmatized'), plainto_tsquery('simple', $1)) AS score, payload
|
||||
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'textLemmatized'), plainto_tsquery('simple', $1)) AS score, payload
|
||||
FROM ${this.collectionName}
|
||||
WHERE to_tsvector('simple', payload->>'text_lemmatized') @@ plainto_tsquery('simple', $1)
|
||||
WHERE to_tsvector('simple', payload->>'textLemmatized') @@ plainto_tsquery('simple', $1)
|
||||
${filterClause}
|
||||
ORDER BY score DESC
|
||||
LIMIT $2
|
||||
|
||||
@@ -469,7 +469,7 @@ describe("Memory – auto-initialization", () => {
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "sk-fake", model: "gpt-4-turbo-preview" },
|
||||
config: { apiKey: "sk-fake", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
disableHistory: true,
|
||||
|
||||
@@ -75,7 +75,7 @@ function createMemory(overrides: Partial<MemoryConfig> = {}): Memory {
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
...overrides,
|
||||
|
||||
@@ -75,7 +75,7 @@ function createMemory(): Memory {
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
});
|
||||
|
||||
@@ -69,7 +69,7 @@ function createMemory(overrides: Partial<MemoryConfig> = {}): Memory {
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-4-turbo-preview" },
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
historyDbPath: ":memory:",
|
||||
...overrides,
|
||||
@@ -94,7 +94,7 @@ describe("Memory - Initialization", () => {
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-4" },
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
};
|
||||
const mem = Memory.fromConfig(config);
|
||||
|
||||
@@ -0,0 +1,299 @@
|
||||
/**
|
||||
* Unit tests for OSS SDK input validation.
|
||||
*
|
||||
* Validates fixes for:
|
||||
* - Undefined/null message handling in add()
|
||||
* - Threshold bounds validation (must be 0-1) in search()
|
||||
* - TopK validation (must be non-negative) in search() and getAll()
|
||||
* - Whitespace-only entity ID rejection in add(), search(), getAll()
|
||||
*/
|
||||
/// <reference types="jest" />
|
||||
import { Memory } from "../src/memory";
|
||||
|
||||
jest.setTimeout(15000);
|
||||
|
||||
// Mock Google modules to prevent @google/genai crash in CI
|
||||
jest.mock("../src/embeddings/google", () => ({
|
||||
GoogleEmbedder: jest.fn(),
|
||||
}));
|
||||
jest.mock("../src/llms/google", () => ({
|
||||
GoogleLLM: jest.fn(),
|
||||
}));
|
||||
|
||||
jest.mock("../src/llms/openai", () => ({
|
||||
OpenAILLM: jest.fn().mockImplementation(() => ({
|
||||
generateResponse: jest.fn().mockResolvedValue(
|
||||
JSON.stringify({
|
||||
memory: [{ id: "0", text: "test memory", attributed_to: "user" }],
|
||||
}),
|
||||
),
|
||||
})),
|
||||
}));
|
||||
|
||||
const mockEmbedding = new Array(1536).fill(0.1);
|
||||
jest.mock("../src/embeddings/openai", () => ({
|
||||
OpenAIEmbedder: jest.fn().mockImplementation(() => ({
|
||||
embed: jest.fn().mockResolvedValue(mockEmbedding),
|
||||
embedBatch: jest.fn().mockResolvedValue([mockEmbedding]),
|
||||
})),
|
||||
}));
|
||||
|
||||
describe("Memory Input Validation", () => {
|
||||
let memory: Memory;
|
||||
const testUserId = "test-user-validation";
|
||||
|
||||
beforeAll(async () => {
|
||||
memory = new Memory({
|
||||
version: "v1.1",
|
||||
embedder: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "text-embedding-3-small" },
|
||||
},
|
||||
llm: {
|
||||
provider: "openai",
|
||||
config: { apiKey: "test-key", model: "gpt-5-mini" },
|
||||
},
|
||||
vectorStore: {
|
||||
provider: "memory",
|
||||
config: { collectionName: "validation-test" },
|
||||
},
|
||||
});
|
||||
// Wait for initialization
|
||||
await new Promise((resolve) => setTimeout(resolve, 2000));
|
||||
});
|
||||
|
||||
afterAll(async () => {
|
||||
try {
|
||||
await memory.reset();
|
||||
} catch (e) {
|
||||
// ignore cleanup errors
|
||||
}
|
||||
});
|
||||
|
||||
describe("add() message validation", () => {
|
||||
it("should throw error when messages is undefined", async () => {
|
||||
await expect(
|
||||
// @ts-ignore - intentionally passing undefined
|
||||
memory.add(undefined, { userId: testUserId }),
|
||||
).rejects.toThrow("messages is required");
|
||||
});
|
||||
|
||||
it("should throw error when messages is null", async () => {
|
||||
await expect(
|
||||
// @ts-ignore - intentionally passing null
|
||||
memory.add(null, { userId: testUserId }),
|
||||
).rejects.toThrow("messages is required");
|
||||
});
|
||||
});
|
||||
|
||||
describe("search() threshold validation", () => {
|
||||
it("should throw error when threshold > 1.0", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: 1.5,
|
||||
}),
|
||||
).rejects.toThrow("Invalid threshold");
|
||||
});
|
||||
|
||||
it("should throw error when threshold = 1.1", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: 1.1,
|
||||
}),
|
||||
).rejects.toThrow("Invalid threshold");
|
||||
});
|
||||
|
||||
it("should throw error when threshold is negative", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: -0.5,
|
||||
}),
|
||||
).rejects.toThrow("Invalid threshold");
|
||||
});
|
||||
|
||||
it("should throw error when threshold = -0.1", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: -0.1,
|
||||
}),
|
||||
).rejects.toThrow("Invalid threshold");
|
||||
});
|
||||
|
||||
it("should accept threshold = 0 (edge case)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: 0,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
|
||||
it("should accept threshold = 1.0 (edge case)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: 1.0,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
|
||||
it("should accept threshold = 0.5 (normal valid value)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
threshold: 0.5,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("search() topK validation", () => {
|
||||
it("should throw error when topK is negative", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
topK: -5,
|
||||
}),
|
||||
).rejects.toThrow("Invalid topK");
|
||||
});
|
||||
|
||||
it("should throw error when topK = -1", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
topK: -1,
|
||||
}),
|
||||
).rejects.toThrow("Invalid topK");
|
||||
});
|
||||
|
||||
it("should accept topK = 0 (returns empty)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
topK: 0,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
|
||||
it("should accept topK = 20 (normal value)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: testUserId },
|
||||
topK: 20,
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("add() entity ID validation", () => {
|
||||
it("should throw error when userId is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { userId: " " }),
|
||||
).rejects.toThrow("Invalid userId");
|
||||
});
|
||||
|
||||
it("should throw error when userId is tabs and newlines", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { userId: "\t\n\t" }),
|
||||
).rejects.toThrow("Invalid userId");
|
||||
});
|
||||
|
||||
it("should throw error when agentId is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { agentId: " " }),
|
||||
).rejects.toThrow("Invalid agentId");
|
||||
});
|
||||
|
||||
it("should throw error when runId is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { runId: " " }),
|
||||
).rejects.toThrow("Invalid runId");
|
||||
});
|
||||
|
||||
it("should throw error when userId contains internal whitespace", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { userId: "user 123" }),
|
||||
).rejects.toThrow("Invalid userId: cannot contain whitespace");
|
||||
});
|
||||
|
||||
it("should throw error when userId contains tab character", async () => {
|
||||
await expect(
|
||||
memory.add("test message", { userId: "user\t123" }),
|
||||
).rejects.toThrow("Invalid userId: cannot contain whitespace");
|
||||
});
|
||||
|
||||
it("should accept userId with leading/trailing whitespace (trimmed)", async () => {
|
||||
// Should not throw - leading/trailing whitespace is trimmed
|
||||
const result = await memory.add("test message", {
|
||||
userId: " valid-user ",
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("search() filter entity ID validation", () => {
|
||||
it("should throw error when user_id in filters is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: " " },
|
||||
}),
|
||||
).rejects.toThrow("Invalid user_id");
|
||||
});
|
||||
|
||||
it("should throw error when agent_id in filters is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { agent_id: " " },
|
||||
}),
|
||||
).rejects.toThrow("Invalid agent_id");
|
||||
});
|
||||
|
||||
it("should throw error when user_id contains internal whitespace", async () => {
|
||||
await expect(
|
||||
memory.search("test query", {
|
||||
filters: { user_id: "user 123" },
|
||||
}),
|
||||
).rejects.toThrow("Invalid user_id: cannot contain whitespace");
|
||||
});
|
||||
|
||||
it("should accept user_id with leading/trailing whitespace (trimmed)", async () => {
|
||||
const result = await memory.search("test query", {
|
||||
filters: { user_id: " valid-user " },
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
});
|
||||
|
||||
describe("getAll() validation", () => {
|
||||
it("should throw error when user_id is whitespace-only", async () => {
|
||||
await expect(
|
||||
memory.getAll({ filters: { user_id: " " } }),
|
||||
).rejects.toThrow("Invalid user_id");
|
||||
});
|
||||
|
||||
it("should throw error when topK is negative", async () => {
|
||||
await expect(
|
||||
memory.getAll({ filters: { user_id: testUserId }, topK: -1 }),
|
||||
).rejects.toThrow("Invalid topK");
|
||||
});
|
||||
|
||||
it("should throw error when user_id contains internal whitespace", async () => {
|
||||
await expect(
|
||||
memory.getAll({ filters: { user_id: "user 123" } }),
|
||||
).rejects.toThrow("Invalid user_id: cannot contain whitespace");
|
||||
});
|
||||
|
||||
it("should accept user_id with leading/trailing whitespace (trimmed)", async () => {
|
||||
const result = await memory.getAll({
|
||||
filters: { user_id: " valid-user " },
|
||||
});
|
||||
expect(result).toBeDefined();
|
||||
expect(result.results).toBeDefined();
|
||||
});
|
||||
});
|
||||
});
|
||||
+150
-17
@@ -110,6 +110,64 @@ def _reject_top_level_entity_params(kwargs: Dict[str, Any], method_name: str) ->
|
||||
)
|
||||
|
||||
|
||||
def _validate_and_trim_entity_id(value: Optional[str], name: str) -> Optional[str]:
|
||||
"""
|
||||
Validates and normalizes an entity ID.
|
||||
- Trims leading/trailing whitespace
|
||||
- Rejects empty or whitespace-only strings
|
||||
- Rejects strings containing internal whitespace
|
||||
|
||||
Args:
|
||||
value: The entity ID value to validate
|
||||
name: The parameter name (for error messages)
|
||||
|
||||
Returns:
|
||||
The trimmed entity ID, or None if input is None
|
||||
|
||||
Raises:
|
||||
ValueError: If entity ID is invalid
|
||||
"""
|
||||
if value is None:
|
||||
return None
|
||||
trimmed = value.strip()
|
||||
if trimmed == "":
|
||||
raise ValueError(
|
||||
f"Invalid {name}: cannot be empty or whitespace-only. Provide a valid identifier."
|
||||
)
|
||||
if any(c.isspace() for c in trimmed):
|
||||
raise ValueError(
|
||||
f"Invalid {name}: cannot contain whitespace. Provide a valid identifier without spaces."
|
||||
)
|
||||
return trimmed
|
||||
|
||||
|
||||
def _validate_search_params(threshold: Optional[float] = None, top_k: Optional[int] = None) -> None:
|
||||
"""
|
||||
Validates search parameters.
|
||||
|
||||
Args:
|
||||
threshold: Similarity threshold (must be between 0 and 1)
|
||||
top_k: Number of results to return (must be non-negative integer)
|
||||
|
||||
Raises:
|
||||
ValueError: If threshold or top_k are invalid
|
||||
"""
|
||||
if threshold is not None:
|
||||
if not isinstance(threshold, (int, float)):
|
||||
raise ValueError("threshold must be a valid number")
|
||||
if threshold < 0 or threshold > 1:
|
||||
raise ValueError(
|
||||
f"Invalid threshold: {threshold}. Must be between 0 and 1 (inclusive)."
|
||||
)
|
||||
if top_k is not None:
|
||||
if not isinstance(top_k, int) or isinstance(top_k, bool):
|
||||
raise ValueError("top_k must be a valid integer")
|
||||
if top_k < 0:
|
||||
raise ValueError(
|
||||
f"Invalid top_k: {top_k}. Must be a non-negative integer."
|
||||
)
|
||||
|
||||
|
||||
def _is_sensitive_field(field_name: str) -> bool:
|
||||
"""Check if a field should be redacted for telemetry safety.
|
||||
|
||||
@@ -217,9 +275,14 @@ def _build_filters_and_metadata(
|
||||
base_metadata_template = deepcopy(input_metadata) if input_metadata else {}
|
||||
effective_query_filters = deepcopy(input_filters) if input_filters else {}
|
||||
|
||||
# ---------- add all provided session ids ----------
|
||||
# ---------- validate and add all provided session ids ----------
|
||||
session_ids_provided = []
|
||||
|
||||
# Validate and trim entity IDs
|
||||
user_id = _validate_and_trim_entity_id(user_id, "user_id")
|
||||
agent_id = _validate_and_trim_entity_id(agent_id, "agent_id")
|
||||
run_id = _validate_and_trim_entity_id(run_id, "run_id")
|
||||
|
||||
if user_id:
|
||||
base_metadata_template["user_id"] = user_id
|
||||
effective_query_filters["user_id"] = user_id
|
||||
@@ -876,7 +939,7 @@ class Memory(MemoryBase):
|
||||
self,
|
||||
*,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
top_k: int = 100,
|
||||
top_k: int = 20,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -886,20 +949,38 @@ class Memory(MemoryBase):
|
||||
filters (dict): Filter dict containing entity IDs and optional metadata filters.
|
||||
Must contain at least one of: user_id, agent_id, run_id.
|
||||
Example: filters={"user_id": "u1", "agent_id": "a1"}
|
||||
top_k (int, optional): The maximum number of memories to return. Defaults to 100.
|
||||
top_k (int, optional): The maximum number of memories to return. Defaults to 20.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing a list of memories under the "results" key.
|
||||
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", ...}]}`
|
||||
|
||||
Raises:
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
|
||||
or if top_k is invalid.
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
_reject_top_level_entity_params(kwargs, "get_all")
|
||||
|
||||
# Validate top_k
|
||||
_validate_search_params(top_k=top_k)
|
||||
|
||||
# Validate and trim entity IDs in filters
|
||||
effective_filters = dict(filters) if filters else {}
|
||||
if "user_id" in effective_filters:
|
||||
effective_filters["user_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["user_id"], "user_id"
|
||||
)
|
||||
if "agent_id" in effective_filters:
|
||||
effective_filters["agent_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["agent_id"], "agent_id"
|
||||
)
|
||||
if "run_id" in effective_filters:
|
||||
effective_filters["run_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["run_id"], "run_id"
|
||||
)
|
||||
|
||||
# Validate filters contains at least one entity ID
|
||||
effective_filters = filters or {}
|
||||
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
@@ -968,7 +1049,7 @@ class Memory(MemoryBase):
|
||||
self,
|
||||
query: str,
|
||||
*,
|
||||
top_k: int = 100,
|
||||
top_k: int = 20,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
threshold: float = 0.1,
|
||||
rerank: bool = False,
|
||||
@@ -979,7 +1060,7 @@ class Memory(MemoryBase):
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
top_k (int, optional): Maximum number of results to return. Defaults to 100.
|
||||
top_k (int, optional): Maximum number of results to return. Defaults to 20.
|
||||
filters (dict): Filter dict containing entity IDs and optional metadata filters.
|
||||
Must contain at least one of: user_id, agent_id, run_id.
|
||||
Example: filters={"user_id": "u1", "agent_id": "a1"}
|
||||
@@ -1008,13 +1089,29 @@ class Memory(MemoryBase):
|
||||
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", "score": 0.8, ...}]}`
|
||||
|
||||
Raises:
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
|
||||
or if threshold/top_k values are invalid.
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
_reject_top_level_entity_params(kwargs, "search")
|
||||
|
||||
# Validate filters contains at least one entity ID
|
||||
# Validate search parameters (before applying defaults)
|
||||
_validate_search_params(threshold=threshold, top_k=top_k)
|
||||
|
||||
# Validate and trim entity IDs in filters
|
||||
effective_filters = filters.copy() if filters else {}
|
||||
if "user_id" in effective_filters:
|
||||
effective_filters["user_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["user_id"], "user_id"
|
||||
)
|
||||
if "agent_id" in effective_filters:
|
||||
effective_filters["agent_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["agent_id"], "agent_id"
|
||||
)
|
||||
if "run_id" in effective_filters:
|
||||
effective_filters["run_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["run_id"], "run_id"
|
||||
)
|
||||
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
@@ -2126,7 +2223,7 @@ class AsyncMemory(MemoryBase):
|
||||
self,
|
||||
*,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
top_k: int = 100,
|
||||
top_k: int = 20,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -2136,20 +2233,38 @@ class AsyncMemory(MemoryBase):
|
||||
filters (dict): Filter dict containing entity IDs and optional metadata filters.
|
||||
Must contain at least one of: user_id, agent_id, run_id.
|
||||
Example: filters={"user_id": "u1", "agent_id": "a1"}
|
||||
top_k (int, optional): The maximum number of memories to return. Defaults to 100.
|
||||
top_k (int, optional): The maximum number of memories to return. Defaults to 20.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing a list of memories under the "results" key.
|
||||
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", ...}]}`
|
||||
|
||||
Raises:
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
|
||||
or if top_k is invalid.
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
_reject_top_level_entity_params(kwargs, "get_all")
|
||||
|
||||
# Validate top_k
|
||||
_validate_search_params(top_k=top_k)
|
||||
|
||||
# Validate and trim entity IDs in filters
|
||||
effective_filters = dict(filters) if filters else {}
|
||||
if "user_id" in effective_filters:
|
||||
effective_filters["user_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["user_id"], "user_id"
|
||||
)
|
||||
if "agent_id" in effective_filters:
|
||||
effective_filters["agent_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["agent_id"], "agent_id"
|
||||
)
|
||||
if "run_id" in effective_filters:
|
||||
effective_filters["run_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["run_id"], "run_id"
|
||||
)
|
||||
|
||||
# Validate filters contains at least one entity ID
|
||||
effective_filters = filters or {}
|
||||
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
@@ -2218,7 +2333,7 @@ class AsyncMemory(MemoryBase):
|
||||
self,
|
||||
query: str,
|
||||
*,
|
||||
top_k: int = 100,
|
||||
top_k: int = 20,
|
||||
filters: Optional[Dict[str, Any]] = None,
|
||||
threshold: float = 0.1,
|
||||
rerank: bool = False,
|
||||
@@ -2229,7 +2344,7 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
Args:
|
||||
query (str): Query to search for.
|
||||
top_k (int, optional): Maximum number of results to return. Defaults to 100.
|
||||
top_k (int, optional): Maximum number of results to return. Defaults to 20.
|
||||
filters (dict): Filter dict containing entity IDs and optional metadata filters.
|
||||
Must contain at least one of: user_id, agent_id, run_id.
|
||||
Example: filters={"user_id": "u1", "agent_id": "a1"}
|
||||
@@ -2258,13 +2373,31 @@ class AsyncMemory(MemoryBase):
|
||||
Example for v1.1+: `{"results": [{"id": "...", "memory": "...", "score": 0.8, ...}]}`
|
||||
|
||||
Raises:
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id.
|
||||
ValueError: If filters doesn't contain at least one of user_id, agent_id, run_id,
|
||||
or if threshold/top_k values are invalid.
|
||||
"""
|
||||
# Reject top-level entity params - must use filters instead
|
||||
_reject_top_level_entity_params(kwargs, "search")
|
||||
|
||||
# Validate filters contains at least one entity ID
|
||||
# Validate search parameters (before applying defaults)
|
||||
_validate_search_params(threshold=threshold, top_k=top_k)
|
||||
|
||||
# Validate and trim entity IDs in filters
|
||||
effective_filters = filters.copy() if filters else {}
|
||||
if "user_id" in effective_filters:
|
||||
effective_filters["user_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["user_id"], "user_id"
|
||||
)
|
||||
if "agent_id" in effective_filters:
|
||||
effective_filters["agent_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["agent_id"], "agent_id"
|
||||
)
|
||||
if "run_id" in effective_filters:
|
||||
effective_filters["run_id"] = _validate_and_trim_entity_id(
|
||||
effective_filters["run_id"], "run_id"
|
||||
)
|
||||
|
||||
# Validate filters contains at least one entity ID
|
||||
if not any(key in effective_filters for key in ("user_id", "agent_id", "run_id")):
|
||||
raise ValueError(
|
||||
"filters must contain at least one of: user_id, agent_id, run_id. "
|
||||
|
||||
+110
-3
@@ -111,9 +111,9 @@ def test_search(memory_instance):
|
||||
# Score is now combined score (semantic only since no BM25/entity), still 0.9
|
||||
assert result["results"][0]["score"] == pytest.approx(0.9)
|
||||
|
||||
# Hybrid pipeline over-fetches: max(100*4, 60) = 400
|
||||
# Hybrid pipeline over-fetches: max(20*4, 60) = 80 (top_k default is now 20)
|
||||
memory_instance.vector_store.search.assert_called_once_with(
|
||||
query="test query", vectors=[0.1, 0.2, 0.3], top_k=400, filters={"user_id": "test_user"}
|
||||
query="test query", vectors=[0.1, 0.2, 0.3], top_k=80, filters={"user_id": "test_user"}
|
||||
)
|
||||
|
||||
|
||||
@@ -200,7 +200,7 @@ def test_get_all(memory_instance):
|
||||
assert result["results"][0]["memory"] == "Memory 1"
|
||||
assert result["results"][0]["user_id"] == "test_user"
|
||||
|
||||
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=100)
|
||||
memory_instance.vector_store.list.assert_called_once_with(filters={"user_id": "test_user"}, top_k=20)
|
||||
|
||||
|
||||
def test_no_telemetry_vector_store_when_disabled():
|
||||
@@ -241,3 +241,110 @@ def test_telemetry_vector_store_created_when_enabled():
|
||||
|
||||
# VectorStoreFactory.create should be called twice — user data + telemetry
|
||||
assert mock_vector_store.create.call_count == 2
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# Input Validation Tests
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class TestEntityIdValidation:
|
||||
"""Tests for entity ID validation (whitespace rejection and trimming)."""
|
||||
|
||||
def test_search_rejects_whitespace_only_user_id(self, memory_instance):
|
||||
"""Search should reject whitespace-only user_id in filters."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
|
||||
memory_instance.search("test query", filters={"user_id": " "})
|
||||
|
||||
def test_search_rejects_internal_whitespace_user_id(self, memory_instance):
|
||||
"""Search should reject user_id with internal whitespace."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
|
||||
memory_instance.search("test query", filters={"user_id": "user 123"})
|
||||
|
||||
def test_search_rejects_tab_in_user_id(self, memory_instance):
|
||||
"""Search should reject user_id with tab character."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
|
||||
memory_instance.search("test query", filters={"user_id": "user\t123"})
|
||||
|
||||
def test_get_all_rejects_whitespace_only_user_id(self, memory_instance):
|
||||
"""get_all should reject whitespace-only user_id in filters."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
|
||||
memory_instance.get_all(filters={"user_id": " "})
|
||||
|
||||
def test_get_all_rejects_internal_whitespace_user_id(self, memory_instance):
|
||||
"""get_all should reject user_id with internal whitespace."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
|
||||
memory_instance.get_all(filters={"user_id": "user 123"})
|
||||
|
||||
def test_add_rejects_whitespace_only_user_id(self, memory_instance):
|
||||
"""add should reject whitespace-only user_id."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot be empty"):
|
||||
memory_instance.add("test message", user_id=" ")
|
||||
|
||||
def test_add_rejects_internal_whitespace_user_id(self, memory_instance):
|
||||
"""add should reject user_id with internal whitespace."""
|
||||
with pytest.raises(ValueError, match="Invalid user_id.*cannot contain whitespace"):
|
||||
memory_instance.add("test message", user_id="user 123")
|
||||
|
||||
|
||||
class TestSearchParamValidation:
|
||||
"""Tests for search parameter validation (threshold and top_k)."""
|
||||
|
||||
def test_search_rejects_threshold_above_1(self, memory_instance):
|
||||
"""Search should reject threshold > 1."""
|
||||
with pytest.raises(ValueError, match="Invalid threshold.*Must be between 0 and 1"):
|
||||
memory_instance.search("test query", filters={"user_id": "test"}, threshold=1.5)
|
||||
|
||||
def test_search_rejects_negative_threshold(self, memory_instance):
|
||||
"""Search should reject negative threshold."""
|
||||
with pytest.raises(ValueError, match="Invalid threshold.*Must be between 0 and 1"):
|
||||
memory_instance.search("test query", filters={"user_id": "test"}, threshold=-0.5)
|
||||
|
||||
def test_search_rejects_negative_top_k(self, memory_instance):
|
||||
"""Search should reject negative top_k."""
|
||||
with pytest.raises(ValueError, match="Invalid top_k.*Must be a non-negative"):
|
||||
memory_instance.search("test query", filters={"user_id": "test"}, top_k=-5)
|
||||
|
||||
def test_get_all_rejects_negative_top_k(self, memory_instance):
|
||||
"""get_all should reject negative top_k."""
|
||||
with pytest.raises(ValueError, match="Invalid top_k.*Must be a non-negative"):
|
||||
memory_instance.get_all(filters={"user_id": "test"}, top_k=-1)
|
||||
|
||||
def test_search_accepts_threshold_zero(self, memory_instance):
|
||||
"""Search should accept threshold=0 (edge case)."""
|
||||
mock_memories = []
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None)
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
result = memory_instance.search("test", filters={"user_id": "test"}, threshold=0)
|
||||
|
||||
assert "results" in result
|
||||
|
||||
def test_search_accepts_threshold_one(self, memory_instance):
|
||||
"""Search should accept threshold=1.0 (edge case)."""
|
||||
mock_memories = []
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None)
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
result = memory_instance.search("test", filters={"user_id": "test"}, threshold=1.0)
|
||||
|
||||
assert "results" in result
|
||||
|
||||
def test_search_accepts_top_k_zero(self, memory_instance):
|
||||
"""Search should accept top_k=0."""
|
||||
mock_memories = []
|
||||
memory_instance.vector_store.search = Mock(return_value=mock_memories)
|
||||
memory_instance.vector_store.keyword_search = Mock(return_value=None)
|
||||
memory_instance.embedding_model.embed = Mock(return_value=[0.1, 0.2, 0.3])
|
||||
|
||||
with patch("mem0.memory.main.lemmatize_for_bm25", return_value="test"), \
|
||||
patch("mem0.memory.main.extract_entities", return_value=[]):
|
||||
result = memory_instance.search("test", filters={"user_id": "test"}, top_k=0)
|
||||
|
||||
assert "results" in result
|
||||
|
||||
Reference in New Issue
Block a user