fix(pgvector, server): add rich filter operators and fix /search 502 (#5263)

This commit is contained in:
Kartik
2026-05-26 22:47:41 +05:30
committed by GitHub
parent 7f6d46050e
commit ad736d9a06
9 changed files with 723 additions and 100 deletions
+15
View File
@@ -7,6 +7,14 @@ mode: "wide"
<Tabs>
<Tab title="Python">
<Update label="2026-05-26" description="v2.0.3">
**Bug Fixes:**
- **Vector Stores:** PGVector adapter now supports rich filter operators (`eq`, `ne`, `gt`, `gte`, `lt`, `lte`, `in`, `nin`, `contains`, `icontains`, wildcard `*`, `$or`, `$not`) in `search()`, `keyword_search()`, and `list()`. Previously only exact-equality filters worked — operator dicts were silently stringified and returned zero results ([#5263](https://github.com/mem0ai/mem0/pull/5263))
- **Server:** Fixed `/search` endpoint returning 502 when `user_id`, `agent_id`, or `run_id` are sent as top-level request fields. The server now maps these into the `filters` dict before calling `Memory.search()`, matching the v3 API contract. Top-level entity ID fields are marked as deprecated in the OpenAPI schema and emit a warning log — clients should migrate to `filters={"user_id": "..."}` ([#5263](https://github.com/mem0ai/mem0/pull/5263))
</Update>
<Update label="2026-05-08" description="v2.0.2">
**Bug Fixes:**
@@ -924,6 +932,13 @@ See the [OSS v1 to v2 migration guide](https://docs.mem0.ai/migration/oss-v1-to-
</Tab>
<Tab title="TypeScript">
<Update label="2026-05-26" description="v3.0.4">
**Bug Fixes:**
- **Vector Stores:** PGVector adapter now supports rich filter operators (`eq`, `ne`, `gt`, `gte`, `lt`, `lte`, `in`, `nin`, `contains`, `icontains`, wildcard `*`, `$or`, `$not`) in `search()`, `keywordSearch()`, and `list()`. Previously only exact-equality filters worked — operator objects were passed as raw values and returned incorrect results ([#5263](https://github.com/mem0ai/mem0/pull/5263))
</Update>
<Update label="2026-05-08" description="v3.0.3">
**Bug Fixes:**
+1 -1
View File
@@ -1,6 +1,6 @@
{
"name": "mem0ai",
"version": "3.0.3",
"version": "3.0.4",
"description": "The Memory Layer For Your AI Apps",
"main": "./dist/index.js",
"module": "./dist/index.mjs",
+136 -48
View File
@@ -28,6 +28,119 @@ function escapeFilterKey(key: string): string {
return key;
}
interface FilterResult {
conditions: string[];
values: any[];
paramIndex: number;
}
const OPERATOR_SQL_MAP: Record<string, { template: string; numeric: boolean }> =
{
eq: { template: "payload->>'%KEY%' = $%IDX%", numeric: false },
ne: { template: "payload->>'%KEY%' != $%IDX%", numeric: false },
gt: { template: "(payload->>'%KEY%')::numeric > $%IDX%", numeric: true },
gte: { template: "(payload->>'%KEY%')::numeric >= $%IDX%", numeric: true },
lt: { template: "(payload->>'%KEY%')::numeric < $%IDX%", numeric: true },
lte: { template: "(payload->>'%KEY%')::numeric <= $%IDX%", numeric: true },
in: { template: "payload->>'%KEY%' = ANY($%IDX%::text[])", numeric: false },
nin: {
template: "NOT (payload->>'%KEY%' = ANY($%IDX%::text[]))",
numeric: false,
},
contains: { template: "payload->>'%KEY%' LIKE $%IDX%", numeric: false },
icontains: { template: "payload->>'%KEY%' ILIKE $%IDX%", numeric: false },
};
export function buildFilterConditions(
filters: Record<string, any> | undefined,
startIndex: number,
): FilterResult {
const conditions: string[] = [];
const values: any[] = [];
let paramIndex = startIndex;
if (!filters) {
return { conditions, values, paramIndex };
}
for (const [key, value] of Object.entries(filters)) {
if (key === "$or") {
const orGroups: string[] = [];
for (const orFilter of value as Record<string, any>[]) {
const sub = buildFilterConditions(orFilter, paramIndex);
if (sub.conditions.length > 0) {
orGroups.push("(" + sub.conditions.join(" AND ") + ")");
values.push(...sub.values);
paramIndex = sub.paramIndex;
}
}
if (orGroups.length > 0) {
conditions.push("(" + orGroups.join(" OR ") + ")");
}
continue;
}
if (key === "$not") {
const notGroups: string[] = [];
for (const notFilter of value as Record<string, any>[]) {
const sub = buildFilterConditions(notFilter, paramIndex);
if (sub.conditions.length > 0) {
notGroups.push("(" + sub.conditions.join(" AND ") + ")");
values.push(...sub.values);
paramIndex = sub.paramIndex;
}
}
if (notGroups.length > 0) {
conditions.push("NOT (" + notGroups.join(" OR ") + ")");
}
continue;
}
const safeKey = escapeFilterKey(key);
if (value === "*") {
conditions.push(`payload ? $${paramIndex}`);
values.push(key);
paramIndex++;
continue;
}
if (typeof value === "object" && value !== null && !Array.isArray(value)) {
for (const [op, opValue] of Object.entries(value)) {
const mapping = OPERATOR_SQL_MAP[op];
if (!mapping) {
throw new Error(`Unsupported filter operator: ${op}`);
}
const clause = mapping.template
.replace("%KEY%", safeKey)
.replace("%IDX%", String(paramIndex));
conditions.push(clause);
if (op === "in" || op === "nin") {
values.push((opValue as any[]).map(String));
} else if (op === "contains" || op === "icontains") {
values.push(`%${opValue}%`);
} else if (mapping.numeric) {
values.push(Number(opValue));
} else {
values.push(String(opValue));
}
paramIndex++;
}
} else if (Array.isArray(value)) {
conditions.push(`payload->>'${safeKey}' = ANY($${paramIndex}::text[])`);
values.push(value.map(String));
paramIndex++;
} else {
conditions.push(`payload->>'${safeKey}' = $${paramIndex}`);
values.push(value);
paramIndex++;
}
}
return { conditions, values, paramIndex };
}
interface PGVectorConfig extends VectorStoreConfig {
dbname?: string;
user: string;
@@ -203,23 +316,15 @@ export class PGVector implements VectorStore {
filters?: SearchFilters,
): Promise<VectorStoreResult[] | null> {
try {
const filterConditions: string[] = [];
const filterValues: any[] = [query, topK];
let filterIndex = 3;
if (filters) {
for (const [key, value] of Object.entries(filters)) {
const safeKey = escapeFilterKey(key);
filterConditions.push(`payload->>'${safeKey}' = $${filterIndex}`);
filterValues.push(value);
filterIndex++;
}
}
const {
conditions,
values,
paramIndex: _,
} = buildFilterConditions(filters, 3);
const filterValues: any[] = [query, topK, ...values];
const filterClause =
filterConditions.length > 0
? "AND " + filterConditions.join(" AND ")
: "";
conditions.length > 0 ? "AND " + conditions.join(" AND ") : "";
const searchQuery = `
SELECT id, ts_rank_cd(to_tsvector('simple', payload->>'textLemmatized'), plainto_tsquery('simple', $1)) AS score, payload
@@ -248,24 +353,16 @@ export class PGVector implements VectorStore {
topK: number = 5,
filters?: SearchFilters,
): Promise<VectorStoreResult[]> {
const filterConditions: string[] = [];
const queryVector = `[${query.join(",")}]`;
const filterValues: any[] = [queryVector, topK];
let filterIndex = 3;
if (filters) {
for (const [key, value] of Object.entries(filters)) {
const safeKey = escapeFilterKey(key);
filterConditions.push(`payload->>'${safeKey}' = $${filterIndex}`);
filterValues.push(value);
filterIndex++;
}
}
const {
conditions,
values,
paramIndex: _,
} = buildFilterConditions(filters, 3);
const filterValues: any[] = [queryVector, topK, ...values];
const filterClause =
filterConditions.length > 0
? "WHERE " + filterConditions.join(" AND ")
: "";
conditions.length > 0 ? "WHERE " + conditions.join(" AND ") : "";
const searchQuery = `
SELECT id, vector <=> $1::vector AS distance, payload
@@ -337,23 +434,14 @@ export class PGVector implements VectorStore {
filters?: SearchFilters,
topK: number = 100,
): Promise<[VectorStoreResult[], number]> {
const filterConditions: string[] = [];
const filterValues: any[] = [];
let paramIndex = 1;
if (filters) {
for (const [key, value] of Object.entries(filters)) {
const safeKey = escapeFilterKey(key);
filterConditions.push(`payload->>'${safeKey}' = $${paramIndex}`);
filterValues.push(value);
paramIndex++;
}
}
const {
conditions,
values: filterValues,
paramIndex,
} = buildFilterConditions(filters, 1);
const filterClause =
filterConditions.length > 0
? "WHERE " + filterConditions.join(" AND ")
: "";
conditions.length > 0 ? "WHERE " + conditions.join(" AND ") : "";
const listQuery = `
SELECT id, payload
@@ -368,11 +456,11 @@ export class PGVector implements VectorStore {
${filterClause}
`;
filterValues.push(topK); // Add limit as the last parameter
const listValues = [...filterValues, topK];
const [listResult, countResult] = await Promise.all([
this.client.query(listQuery, filterValues),
this.client.query(countQuery, filterValues.slice(0, -1)), // Remove limit parameter for count query
this.client.query(listQuery, listValues),
this.client.query(countQuery, filterValues),
]);
const results = listResult.rows.map((row) => ({
@@ -0,0 +1,219 @@
/// <reference types="jest" />
jest.mock("pg", () => {
const Client = jest.fn().mockImplementation(() => ({
connect: jest.fn().mockResolvedValue(undefined),
end: jest.fn().mockResolvedValue(undefined),
query: jest.fn().mockResolvedValue({ rows: [] }),
}));
const escapeIdentifier = (str: string) => `"${str.replace(/"/g, '""')}"`;
return {
__esModule: true,
default: { Client, escapeIdentifier },
Client,
escapeIdentifier,
};
});
import { buildFilterConditions } from "../src/vector_stores/pgvector";
describe("buildFilterConditions", () => {
test("returns empty for undefined filters", () => {
const result = buildFilterConditions(undefined, 1);
expect(result.conditions).toEqual([]);
expect(result.values).toEqual([]);
expect(result.paramIndex).toBe(1);
});
test("returns empty for empty filters", () => {
const result = buildFilterConditions({}, 1);
expect(result.conditions).toEqual([]);
expect(result.values).toEqual([]);
});
test("simple equality", () => {
const result = buildFilterConditions({ user_id: "alice" }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("payload->>'user_id' = $1");
expect(result.values).toEqual(["alice"]);
expect(result.paramIndex).toBe(2);
});
test("multiple equalities", () => {
const result = buildFilterConditions(
{ user_id: "alice", agent_id: "bot1" },
1,
);
expect(result.conditions).toHaveLength(2);
expect(result.values).toEqual(["alice", "bot1"]);
expect(result.paramIndex).toBe(3);
});
test("eq operator", () => {
const result = buildFilterConditions({ status: { eq: "active" } }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("= $1");
expect(result.values).toEqual(["active"]);
});
test("ne operator", () => {
const result = buildFilterConditions({ status: { ne: "deleted" } }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("!= $1");
expect(result.values).toEqual(["deleted"]);
});
test("gt operator", () => {
const result = buildFilterConditions({ price: { gt: 100 } }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("::numeric > $1");
expect(result.values).toEqual([100]);
});
test("gte operator", () => {
const result = buildFilterConditions({ price: { gte: 100 } }, 1);
expect(result.conditions[0]).toContain("::numeric >= $1");
expect(result.values).toEqual([100]);
});
test("lt operator", () => {
const result = buildFilterConditions({ price: { lt: 50 } }, 1);
expect(result.conditions[0]).toContain("::numeric < $1");
expect(result.values).toEqual([50]);
});
test("lte operator", () => {
const result = buildFilterConditions({ price: { lte: 50 } }, 1);
expect(result.conditions[0]).toContain("::numeric <= $1");
expect(result.values).toEqual([50]);
});
test("range combination (gte + lte)", () => {
const result = buildFilterConditions({ score: { gte: 1, lte: 10 } }, 1);
expect(result.conditions).toHaveLength(2);
expect(result.conditions[0]).toContain("::numeric >= $1");
expect(result.conditions[1]).toContain("::numeric <= $2");
expect(result.values).toEqual([1, 10]);
});
test("in operator", () => {
const result = buildFilterConditions(
{ status: { in: ["active", "pending"] } },
1,
);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("= ANY($1::text[])");
expect(result.values).toEqual([["active", "pending"]]);
});
test("nin operator", () => {
const result = buildFilterConditions(
{ status: { nin: ["deleted", "archived"] } },
1,
);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("NOT");
expect(result.conditions[0]).toContain("= ANY($1::text[])");
expect(result.values).toEqual([["deleted", "archived"]]);
});
test("contains operator", () => {
const result = buildFilterConditions({ name: { contains: "alice" } }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("LIKE $1");
expect(result.values).toEqual(["%alice%"]);
});
test("icontains operator", () => {
const result = buildFilterConditions({ name: { icontains: "Alice" } }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("ILIKE $1");
expect(result.values).toEqual(["%Alice%"]);
});
test("wildcard value", () => {
const result = buildFilterConditions({ metadata_key: "*" }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("payload ? $1");
expect(result.values).toEqual(["metadata_key"]);
});
test("list shorthand (array value)", () => {
const result = buildFilterConditions({ tags: ["a", "b", "c"] }, 1);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain("= ANY($1::text[])");
expect(result.values).toEqual([["a", "b", "c"]]);
});
test("$or operator", () => {
const result = buildFilterConditions(
{
$or: [{ user_id: "alice" }, { user_id: "bob" }],
},
1,
);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain(" OR ");
expect(result.conditions[0]).toMatch(/^\(/);
expect(result.values).toEqual(["alice", "bob"]);
});
test("$not operator", () => {
const result = buildFilterConditions(
{
$not: [{ status: "deleted" }],
},
1,
);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toMatch(/^NOT/);
expect(result.values).toEqual(["deleted"]);
});
test("$or with operators", () => {
const result = buildFilterConditions(
{
$or: [{ price: { gt: 100 } }, { price: { lt: 10 } }],
},
1,
);
expect(result.conditions).toHaveLength(1);
expect(result.conditions[0]).toContain(" OR ");
expect(result.values).toEqual([100, 10]);
});
test("mixed simple and operator filters", () => {
const result = buildFilterConditions(
{
user_id: "alice",
score: { gte: 5 },
},
1,
);
expect(result.conditions).toHaveLength(2);
expect(result.values[0]).toBe("alice");
expect(result.values[1]).toBe(5);
});
test("unsupported operator throws", () => {
expect(() => buildFilterConditions({ x: { badop: 1 } }, 1)).toThrow(
"Unsupported filter operator",
);
});
test("in with numeric values converts to strings", () => {
const result = buildFilterConditions({ priority: { in: [1, 2, 3] } }, 1);
expect(result.values).toEqual([["1", "2", "3"]]);
});
test("paramIndex increments correctly across multiple fields", () => {
const result = buildFilterConditions(
{ a: "x", b: { gt: 5 }, c: { in: [1, 2] } },
3,
);
expect(result.conditions[0]).toContain("$3");
expect(result.conditions[1]).toContain("$4");
expect(result.conditions[2]).toContain("$5");
expect(result.paramIndex).toBe(6);
});
});
+80 -24
View File
@@ -31,6 +31,83 @@ from mem0.vector_stores.base import VectorStoreBase
logger = logging.getLogger(__name__)
OPERATOR_SQL_MAP = {
"eq": ("payload->>%s = %s", False),
"ne": ("payload->>%s != %s", False),
"gt": ("(payload->>%s)::numeric > %s", True),
"gte": ("(payload->>%s)::numeric >= %s", True),
"lt": ("(payload->>%s)::numeric < %s", True),
"lte": ("(payload->>%s)::numeric <= %s", True),
"in": ("payload->>%s = ANY(%s)", False),
"nin": ("NOT (payload->>%s = ANY(%s))", False),
"contains": ("payload->>%s LIKE %s", False),
"icontains": ("payload->>%s ILIKE %s", False),
}
def _build_filter_conditions(filters):
"""Translate a processed filter dict into SQL WHERE fragments and parameter list."""
conditions = []
params = []
if not filters:
return conditions, params
for key, value in filters.items():
if key == "$or":
or_groups = []
for or_filter in value:
sub_conds, sub_params = _build_filter_conditions(or_filter)
if sub_conds:
or_groups.append("(" + " AND ".join(sub_conds) + ")")
params.extend(sub_params)
if or_groups:
conditions.append("(" + " OR ".join(or_groups) + ")")
continue
if key == "$not":
not_groups = []
for not_filter in value:
sub_conds, sub_params = _build_filter_conditions(not_filter)
if sub_conds:
not_groups.append("(" + " AND ".join(sub_conds) + ")")
params.extend(sub_params)
if not_groups:
conditions.append("NOT (" + " OR ".join(not_groups) + ")")
continue
if value == "*":
conditions.append("payload ? %s")
params.append(key)
continue
if isinstance(value, dict):
for op, op_value in value.items():
if op not in OPERATOR_SQL_MAP:
raise ValueError(f"Unsupported filter operator: {op}")
template, is_numeric = OPERATOR_SQL_MAP[op]
if op in ("in", "nin"):
str_list = [str(v) for v in op_value]
conditions.append(template)
params.extend([key, str_list])
elif op in ("contains", "icontains"):
conditions.append(template)
params.extend([key, f"%{op_value}%"])
else:
conditions.append(template)
if is_numeric:
params.extend([key, float(op_value)])
else:
params.extend([key, str(op_value)])
elif isinstance(value, list):
conditions.append("payload->>%s = ANY(%s)")
params.extend([key, [str(v) for v in value]])
else:
conditions.append("payload->>%s = %s")
params.extend([key, str(value)])
return conditions, params
class OutputData(BaseModel):
id: Optional[str]
@@ -237,14 +314,7 @@ class PGVector(VectorStoreBase):
Returns:
list: Search results.
"""
filter_conditions = []
filter_params = []
if filters:
for k, v in filters.items():
filter_conditions.append("payload->>%s = %s")
filter_params.extend([k, str(v)])
filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
with self._get_cursor() as cur:
@@ -274,14 +344,7 @@ class PGVector(VectorStoreBase):
Returns:
List[OutputData]: Search results ranked by text relevance.
"""
filter_conditions = []
filter_params = []
if filters:
for k, v in filters.items():
filter_conditions.append("payload->>%s = %s")
filter_params.extend([k, str(v)])
filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("AND " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
try:
@@ -423,14 +486,7 @@ class PGVector(VectorStoreBase):
Returns:
List[OutputData]: List of vectors.
"""
filter_conditions = []
filter_params = []
if filters:
for k, v in filters.items():
filter_conditions.append("payload->>%s = %s")
filter_params.extend([k, str(v)])
filter_conditions, filter_params = _build_filter_conditions(filters)
filter_clause = sql.SQL("WHERE " + " AND ".join(filter_conditions)) if filter_conditions else sql.SQL("")
with self._get_cursor() as cur:
+1 -1
View File
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "mem0ai"
version = "2.0.2"
version = "2.0.3"
description = "Long-term memory for AI Agents"
authors = [
{ name = "Mem0", email = "support@mem0.ai" }
+41 -19
View File
@@ -4,16 +4,10 @@ import os
import time
from typing import Any, Dict, List, Optional
from dotenv import load_dotenv
from fastapi import Depends, FastAPI, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, RedirectResponse
from pydantic import BaseModel, Field
from slowapi import _rate_limit_exceeded_handler
from slowapi.errors import RateLimitExceeded
from sqlalchemy import func, select
import telemetry
from auth import ADMIN_API_KEY, AUTH_DISABLED, JWT_SECRET, verify_auth
from db import SessionLocal
from dotenv import load_dotenv
from errors import (
UpstreamError,
install_request_id_logging,
@@ -22,16 +16,27 @@ from errors import (
upstream_error,
upstream_error_handler,
)
from rate_limit import limiter
from db import SessionLocal
from fastapi import Depends, FastAPI, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse, RedirectResponse
from models import RequestLog, User
import telemetry
from routers import auth as auth_router
from pydantic import BaseModel, Field
from rate_limit import limiter
from routers import api_keys as api_keys_router
from routers import auth as auth_router
from routers import entities as entities_router
from routers import requests as requests_router
from schemas import MessageResponse
from server_state import get_current_config, get_memory_instance, initialize_state, set_session_factory, update_config
from server_state import (
get_current_config,
get_memory_instance,
initialize_state,
set_session_factory,
update_config,
)
from slowapi import _rate_limit_exceeded_handler
from slowapi.errors import RateLimitExceeded
from sqlalchemy import func, select
load_dotenv()
@@ -188,9 +193,9 @@ class MemoryUpdate(BaseModel):
class SearchRequest(BaseModel):
query: str = Field(..., description="Search query.")
user_id: Optional[str] = None
run_id: Optional[str] = None
agent_id: Optional[str] = None
user_id: Optional[str] = Field(None, description="Deprecated: pass inside `filters` instead.", deprecated=True)
run_id: Optional[str] = Field(None, description="Deprecated: pass inside `filters` instead.", deprecated=True)
agent_id: Optional[str] = Field(None, description="Deprecated: pass inside `filters` instead.", deprecated=True)
filters: Optional[Dict[str, Any]] = None
top_k: Optional[int] = Field(None, description="Maximum number of results to return.")
threshold: Optional[float] = Field(None, description="Minimum similarity score for results.")
@@ -416,8 +421,25 @@ def get_memory(memory_id: str, _auth=Depends(verify_auth)):
def search_memories(search_req: SearchRequest, _auth=Depends(verify_auth)):
"""Search for memories based on a query."""
try:
params = {k: v for k, v in search_req.model_dump().items() if v is not None and k != "query"}
return get_memory_instance().search(query=search_req.query, **params)
filters = search_req.filters or {}
deprecated_keys = []
for entity_key in ("user_id", "agent_id", "run_id"):
entity_val = getattr(search_req, entity_key, None)
if entity_val is not None:
filters[entity_key] = entity_val
deprecated_keys.append(entity_key)
if deprecated_keys:
logging.warning(
"Top-level %s in /search is deprecated. Use filters={%s} instead.",
", ".join(deprecated_keys),
", ".join(f'"{k}": "..."' for k in deprecated_keys),
)
params = {}
if search_req.top_k is not None:
params["top_k"] = search_req.top_k
if search_req.threshold is not None:
params["threshold"] = search_req.threshold
return get_memory_instance().search(query=search_req.query, filters=filters, **params)
except Exception:
raise upstream_error()
+74 -6
View File
@@ -16,7 +16,6 @@ pytest.importorskip("fastapi", reason="fastapi not installed")
from fastapi.testclient import TestClient
# ---------------------------------------------------------------------------
# Fixtures
# ---------------------------------------------------------------------------
@@ -306,9 +305,9 @@ class TestExistingParamsUnchanged:
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["user_id"] == "u1"
assert kwargs["agent_id"] == "a1"
assert kwargs["filters"] == {"category": "food"}
assert kwargs["filters"]["user_id"] == "u1"
assert kwargs["filters"]["agent_id"] == "a1"
assert kwargs["filters"]["category"] == "food"
def test_add_metadata_still_forwarded(self, client, mock_memory):
resp = client.post("/memories", json={
@@ -465,11 +464,14 @@ class TestCallSignatureMatch:
"top_k": 10, "threshold": 0.5,
})
assert resp.status_code == 200
# The handler passes query= as a keyword arg, so it appears in kwargs too
_, kwargs = mock_memory.search.call_args
valid_params = {"query", "user_id", "agent_id", "run_id", "top_k", "filters", "threshold", "rerank"}
valid_params = {"query", "top_k", "filters", "threshold", "rerank"}
for key in kwargs:
assert key in valid_params, f"Unexpected kwarg '{key}' forwarded to Memory.search()"
assert kwargs["filters"]["user_id"] == "u1"
assert kwargs["filters"]["agent_id"] == "a1"
assert kwargs["filters"]["run_id"] == "r1"
assert kwargs["filters"]["k"] == "v"
def test_add_kwargs_are_valid(self, client, mock_memory):
"""All kwargs forwarded to Memory.add() must be in its signature."""
@@ -588,3 +590,69 @@ class TestGetMemories:
# 3. Verify the core logic: the param was mapped to the filters dict!
_, kwargs = mock_memory.get_all.call_args
assert kwargs["filters"] == {"user_id": "test_routing_user"}
# ===========================================================================
# SearchRequest: entity IDs mapped into filters (fix for server 502)
# ===========================================================================
class TestSearchEntityIdMapping:
"""Verify that POST /search maps top-level user_id / agent_id / run_id
into the filters dict instead of forwarding them as kwargs, which would
cause Memory.search() to raise ValueError in v3."""
def test_user_id_mapped_to_filters(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "user_id": "u1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "user_id" not in kwargs
assert kwargs["filters"]["user_id"] == "u1"
def test_agent_id_mapped_to_filters(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "agent_id": "a1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "agent_id" not in kwargs
assert kwargs["filters"]["agent_id"] == "a1"
def test_run_id_mapped_to_filters(self, client, mock_memory):
resp = client.post("/search", json={"query": "food", "run_id": "r1"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert "run_id" not in kwargs
assert kwargs["filters"]["run_id"] == "r1"
def test_all_entity_ids_mapped(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food", "user_id": "u1", "agent_id": "a1", "run_id": "r1",
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["filters"] == {"user_id": "u1", "agent_id": "a1", "run_id": "r1"}
def test_entity_ids_merged_with_explicit_filters(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food",
"user_id": "u1",
"filters": {"category": "food"},
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["filters"]["user_id"] == "u1"
assert kwargs["filters"]["category"] == "food"
def test_no_entity_ids_no_filters(self, client, mock_memory):
resp = client.post("/search", json={"query": "food"})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["filters"] == {}
def test_only_filters_no_entity_ids(self, client, mock_memory):
resp = client.post("/search", json={
"query": "food",
"filters": {"user_id": "u1", "category": "food"},
})
assert resp.status_code == 200
_, kwargs = mock_memory.search.call_args
assert kwargs["filters"]["user_id"] == "u1"
assert kwargs["filters"]["category"] == "food"
+156 -1
View File
@@ -4,7 +4,7 @@ import unittest
import uuid
from unittest.mock import MagicMock, patch
from mem0.vector_stores.pgvector import PGVector
from mem0.vector_stores.pgvector import PGVector, _build_filter_conditions
class TestPGVector(unittest.TestCase):
@@ -2233,3 +2233,158 @@ class TestPGVector(unittest.TestCase):
def tearDown(self):
"""Clean up after each test."""
pass
class TestBuildFilterConditions(unittest.TestCase):
"""Tests for the _build_filter_conditions helper that translates filter dicts to SQL."""
def test_none_filters(self):
conditions, params = _build_filter_conditions(None)
self.assertEqual(conditions, [])
self.assertEqual(params, [])
def test_empty_filters(self):
conditions, params = _build_filter_conditions({})
self.assertEqual(conditions, [])
self.assertEqual(params, [])
def test_simple_equality(self):
conditions, params = _build_filter_conditions({"user_id": "alice"})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s = %s", conditions[0])
self.assertEqual(params, ["user_id", "alice"])
def test_multiple_equalities(self):
conditions, params = _build_filter_conditions({"user_id": "alice", "agent_id": "bot1"})
self.assertEqual(len(conditions), 2)
self.assertEqual(params, ["user_id", "alice", "agent_id", "bot1"])
def test_eq_operator(self):
conditions, params = _build_filter_conditions({"status": {"eq": "active"}})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s = %s", conditions[0])
self.assertEqual(params, ["status", "active"])
def test_ne_operator(self):
conditions, params = _build_filter_conditions({"status": {"ne": "deleted"}})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s != %s", conditions[0])
self.assertEqual(params, ["status", "deleted"])
def test_gt_operator(self):
conditions, params = _build_filter_conditions({"price": {"gt": 100}})
self.assertEqual(len(conditions), 1)
self.assertIn("(payload->>%s)::numeric > %s", conditions[0])
self.assertEqual(params, ["price", 100.0])
def test_gte_operator(self):
conditions, params = _build_filter_conditions({"price": {"gte": 100}})
self.assertEqual(len(conditions), 1)
self.assertIn("(payload->>%s)::numeric >= %s", conditions[0])
self.assertEqual(params, ["price", 100.0])
def test_lt_operator(self):
conditions, params = _build_filter_conditions({"price": {"lt": 50}})
self.assertEqual(len(conditions), 1)
self.assertIn("(payload->>%s)::numeric < %s", conditions[0])
self.assertEqual(params, ["price", 50.0])
def test_lte_operator(self):
conditions, params = _build_filter_conditions({"price": {"lte": 50}})
self.assertEqual(len(conditions), 1)
self.assertIn("(payload->>%s)::numeric <= %s", conditions[0])
self.assertEqual(params, ["price", 50.0])
def test_range_combination(self):
conditions, params = _build_filter_conditions({"score": {"gte": 1, "lte": 10}})
self.assertEqual(len(conditions), 2)
self.assertIn("(payload->>%s)::numeric >= %s", conditions[0])
self.assertIn("(payload->>%s)::numeric <= %s", conditions[1])
self.assertEqual(params, ["score", 1.0, "score", 10.0])
def test_in_operator(self):
conditions, params = _build_filter_conditions({"status": {"in": ["active", "pending"]}})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s = ANY(%s)", conditions[0])
self.assertEqual(params, ["status", ["active", "pending"]])
def test_nin_operator(self):
conditions, params = _build_filter_conditions({"status": {"nin": ["deleted", "archived"]}})
self.assertEqual(len(conditions), 1)
self.assertIn("NOT (payload->>%s = ANY(%s))", conditions[0])
self.assertEqual(params, ["status", ["deleted", "archived"]])
def test_contains_operator(self):
conditions, params = _build_filter_conditions({"name": {"contains": "alice"}})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s LIKE %s", conditions[0])
self.assertEqual(params, ["name", "%alice%"])
def test_icontains_operator(self):
conditions, params = _build_filter_conditions({"name": {"icontains": "Alice"}})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s ILIKE %s", conditions[0])
self.assertEqual(params, ["name", "%Alice%"])
def test_wildcard(self):
conditions, params = _build_filter_conditions({"metadata_key": "*"})
self.assertEqual(len(conditions), 1)
self.assertIn("payload ? %s", conditions[0])
self.assertEqual(params, ["metadata_key"])
def test_list_shorthand(self):
conditions, params = _build_filter_conditions({"tags": ["a", "b", "c"]})
self.assertEqual(len(conditions), 1)
self.assertIn("payload->>%s = ANY(%s)", conditions[0])
self.assertEqual(params, ["tags", ["a", "b", "c"]])
def test_or_operator(self):
conditions, params = _build_filter_conditions({
"$or": [
{"user_id": "alice"},
{"user_id": "bob"},
]
})
self.assertEqual(len(conditions), 1)
self.assertIn(" OR ", conditions[0])
self.assertTrue(conditions[0].startswith("("))
self.assertEqual(params, ["user_id", "alice", "user_id", "bob"])
def test_not_operator(self):
conditions, params = _build_filter_conditions({
"$not": [
{"status": "deleted"},
]
})
self.assertEqual(len(conditions), 1)
self.assertTrue(conditions[0].startswith("NOT"))
self.assertEqual(params, ["status", "deleted"])
def test_or_with_operators(self):
conditions, params = _build_filter_conditions({
"$or": [
{"price": {"gt": 100}},
{"price": {"lt": 10}},
]
})
self.assertEqual(len(conditions), 1)
self.assertIn(" OR ", conditions[0])
self.assertEqual(params, ["price", 100.0, "price", 10.0])
def test_mixed_simple_and_operator_filters(self):
conditions, params = _build_filter_conditions({
"user_id": "alice",
"score": {"gte": 5},
})
self.assertEqual(len(conditions), 2)
self.assertIn("payload->>%s = %s", conditions[0])
self.assertIn("(payload->>%s)::numeric >= %s", conditions[1])
def test_unsupported_operator_raises(self):
with self.assertRaises(ValueError) as ctx:
_build_filter_conditions({"x": {"badop": 1}})
self.assertIn("Unsupported filter operator", str(ctx.exception))
def test_in_with_numeric_values(self):
conditions, params = _build_filter_conditions({"priority": {"in": [1, 2, 3]}})
self.assertEqual(params, ["priority", ["1", "2", "3"]])