refactor: add entity ID and search param validation, rename textLemmatized field, update tests (#4843)

This commit is contained in:
Kartik
2026-04-15 20:57:09 +05:30
committed by GitHub
parent 9692726db4
commit e6d6276bb9
10 changed files with 667 additions and 41 deletions
+100 -13
View File
@@ -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,
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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,
+1 -1
View File
@@ -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:",
});
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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